From 040ffaf1271811617548474a0170471f991f1465 Mon Sep 17 00:00:00 2001 From: Haoran Qian Date: Fri, 17 Jul 2026 19:48:17 +0800 Subject: [PATCH] Fix FlexAttention text mask length Signed-off-by: Haoran Qian --- wan_va/modules/model.py | 7 ++++++- 1 file changed, 6 insertions(+), 1 deletion(-) diff --git a/wan_va/modules/model.py b/wan_va/modules/model.py index 25b45e57..70bd331d 100644 --- a/wan_va/modules/model.py +++ b/wan_va/modules/model.py @@ -95,6 +95,7 @@ def half(x): def init_mask( latent_shape, action_shape, + text_sequence_length, padded_length, chunk_size, window_size, @@ -133,7 +134,9 @@ def init_mask( ) FlexAttnFunc.attention_mask = block_mask - text_seq_ids = torch.arange(B)[:, None].expand(-1, 512).flatten() + text_seq_ids = torch.arange(B)[:, None].expand( + -1, text_sequence_length + ).flatten() mask_mod_cross = FlexAttnFunc._get_cross_mask_mod(seq_ids.long().to(device), text_seq_ids.long().to(device)) block_mask_cross = FlexAttnFunc.compiled_create_block_mask( mask_mod_cross, 1, 1, len(seq_ids), len(text_seq_ids), device=device, _compile=True @@ -712,6 +715,7 @@ def forward_train(self, input_dict): latent_hidden_states = self._input_embed(latent_dict['noisy_latents'], input_type='latent').flatten(0, 1)[None] action_hidden_states = self._input_embed(action_dict['noisy_latents'], input_type='action').flatten(0, 1)[None] text_hidden_states = self._input_embed(latent_dict["text_emb"], input_type='text') + text_sequence_length = text_hidden_states.shape[1] text_hidden_states = text_hidden_states.flatten(0, 1)[None] @@ -764,6 +768,7 @@ def forward_train(self, input_dict): FlexAttnFunc.init_mask(latent_dict['noisy_latents'].shape, action_dict['noisy_latents'].shape, + text_sequence_length, padded_length, input_dict["chunk_size"], window_size=input_dict['window_size'],