| 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( |