| import torch.nn as nn |
|
|
| from .modules.fused_norm_gate import rms_norm_gated |
| from .ops.gated_delta_rule import chunk_gated_delta_rule as chunk_gdn |
| from .ops.gated_delta_rule import fused_recurrent_gated_delta_rule as fused_recurrent_gdn |
| from .ops.kda import chunk_kda, fused_recurrent_kda |
|
|
|
|
| class FusedRMSNormGated(nn.Module): |
| def forward(self, hidden_states, gate=None): |
| return rms_norm_gated( |
| hidden_states, |
| gate, |
| self.weight, |
| None, |
| self.activation, |
| residual=None, |
| eps=self.variance_epsilon, |
| prenorm=False, |
| residual_in_fp32=False, |
| ) |
|
|
|
|
| class chunk_kimi_delta_attention(nn.Module): |
| def forward( |
| self, |
| query, |
| key, |
| value, |
| g, |
| beta, |
| chunk_size=64, |
| initial_state=None, |
| output_final_state=False, |
| use_qk_l2norm_in_kernel=False, |
| **kwargs, |
| ): |
| |
| cu_seqlens = kwargs.pop("cu_seq_lens_q", kwargs.pop("cu_seqlens", None)) |
|
|
| return chunk_kda( |
| query, |
| key, |
| value, |
| g=g, |
| beta=beta, |
| chunk_size=chunk_size, |
| initial_state=initial_state, |
| output_final_state=output_final_state, |
| use_qk_l2norm_in_kernel=use_qk_l2norm_in_kernel, |
| cu_seqlens=cu_seqlens, |
| **kwargs, |
| ) |
|
|
|
|
| class recurrent_kimi_delta_attention(nn.Module): |
| def forward( |
| self, query, key, value, g, beta, initial_state, output_final_state, use_qk_l2norm_in_kernel=False, **kwargs, |
| ): |
| |
| cu_seqlens = kwargs.pop("cu_seq_lens_q", kwargs.pop("cu_seqlens", None)) |
|
|
| return fused_recurrent_kda( |
| query, |
| key, |
| value, |
| g=g, |
| beta=beta, |
| initial_state=initial_state, |
| output_final_state=output_final_state, |
| use_qk_l2norm_in_kernel=use_qk_l2norm_in_kernel, |
| cu_seqlens=cu_seqlens, |
| **kwargs, |
| ) |
|
|
|
|
| class chunk_gated_delta_rule(nn.Module): |
| def forward( |
| self, |
| query, |
| key, |
| value, |
| g, |
| beta, |
| chunk_size=64, |
| initial_state=None, |
| output_final_state=False, |
| use_qk_l2norm_in_kernel=False, |
| **kwargs, |
| ): |
| |
| cu_seqlens = kwargs.pop("cu_seq_lens_q", kwargs.pop("cu_seqlens", None)) |
|
|
| return chunk_gdn( |
| query, |
| key, |
| value, |
| g=g, |
| beta=beta, |
| chunk_size=chunk_size, |
| initial_state=initial_state, |
| output_final_state=output_final_state, |
| use_qk_l2norm_in_kernel=use_qk_l2norm_in_kernel, |
| cu_seqlens=cu_seqlens, |
| **kwargs, |
| ) |
|
|
|
|
| class recurrent_gated_delta_rule(nn.Module): |
| def forward( |
| self, query, key, value, g, beta, initial_state, output_final_state, use_qk_l2norm_in_kernel=False, **kwargs, |
| ): |
| |
| cu_seqlens = kwargs.pop("cu_seq_lens_q", kwargs.pop("cu_seqlens", None)) |
|
|
| return fused_recurrent_gdn( |
| query, |
| key, |
| value, |
| g=g, |
| beta=beta, |
| initial_state=initial_state, |
| output_final_state=output_final_state, |
| use_qk_l2norm_in_kernel=use_qk_l2norm_in_kernel, |
| cu_seqlens=cu_seqlens, |
| **kwargs, |
| ) |
|
|
|
|
| __all__ = [ |
| "FusedRMSNormGated", |
| "chunk_kimi_delta_attention", |
| "recurrent_kimi_delta_attention", |
| "chunk_gated_delta_rule", |
| "recurrent_gated_delta_rule" |
| ] |
|
|