Compare commits

...
Author SHA1 Message Date
Will Lin e4ceadb5d5 fix denoising stage init 2025-06-06 11:44:36 -07:00
+11 -8
View File
@@ -47,14 +47,17 @@ class DenoisingStage(PipelineStage):
super().__init__()
self.transformer = transformer
self.scheduler = scheduler
attn_head_size = self.transformer.hidden_size // self.transformer.num_attention_heads
self.attn_backend = get_attn_backend(
head_size=attn_head_size,
dtype=torch.float16, # TODO(will): hack
supported_attention_backends=(_Backend.SLIDING_TILE_ATTN,
_Backend.FLASH_ATTN,
_Backend.TORCH_SDPA) # hack
)
# when used for validation, transformer is None as it is taking from the
# training loop
if transformer is not None:
attn_head_size = self.transformer.hidden_size // self.transformer.num_attention_heads
self.attn_backend = get_attn_backend(
head_size=attn_head_size,
dtype=torch.float16, # TODO(will): hack
supported_attention_backends=(_Backend.SLIDING_TILE_ATTN,
_Backend.FLASH_ATTN,
_Backend.TORCH_SDPA) # hack
)
def forward(
self,