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"
]