MCPcopy Create free account
hub / github.com/OpenImagingLab/FlashVSR / prepare_attn_mask

Method prepare_attn_mask

diffsynth/models/stepvideo_dit.py:817–823  ·  view source on GitHub ↗
(self, encoder_attention_mask, encoder_hidden_states, q_seqlen)

Source from the content-addressed store, hash-verified

815 return hidden_states
816
817 def prepare_attn_mask(self, encoder_attention_mask, encoder_hidden_states, q_seqlen):
818 kv_seqlens = encoder_attention_mask.sum(dim=1).int()
819 mask = torch.zeros([len(kv_seqlens), q_seqlen, max(kv_seqlens)], dtype=torch.bool, device=encoder_attention_mask.device)
820 encoder_hidden_states = encoder_hidden_states[:,: max(kv_seqlens)]
821 for i, kv_len in enumerate(kv_seqlens):
822 mask[i, :, :kv_len] = 1
823 return encoder_hidden_states, mask
824
825
826 def block_forward(

Callers 1

forwardMethod · 0.95

Calls

no outgoing calls

Tested by

no test coverage detected