Kernels:
Trusted publisher
File size: 3,911 Bytes
e19323e 888b93f e19323e 888b93f | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 | 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, # bias
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,
):
# Keep internal consistency between transformers and fla to allow both
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,
):
# Keep internal consistency between transformers and fla to allow both
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,
):
# Keep internal consistency between transformers and fla to allow both
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,
):
# Keep internal consistency between transformers and fla to allow both
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"
]
|