From c19c5b4080941a588d985a7a5c41ed198c9bfab3 Mon Sep 17 00:00:00 2001 From: ciaranbor Date: Fri, 5 Dec 2025 19:39:29 +0000 Subject: [PATCH] Prune unused transformer blocks --- .../mflux/pipefusion/distributed_transformer.py | 13 +++++++++---- 1 file changed, 9 insertions(+), 4 deletions(-) diff --git a/src/exo/worker/engines/mflux/pipefusion/distributed_transformer.py b/src/exo/worker/engines/mflux/pipefusion/distributed_transformer.py index ca36665a..a2531809 100644 --- a/src/exo/worker/engines/mflux/pipefusion/distributed_transformer.py +++ b/src/exo/worker/engines/mflux/pipefusion/distributed_transformer.py @@ -82,6 +82,13 @@ class DistributedTransformer: self.has_single_blocks or self.end_layer == self.total_joint ) + self.transformer_blocks = self.transformer_blocks[ + self.joint_start : self.joint_end + ] + self.single_transformer_blocks = self.single_transformer_blocks[ + self.single_start : self.single_end + ] + @property def is_first_stage(self) -> bool: return self.rank == 0 @@ -126,8 +133,7 @@ class DistributedTransformer: ) # Run assigned joint blocks - for idx in range(self.joint_start, self.joint_end): - block = transformer.transformer_blocks[idx] + for block in transformer.transformer_blocks: encoder_hidden_states, hidden_states = block( hidden_states=hidden_states, encoder_hidden_states=encoder_hidden_states, @@ -167,8 +173,7 @@ class DistributedTransformer: mx.eval(hidden_states) # Run assigned single blocks - for idx in range(self.single_start, self.single_end): - block = transformer.single_transformer_blocks[idx] + for block in transformer.single_transformer_blocks: hidden_states = block( hidden_states=hidden_states, text_embeddings=text_embeddings,