Diff Coverage

Diff: origin/master...HEAD, staged and unstaged changes

Source File Diff Coverage (%) Missing Lines
hyper_parallel/components/functional/_triton/gated_delta_net/chunk_delta_h.py 0.0% 22,24-26,28,30,33,41,45-46,72-80,82-84,87-93,95-102,104-110,112-115,118-129,132,134-144,147-161,163-164,166-168,170,173-186,189-205,207,210-215,217-223,225-231,233-239,242-253,256,268-269,271,273-274,276-278,280-281,283,285-290,310-312,315,322,326-327,354-361,363-365,367-373,375-383,385-388,390-392,394-405,407-418,420-423,425-430,432-434,436,439-444,446-452,454-460,462-468,470-473,475,477-519,521-532,535,551,553-555,557-560,562,564-566,568,570,572,594
hyper_parallel/components/functional/_triton/gated_delta_net/chunk_o.py 0.0% 22,24-26,28,31,37-38,70-77,79-82,84-87,89-97,99-102,104-107,109-111,113-116,118-121,123-127,129-132,134-135,137-139,141-144,146-148,150,152-155,157-158,160-162,164-169,171-172,174-176,178-179,181,183-184,186-188,190-193,195-198,200,202-204,206-214,217-223,226,231-232,253-254,256-259,261,263-265,267-273,275-277,279,281-282,284-286,288-290,292-293,295,297-302,305,310-311,335-339,342-344,346-348,350-353,355-357,359-360,363-367,369-375,378,381,384-386,389,391,393-395,397-403,405-407,409,412,415,419-420,423,438-441,443-451,453-458,460-469,471,473,502-505,508,518-520,522-524,526-529,547,550,561-564,567-569,571-573,575,580-583,604,606-607
hyper_parallel/components/functional/_triton/gated_delta_net/chunk_scaled_dot_kkt.py 0.0% 22,24-26,28,31,35-36,54-56,58-59,61-63,65-66,68-70,72-77,79-81,83-87,89-90,92-96,98-100,102-107,109-115,117-118,121,124,131-132,148-149,151-155,157-158,160,162-165,167-168,170-172,174,176,178,181-189,191-192,195,198-199,214,216-220,222-223,225,227-232,234,236,238,240-241,243-244,246-249,251-252,254-256,259,291-298,300-302,318,320-325,340-341,355
hyper_parallel/components/functional/_triton/gated_delta_net/cumsum.py 0.0% 22,24-26,28,31,35-36,52-53,55-56,60,63,65,67,70,73-83,85-86,89,99-101,110-115,129,132,142-144,148-149,159
hyper_parallel/components/functional/_triton/gated_delta_net/solve_tril.py 0.0% 24-25,27-29,31,34,36-44,46-47,50-52,65-76,79,82,84,86-87,89,91-92,94-95,97-100,102-104,111-114,117,126-127,130-134,136-137,140,142-144,146-147,155-156,158,160-162,169-171,178-180,193-204,207,210,212-213,215-217,219,222,225,228,231,234,238-241,244,249,254,261,266-267,280-291,294,297,299-300,302-304,306,310,313,317,320,323,327-330,333,338,343,350-352,364-369,371-373,375-376,378-381,383-385,387-394,396,399,404-406,408-409,420,423-424,445-448,451-453,457,462,466,468,473-474,487-488,491,495,498-499,511-512,514,516,519,522-523,535-537
hyper_parallel/components/functional/_triton/gated_delta_net/state_summary.py 0.0% 23,25-26,28,31,40-41,56-59,61-67,69-74,77,81-82,85,88-91,93,101,103-112,114,117,120-121,123,131,139-140,143,152-153,171-174,176-184,186-187,189-195,197,200,203-207,209,217,219,227-230,232,235,238,241,244-254,256,259,262-263
hyper_parallel/components/functional/_triton/gated_delta_net/utils.py 0.0% 24-33,35-39,41,43,46,67-68,70-76,78-79,81-85,88-90,92-96,98,100-102,105-107,110-113,116-117,120-123,126-136,138-139,142,144-145,147,149,155-156,163,166-172,175-177,180-185,189-190,199,202,205-207,210-213,215-216,218-221,223-224,227,232-235,237-246,248-249,251,253-254,256,259-260,263-265,268-272,274-277,280-282,286-288,291-295,297-302,305-312,315,325-326,337,344-345,347,358-359,367-368,371-372,375-377,379-381,384-387,389-391
hyper_parallel/components/functional/_triton/gated_delta_net/wy_fast.py 0.0% 22,24-26,28,31,34-35,60-62,64-65,67-69,71-72,74-79,81-82,84-89,91,93-95,97-100,102-104,106-118,120-131,133-138,140-151,153-158,161,166-167,191-193,195-196,198-200,202-203,205-207,209-213,215-218,220-221,223-224,226-232,234-236,238-249,252,261-264,266-269,271-274,295,298,309-316,318-321,323-324,349-350,352,355,357
hyper_parallel/components/functional/_triton/kimi_delta_attention/state_summary.py 0.0% 22,27-28,31-32,34,37-38,57-60,62-67,69-73,75-83,85-87,92,98-100,108,116-119,121-122,124-125,133,138-139,141-146,148,156,164-167,169-171,176,181,186,193-194,213-216,218-226,228-231,233-235,237,245,253-256,258,266,268-270,275,280-281,283,291,293,301,309,317,325-332,334,342,350,355
hyper_parallel/components/functional/gated_delta_net.py 0.0% 29-35,53,61,70,131,149,159,374,707,722,727,737,882,895
hyper_parallel/components/functional/gated_delta_net_state_summary.py 0.0% 19,31,60,108,127,136,139-140,152
hyper_parallel/components/functional/kimi_delta_attention.py 0.0% 20,36,44,60,63,76,461,485
hyper_parallel/components/functional/kimi_delta_attention_fla_adapter.py 0.0% 20
hyper_parallel/components/functional/kimi_delta_attention_state_summary.py 0.0% 19,40,153,214
hyper_parallel/components/modules/gated_delta_net.py 14.3% 143-147,159
hyper_parallel/components/modules/kimi_delta_attention.py 33.3% 61-63,106,109-110
hyper_parallel/distributed/context_parallel/gated_delta_net.py 100%  
hyper_parallel/distributed/context_parallel/kimi_delta_attention.py 14.3% 217-219,221-223,232-234,252,274,570
hyper_parallel/components/functional/_triton/gated_delta_net/chunk_delta_h.py
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
# pylint: disable=used-before-assignment,unsupported-binary-operation,unused-argument
# pylint: disable=invalid-name,missing-module-docstring,missing-function-docstring
# pylint: disable=forbidden-backend-import

from typing import Optional, Tuple

import torch
import triton
import triton.language as tl

from .utils import prepare_chunk_indices, prepare_chunk_offsets, get_autotune_config, get_npu_properties

CUBE_CORE_NUM = get_npu_properties()['num_aicore']


@triton.heuristics({
    'USE_G': lambda args: args['g'] is not None,
    'USE_GK': lambda args: args['gk'] is not None,
    'USE_INITIAL_STATE': lambda args: args['h0'] is not None,
    'STORE_FINAL_STATE': lambda args: args['ht'] is not None,
37
38
39
40
41
42
43
44
45
46
47
48
49
50
    'STORE_FINAL_STATE': lambda args: args['ht'] is not None,
    'SAVE_NEW_VALUE': lambda args: args['v_new'] is not None,
    'IS_VARLEN': lambda args: args['cu_seqlens'] is not None,
})
@triton.autotune(
    configs=get_autotune_config(multibuffer_list=(False,)),
    key=['H', 'K', 'V', 'BT'],
)
@triton.jit(do_not_specialize=['T'])
def chunk_gated_delta_rule_fwd_kernel_h_blockdim64(
        k,
        v,
        w,
        v_new,
 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
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
        STORE_FINAL_STATE: tl.constexpr,
        SAVE_NEW_VALUE: tl.constexpr,
        IS_VARLEN: tl.constexpr,
):
    T_all = T
    NT_all = NT
    i_v, i_nh = tl.program_id(0), tl.program_id(1)
    i_n, i_h = i_nh // H, i_nh % H
    if IS_VARLEN:
        bos, eos = tl.load(cu_seqlens + i_n).to(tl.int32), tl.load(cu_seqlens + i_n + 1).to(tl.int32)
        T = eos - bos
        NT = tl.cdiv(T, BT)
        boh = tl.load(chunk_offsets + i_n).to(tl.int32)
    else:
        bos, eos = i_n * T, i_n * T + T
        NT = tl.cdiv(T, BT)
        boh = i_n * NT

    # Initialize hidden states
    b_h1 = tl.zeros([64, BV], dtype=tl.float32)
    if K > 64:
        b_h2 = tl.zeros([64, BV], dtype=tl.float32)
    if K > 128:
        b_h3 = tl.zeros([64, BV], dtype=tl.float32)
    if K > 192:
        b_h4 = tl.zeros([64, BV], dtype=tl.float32)

    if IS_VARLEN:
        v = v + (i_h * T_all + bos) * V
        k = k + (i_h * T_all + bos) * K
        w = w + (i_h * T_all + bos) * K
        g = g + i_h * T_all + bos
        h = h + (i_h * NT_all + boh) * K * V
        if SAVE_NEW_VALUE:
            v_new_base = v_new + (i_h * T_all + bos) * V
    else:
        v = v + (i_n * H + i_h) * T * V
        k = k + (i_n * H + i_h) * T * K
        w = w + (i_n * H + i_h) * T * K
        g = g + (i_n * H + i_h) * T
        h = h + (i_n * H + i_h) * NT * K * V
        if SAVE_NEW_VALUE:
            v_new_base = v_new + (i_n * H + i_h) * T * V

    if USE_INITIAL_STATE:
        h0_ptr = h0 + i_nh * K * V
    if STORE_FINAL_STATE:
        ht_ptr = ht + i_nh * K * V

    # Load initial state
    if USE_INITIAL_STATE:
        p_h0_1 = tl.make_block_ptr(h0_ptr, (K, V), (V, 1), (0, i_v * BV), (64, BV), (1, 0))
        b_h1 += tl.load(p_h0_1, boundary_check=(0, 1)).to(tl.float32)
        if K > 64:
            p_h0_2 = tl.make_block_ptr(h0_ptr, (K, V), (V, 1), (64, i_v * BV), (64, BV), (1, 0))
            b_h2 += tl.load(p_h0_2, boundary_check=(0, 1)).to(tl.float32)
        if K > 128:
            p_h0_3 = tl.make_block_ptr(h0_ptr, (K, V), (V, 1), (128, i_v * BV), (64, BV), (1, 0))
            b_h3 += tl.load(p_h0_3, boundary_check=(0, 1)).to(tl.float32)
        if K > 192:
            p_h0_4 = tl.make_block_ptr(h0_ptr, (K, V), (V, 1), (192, i_v * BV), (64, BV), (1, 0))
            b_h4 += tl.load(p_h0_4, boundary_check=(0, 1)).to(tl.float32)

    # Main recurrence over chunks
    for i_t in range(NT):
        # Store current hidden state h_t
        p_h1 = tl.make_block_ptr(h + i_t * K * V, (K, V), (V, 1), (0, i_v * BV), (64, BV), (1, 0))
        tl.store(p_h1, b_h1.to(p_h1.dtype.element_ty), boundary_check=(0, 1))
        if K > 64:
            p_h2 = tl.make_block_ptr(h + i_t * K * V, (K, V), (V, 1), (64, i_v * BV), (64, BV), (1, 0))
            tl.store(p_h2, b_h2.to(p_h2.dtype.element_ty), boundary_check=(0, 1))
        if K > 128:
            p_h3 = tl.make_block_ptr(h + i_t * K * V, (K, V), (V, 1), (128, i_v * BV), (64, BV), (1, 0))
            tl.store(p_h3, b_h3.to(p_h3.dtype.element_ty), boundary_check=(0, 1))
        if K > 192:
            p_h4 = tl.make_block_ptr(h + i_t * K * V, (K, V), (V, 1), (192, i_v * BV), (64, BV), (1, 0))
            tl.store(p_h4, b_h4.to(p_h4.dtype.element_ty), boundary_check=(0, 1))

        # Compute v_residual = v - w @ h
        p_w = tl.make_block_ptr(w, (T, K), (K, 1), (i_t * BT, 0), (BT, 64), (1, 0))
        b_w = tl.load(p_w, boundary_check=(0, 1))
        b_v = tl.dot(b_w, b_h1.to(b_w.dtype))
        if K > 64:
            p_w = tl.make_block_ptr(w, (T, K), (K, 1), (i_t * BT, 64), (BT, 64), (1, 0))
            b_w = tl.load(p_w, boundary_check=(0, 1))
            b_v += tl.dot(b_w, b_h2.to(b_w.dtype))
        if K > 128:
            p_w = tl.make_block_ptr(w, (T, K), (K, 1), (i_t * BT, 128), (BT, 64), (1, 0))
            b_w = tl.load(p_w, boundary_check=(0, 1))
            b_v += tl.dot(b_w, b_h3.to(b_w.dtype))
        if K > 192:
            p_w = tl.make_block_ptr(w, (T, K), (K, 1), (i_t * BT, 192), (BT, 64), (1, 0))
            b_w = tl.load(p_w, boundary_check=(0, 1))
            b_v += tl.dot(b_w, b_h4.to(b_w.dtype))

        p_v = tl.make_block_ptr(v, (T, V), (V, 1), (i_t * BT, i_v * BV), (BT, BV), (1, 0))
        b_v = tl.load(p_v, boundary_check=(0, 1)) - b_v

        if SAVE_NEW_VALUE:
            p_v_new = tl.make_block_ptr(v_new_base, (T, V), (V, 1), (i_t * BT, i_v * BV), (BT, BV), (1, 0))
            tl.store(p_v_new, b_v.to(p_v_new.dtype.element_ty), boundary_check=(0, 1))

        last_idx = min((i_t + 1) * BT, T) - 1

        # Apply output gate g
        if USE_G:
            m_t = (i_t * BT + tl.arange(0, BT)).to(tl.float32) < T
            b_g_last = tl.load(g + last_idx)
            p_g = tl.make_block_ptr(g, (T,), (1,), (i_t * BT,), (BT,), (0,))
            b_g = tl.load(p_g, boundary_check=(0,))
            b_v *= (m_t * tl.exp(b_g_last - b_g))[:, None]
            b_g_last_exp = tl.exp(b_g_last)
            b_h1 *= b_g_last_exp
            if K > 64:
                b_h2 *= b_g_last_exp
            if K > 128:
                b_h3 *= b_g_last_exp
            if K > 192:
                b_h4 *= b_g_last_exp

        # Apply key gate gk
        if USE_GK:
            o_k1 = tl.arange(0, 64).to(tl.float32)
            gk_base_ptr = gk + (i_n * H + i_h) * T * K
            b_gk_last1 = tl.load(gk_base_ptr + last_idx * K + o_k1, mask=(o_k1 < K), other=0.)
            b_h1 *= tl.exp(b_gk_last1)[:, None]
            if K > 64:
                o_k2 = 64 + o_k1
                b_gk_last2 = tl.load(gk_base_ptr + last_idx * K + o_k2, mask=(o_k2 < K), other=0.)
                b_h2 *= tl.exp(b_gk_last2)[:, None]
            if K > 128:
                o_k3 = 128 + o_k1
                b_gk_last3 = tl.load(gk_base_ptr + last_idx * K + o_k3, mask=(o_k3 < K), other=0.)
                b_h3 *= tl.exp(b_gk_last3)[:, None]
            if K > 192:
                o_k4 = 192 + o_k1
                b_gk_last4 = tl.load(gk_base_ptr + last_idx * K + o_k4, mask=(o_k4 < K), other=0.)
                b_h4 *= tl.exp(b_gk_last4)[:, None]

        b_v = b_v.to(k.dtype.element_ty)

        # Update hidden state: h += k @ v
        p_k = tl.make_block_ptr(k, (K, T), (1, K), (0, i_t * BT), (64, BT), (0, 1))
        b_k = tl.load(p_k, boundary_check=(0, 1))
        if USE_GK:
            p_gk = tl.make_block_ptr(gk_base_ptr, (K, T), (1, K), (0, i_t * BT), (64, BT), (0, 1))
            b_k = (b_k * tl.exp(b_gk_last1[:, None] - tl.load(p_gk, boundary_check=(0, 1)))).to(b_k.dtype)
        b_h1 += tl.dot(b_k, b_v)

        if K > 64:
            p_k = tl.make_block_ptr(k, (K, T), (1, K), (64, i_t * BT), (64, BT), (0, 1))
            b_k = tl.load(p_k, boundary_check=(0, 1))
            if USE_GK:
                p_gk = tl.make_block_ptr(gk_base_ptr, (K, T), (1, K), (64, i_t * BT), (64, BT), (0, 1))
                b_k = (b_k * tl.exp(b_gk_last2[:, None] - tl.load(p_gk, boundary_check=(0, 1)))).to(b_k.dtype)
            b_h2 += tl.dot(b_k, b_v)

        if K > 128:
            p_k = tl.make_block_ptr(k, (K, T), (1, K), (128, i_t * BT), (64, BT), (0, 1))
            b_k = tl.load(p_k, boundary_check=(0, 1))
            if USE_GK:
                p_gk = tl.make_block_ptr(gk_base_ptr, (K, T), (1, K), (128, i_t * BT), (64, BT), (0, 1))
                b_k = (b_k * tl.exp(b_gk_last3[:, None] - tl.load(p_gk, boundary_check=(0, 1)))).to(b_k.dtype)
            b_h3 += tl.dot(b_k, b_v)

        if K > 192:
            p_k = tl.make_block_ptr(k, (K, T), (1, K), (192, i_t * BT), (64, BT), (0, 1))
            b_k = tl.load(p_k, boundary_check=(0, 1))
            if USE_GK:
                p_gk = tl.make_block_ptr(gk_base_ptr, (K, T), (1, K), (192, i_t * BT), (64, BT), (0, 1))
                b_k = (b_k * tl.exp(b_gk_last4[:, None] - tl.load(p_gk, boundary_check=(0, 1)))).to(b_k.dtype)
            b_h4 += tl.dot(b_k, b_v)

    # Store final state
    if STORE_FINAL_STATE:
        p_ht = tl.make_block_ptr(ht_ptr, (K, V), (V, 1), (0, i_v * BV), (64, BV), (1, 0))
        tl.store(p_ht, b_h1.to(p_ht.dtype.element_ty), boundary_check=(0, 1))
        if K > 64:
            p_ht = tl.make_block_ptr(ht_ptr, (K, V), (V, 1), (64, i_v * BV), (64, BV), (1, 0))
            tl.store(p_ht, b_h2.to(p_ht.dtype.element_ty), boundary_check=(0, 1))
        if K > 128:
            p_ht = tl.make_block_ptr(ht_ptr, (K, V), (V, 1), (128, i_v * BV), (64, BV), (1, 0))
            tl.store(p_ht, b_h3.to(p_ht.dtype.element_ty), boundary_check=(0, 1))
        if K > 192:
            p_ht = tl.make_block_ptr(ht_ptr, (K, V), (V, 1), (192, i_v * BV), (64, BV), (1, 0))
            tl.store(p_ht, b_h4.to(p_ht.dtype.element_ty), boundary_check=(0, 1))


def chunk_gated_delta_rule_fwd_h(
        k: torch.Tensor,
        w: torch.Tensor,
        u: torch.Tensor,
        g: Optional[torch.Tensor] = None,
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
        chunk_size: int = 64,  # default:64
        save_new_value: bool = True,
        cu_seqlens: Optional[torch.LongTensor] = None,
) -> Tuple[torch.Tensor, torch.Tensor, Optional[torch.Tensor]]:
    B, T, H, K, V = *k.shape, u.shape[-1]
    BT = chunk_size

    chunk_indices = prepare_chunk_indices(cu_seqlens, chunk_size) if cu_seqlens is not None else None
    # N: the actual number of sequences in the batch with either equal or variable lengths
    if cu_seqlens is None:
        N, NT, chunk_offsets = B, triton.cdiv(T, BT), None
    else:
        N, NT, chunk_offsets = len(cu_seqlens) - 1, len(chunk_indices), prepare_chunk_offsets(cu_seqlens, BT)
    if K > 256:
        raise ValueError("current kernel does not support head dimension larger than 256.")

    h = k.new_empty(B, NT, H, K, V).permute(0, 2, 1, 3, 4).contiguous()
    final_state = k.new_empty(N, H, K, V, dtype=torch.float32) if output_final_state else None

    BV = 128

    v_new = torch.empty_like(u).permute(0, 2, 1, 3).contiguous() if save_new_value else None
    k = k.permute(0, 2, 1, 3).contiguous()
    w = w.permute(0, 2, 1, 3).contiguous()
    u = u.permute(0, 2, 1, 3).contiguous()
    g = g.permute(0, 2, 1).contiguous()
    chunk_gated_delta_rule_fwd_kernel_h_blockdim64[(triton.cdiv(V, BV), N * H)](
        k=k,
        v=u,
        w=w,
        v_new=v_new,
306
307
308
309
310
311
312
313
314
315
316
317
318
319
        BT=BT,
        BV=BV,
        NT=NT,
    )
    h = h.permute(0, 2, 1, 3, 4).contiguous()
    v_new = v_new.permute(0, 2, 1, 3).contiguous()
    return h, v_new, final_state


@triton.heuristics({
    'USE_G': lambda args: args['g'] is not None,
    'USE_GK': lambda args: args['gk'] is not None,
    'USE_INITIAL_STATE': lambda args: args['dh0'] is not None,
    'USE_FINAL_STATE_GRADIENT': lambda args: args['dht'] is not None,
318
319
320
321
322
323
324
325
326
327
328
329
330
331
    'USE_INITIAL_STATE': lambda args: args['dh0'] is not None,
    'USE_FINAL_STATE_GRADIENT': lambda args: args['dht'] is not None,
    'IS_VARLEN': lambda args: args['cu_seqlens'] is not None,
})
@triton.autotune(
    configs=get_autotune_config(multibuffer_list=(True, False)),
    key=['H', 'K', 'V', 'BT', 'BV', 'USE_G', 'IS_VARLEN'],
)
@triton.jit(do_not_specialize=['T'])
def chunk_gated_delta_rule_bwd_kernel_dhu_blockdim64(
        q,
        k,
        w,
        g,
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
416
417
418
419
420
421
422
423
424
425
426
427
428
429
430
431
432
433
434
435
436
437
438
439
440
441
442
443
444
445
446
447
448
449
450
451
452
453
454
455
456
457
458
459
460
461
462
463
464
465
466
467
468
469
470
471
472
473
474
475
476
477
478
479
480
481
482
483
484
485
486
487
488
489
490
491
492
493
494
495
496
497
498
499
500
501
502
503
504
505
506
507
508
509
510
511
512
513
514
515
516
517
518
519
520
521
522
523
524
525
526
527
528
529
530
531
532
533
534
535
536
537
538
539
        USE_INITIAL_STATE: tl.constexpr,
        USE_FINAL_STATE_GRADIENT: tl.constexpr,
        IS_VARLEN: tl.constexpr,
):
    T_all = T
    i_v, i_nh = tl.program_id(0), tl.program_id(1)
    i_n, i_h = i_nh // H, i_nh % H
    if IS_VARLEN:
        bos, eos = tl.load(cu_seqlens + i_n).to(tl.int32), tl.load(cu_seqlens + i_n + 1).to(tl.int32)
        T = eos - bos
        NT = tl.cdiv(T, BT)
        boh = tl.load(chunk_offsets + i_n).to(tl.int32)
    else:
        bos, eos = i_n * T, i_n * T + T
        NT = tl.cdiv(T, BT)
        boh = i_n * NT

    b_dh1 = tl.zeros([64, BV], dtype=tl.float32)
    if K > 64:
        b_dh2 = tl.zeros([64, BV], dtype=tl.float32)
    if K > 128:
        b_dh3 = tl.zeros([64, BV], dtype=tl.float32)
    if K > 192:
        b_dh4 = tl.zeros([64, BV], dtype=tl.float32)

    q += (bos * H + i_h) * K
    k += (bos * H + i_h) * K
    w += (bos * H + i_h) * K
    do += (bos * H + i_h) * V
    dv += (bos * H + i_h) * V
    dv2 += (bos * H + i_h) * V
    dh += (boh * H + i_h) * K * V
    if USE_GK:
        gk += (bos * H + i_h) * K

    if USE_INITIAL_STATE:
        dh0 += i_nh * K * V
    if USE_FINAL_STATE_GRADIENT:
        dht += i_nh * K * V

    stride_v = H * V
    stride_h = H * K * V
    stride_k = H * K

    if USE_FINAL_STATE_GRADIENT:
        p_dht1 = tl.make_block_ptr(dht, (K, V), (V, 1), (0, i_v * BV), (64, BV), (1, 0))
        b_dh1 += tl.load(p_dht1, boundary_check=(0, 1))
        if K > 64:
            p_dht2 = tl.make_block_ptr(dht, (K, V), (V, 1), (64, i_v * BV), (64, BV), (1, 0))
            b_dh2 += tl.load(p_dht2, boundary_check=(0, 1))
        if K > 128:
            p_dht3 = tl.make_block_ptr(dht, (K, V), (V, 1), (128, i_v * BV), (64, BV), (1, 0))
            b_dh3 += tl.load(p_dht3, boundary_check=(0, 1))
        if K > 192:
            p_dht4 = tl.make_block_ptr(dht, (K, V), (V, 1), (192, i_v * BV), (64, BV), (1, 0))
            b_dh4 += tl.load(p_dht4, boundary_check=(0, 1))

    for i_t in range(NT - 1, -1, -1):
        p_dh1 = tl.make_block_ptr(dh + i_t * stride_h, (K, V), (V, 1), (0, i_v * BV), (64, BV), (1, 0))
        tl.store(p_dh1, b_dh1.to(p_dh1.dtype.element_ty), boundary_check=(0, 1))
        if K > 64:
            p_dh2 = tl.make_block_ptr(dh + i_t * stride_h, (K, V), (V, 1), (64, i_v * BV), (64, BV), (1, 0))
            tl.store(p_dh2, b_dh2.to(p_dh2.dtype.element_ty), boundary_check=(0, 1))
        if K > 128:
            p_dh3 = tl.make_block_ptr(dh + i_t * stride_h, (K, V), (V, 1), (128, i_v * BV), (64, BV), (1, 0))
            tl.store(p_dh3, b_dh3.to(p_dh3.dtype.element_ty), boundary_check=(0, 1))
        if K > 192:
            p_dh4 = tl.make_block_ptr(dh + i_t * stride_h, (K, V), (V, 1), (192, i_v * BV), (64, BV), (1, 0))
            tl.store(p_dh4, b_dh4.to(p_dh4.dtype.element_ty), boundary_check=(0, 1))

        last_idx = min((i_t + 1) * BT, T) - 1
        if USE_G:
            if IS_VARLEN:
                bos_g = i_h * T_all + bos
            else:
                bos_g = (i_n * H + i_h) * T_all
            bg_last = tl.load(g + bos_g + last_idx)
            bg_last_exp = tl.exp(bg_last)
            p_g = tl.make_block_ptr(base=g + bos_g, shape=(T,), strides=(1,), offsets=(i_t * BT,), block_shape=(BT,), order=(0,))
            b_g = tl.load(p_g, boundary_check=(0,))
            b_g_exp = tl.exp(b_g)

        p_dv = tl.make_block_ptr(dv, (T, V), (stride_v, 1), (i_t * BT, i_v * BV), (BT, BV), (1, 0))
        p_dv2 = tl.make_block_ptr(dv2, (T, V), (stride_v, 1), (i_t * BT, i_v * BV), (BT, BV), (1, 0))
        p_do = tl.make_block_ptr(do, (T, V), (stride_v, 1), (i_t * BT, i_v * BV), (BT, BV), (1, 0))

        b_do = tl.load(p_do, boundary_check=(0, 1))

        # Update dv
        p_k = tl.make_block_ptr(k, (T, K), (stride_k, 1), (i_t * BT, 0), (BT, 64), (1, 0))
        b_k = tl.load(p_k, boundary_check=(0, 1))
        if USE_GK:
            o_k1 = tl.arange(0, 64)
            b_gk_last1 = tl.load(gk + last_idx * H * K + o_k1, mask=(o_k1 < K), other=0.)
        b_dv = tl.dot(b_k, b_dh1.to(b_k.dtype))

        if K > 64:
            p_k = tl.make_block_ptr(k, (T, K), (stride_k, 1), (i_t * BT, 64), (BT, 64), (1, 0))
            b_k = tl.load(p_k, boundary_check=(0, 1))
            if USE_GK:
                o_k2 = 64 + o_k1
                b_gk_last2 = tl.load(gk + last_idx * H * K + o_k2, mask=(o_k2 < K), other=0.)
            b_dv += tl.dot(b_k, b_dh2.to(b_k.dtype))

        if K > 128:
            p_k = tl.make_block_ptr(k, (T, K), (stride_k, 1), (i_t * BT, 128), (BT, 64), (1, 0))
            b_k = tl.load(p_k, boundary_check=(0, 1))
            if USE_GK:
                o_k3 = 128 + o_k1
                b_gk_last3 = tl.load(gk + last_idx * H * K + o_k3, mask=(o_k3 < K), other=0.)
            b_dv += tl.dot(b_k, b_dh3.to(b_k.dtype))

        if K > 192:
            p_k = tl.make_block_ptr(k, (T, K), (stride_k, 1), (i_t * BT, 192), (BT, 64), (1, 0))
            b_k = tl.load(p_k, boundary_check=(0, 1))
            if USE_GK:
                o_k4 = 192 + o_k1
                b_gk_last4 = tl.load(gk + last_idx * H * K + o_k4, mask=(o_k4 < K), other=0.)
            b_dv += tl.dot(b_k, b_dh4.to(b_k.dtype))

        if USE_G:
            m_t = (i_t * BT + tl.arange(0, BT)).to(tl.float32) < T
            b_dv *= (m_t * tl.exp(bg_last - b_g))[:, None]
        b_dv += tl.load(p_dv, boundary_check=(0, 1))

        tl.store(p_dv2, b_dv.to(p_dv.dtype.element_ty), boundary_check=(0, 1))
        # Update dh
        p_w = tl.make_block_ptr(w, (K, T), (1, stride_k), (0, i_t * BT), (64, BT), (0, 1))
        p_q = tl.make_block_ptr(q, (K, T), (1, stride_k), (0, i_t * BT), (64, BT), (0, 1))
        b_w = tl.load(p_w, boundary_check=(0, 1))
        b_q = tl.load(p_q, boundary_check=(0, 1))
        if USE_G:
            b_dh1 *= bg_last_exp
            b_q = b_q * b_g_exp[None, :]
        if USE_GK:
            b_dh1 *= tl.exp(b_gk_last1[:, None])
        b_dh1 += tl.dot(b_q.to(b_q.dtype), b_do.to(b_q.dtype)) * scale - tl.dot(b_w, b_dv.to(b_w.dtype))
        if K > 64:
            p_q = tl.make_block_ptr(q, (K, T), (1, stride_k), (64, i_t * BT), (64, BT), (0, 1))
            p_w = tl.make_block_ptr(w, (K, T), (1, stride_k), (64, i_t * BT), (64, BT), (0, 1))
            b_q = tl.load(p_q, boundary_check=(0, 1))
            b_w = tl.load(p_w, boundary_check=(0, 1))
            if USE_G:
                b_dh2 *= bg_last_exp
                b_q = b_q * b_g_exp[None, :]
            if USE_GK:
                b_dh2 *= tl.exp(b_gk_last2[:, None])
            b_dh2 += tl.dot(b_q.to(b_q.dtype), b_do.to(b_q.dtype)) * scale - tl.dot(b_w, b_dv.to(b_w.dtype))
        if K > 128:
            p_q = tl.make_block_ptr(q, (K, T), (1, stride_k), (128, i_t * BT), (64, BT), (0, 1))
            p_w = tl.make_block_ptr(w, (K, T), (1, stride_k), (128, i_t * BT), (64, BT), (0, 1))
            b_q = tl.load(p_q, boundary_check=(0, 1))
            b_w = tl.load(p_w, boundary_check=(0, 1))
            if USE_G:
                b_dh3 *= bg_last_exp
                b_q = b_q * b_g_exp[None, :]
            if USE_GK:
                b_dh3 *= tl.exp(b_gk_last3[:, None])
            b_dh3 += tl.dot(b_q.to(b_q.dtype), b_do.to(b_q.dtype)) * scale - tl.dot(b_w, b_dv.to(b_w.dtype))
        if K > 192:
            p_q = tl.make_block_ptr(q, (K, T), (1, stride_k), (192, i_t * BT), (64, BT), (0, 1))
            p_w = tl.make_block_ptr(w, (K, T), (1, stride_k), (192, i_t * BT), (64, BT), (0, 1))
            b_q = tl.load(p_q, boundary_check=(0, 1))
            b_w = tl.load(p_w, boundary_check=(0, 1))
            if USE_G:
                b_dh4 *= bg_last_exp
                b_q = b_q * b_g_exp[None, :]
            if USE_GK:
                b_dh4 *= tl.exp(b_gk_last4[:, None])
            b_dh4 += tl.dot(b_q.to(b_q.dtype), b_do.to(b_q.dtype)) * scale - tl.dot(b_w, b_dv.to(b_w.dtype))

    if USE_INITIAL_STATE:
        p_dh0 = tl.make_block_ptr(dh0, (K, V), (V, 1), (0, i_v * BV), (64, BV), (1, 0))
        tl.store(p_dh0, b_dh1.to(p_dh0.dtype.element_ty), boundary_check=(0, 1))
        if K > 64:
            p_dh1 = tl.make_block_ptr(dh0, (K, V), (V, 1), (64, i_v * BV), (64, BV), (1, 0))
            tl.store(p_dh1, b_dh2.to(p_dh1.dtype.element_ty), boundary_check=(0, 1))
        if K > 128:
            p_dh2 = tl.make_block_ptr(dh0, (K, V), (V, 1), (128, i_v * BV), (64, BV), (1, 0))
            tl.store(p_dh2, b_dh3.to(p_dh2.dtype.element_ty), boundary_check=(0, 1))
        if K > 192:
            p_dh3 = tl.make_block_ptr(dh0, (K, V), (V, 1), (192, i_v * BV), (64, BV), (1, 0))
            tl.store(p_dh3, b_dh4.to(p_dh3.dtype.element_ty), boundary_check=(0, 1))


def chunk_gated_delta_rule_bwd_dhu(
    q: torch.Tensor,
    k: torch.Tensor,
    w: torch.Tensor,
    do: torch.Tensor,
547
548
549
550
551
552
553
554
555
556
557
558
559
560
561
562
563
564
565
566
567
568
569
570
571
572
573
574
575
576
    chunk_size: int = 64,  # SY: remove this argument and force chunk size 64?
    chunk_indices: torch.LongTensor | None = None,
    use_exp2: bool = False,
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
    B, T, H, K, V = *q.shape, do.shape[-1]
    # N: the actual number of sequences in the batch with either equal or variable lengths
    BT = 64
    if K > 256:
        raise ValueError("current kernel does not support head dimension being larger than 256.")

    if chunk_indices is None and cu_seqlens is not None:
        chunk_indices = prepare_chunk_indices(cu_seqlens, chunk_size)
    if cu_seqlens is None:
        N, NT, chunk_offsets = B, triton.cdiv(T, BT), None
    else:
        N, NT, chunk_offsets = len(cu_seqlens) - 1, len(chunk_indices), prepare_chunk_offsets(cu_seqlens, BT)

    dh = q.new_empty(B, NT, H, K, V)
    dh0 = torch.empty_like(h0, dtype=torch.float32) if h0 is not None else None
    dv2 = torch.empty_like(dv)

    BV = 128

    g = g.permute(0, 2, 1).contiguous()

    chunk_gated_delta_rule_bwd_kernel_dhu_blockdim64[(triton.cdiv(V, BV), N * H)](
        q=q,
        k=k,
        w=w,
        g=g,
590
591
592
593
594
        V=V,
        BT=BT,
        BV=BV,
    )
    return dh, dh0, dv2
hyper_parallel/components/functional/_triton/gated_delta_net/chunk_o.py
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
# pylint: disable=invalid-name,missing-module-docstring,missing-function-docstring
# pylint: disable=unused-variable,too-many-nested-blocks
# pylint: disable=forbidden-backend-import

from typing import Optional, Tuple

import torch
import triton
import triton.language as tl

from .utils import prepare_chunk_indices, exp, prepare_chunk_offsets


@triton.heuristics({
    'USE_G': lambda args: args['g'] is not None,
    'USE_G_GAMMA': lambda args: args['g_gamma'] is not None,
    'USE_DW': lambda args: args['dw'] is not None,
    'IS_VARLEN': lambda args: args['cu_seqlens'] is not None,
33
34
35
36
37
38
39
40
41
42
    'USE_G_GAMMA': lambda args: args['g_gamma'] is not None,
    'USE_DW': lambda args: args['dw'] is not None,
    'IS_VARLEN': lambda args: args['cu_seqlens'] is not None,
})
@triton.jit(do_not_specialize=['T'])
def chunk_bwd_kernel_dqkwg(
    q,
    k,
    v,
    h,
 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
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
    USE_DW: tl.constexpr,
    IS_VARLEN: tl.constexpr,
    gdiff,
):
    i_t, i_b = tl.program_id(0), tl.program_id(1)
    T_max = T
    if IS_VARLEN:
        i_tg = i_t
        i_n, i_t = tl.load(chunk_indices + i_t * 2).to(tl.int32), tl.load(chunk_indices + i_t * 2 + 1).to(tl.int32)
        bos, eos = tl.load(cu_seqlens + i_n).to(tl.int32), tl.load(cu_seqlens + i_n + 1).to(tl.int32)
        total = B * T_max
        T = eos - bos
    else:
        NT = tl.cdiv(T, BT)
        i_tg = i_b * NT + i_t
        bos, eos = i_b * T, i_b * T + T
        total = B * T_max

    NK = tl.cdiv(K, BK)
    for i_k in range(NK):
        if USE_G:
            dg_k = dg + i_k * total * H

        for i_h in range(H):
            v_h = v + (bos * H + i_h) * V
            do_h = do + (bos * H + i_h) * V
            h_h = h + (i_tg * H + i_h).to(tl.int64) * K * V
            dh_h = dh + (i_tg * H + i_h).to(tl.int64) * K * V
            q_h = q + (bos * H + i_h) * K
            k_h = k + (bos * H + i_h) * K
            dq_h = dq + (bos * H + i_h) * K
            dk_h = dk + (bos * H + i_h) * K

            if USE_DW:
                w_h = w + (bos * H + i_h) * K
                dw_h = dw + (bos * H + i_h) * K
                dv_h = dv + (bos * H + i_h) * V

            if USE_G:
                if IS_VARLEN:
                    dg_h = dg_k + i_h * T_max + bos
                    g_h = g + i_h * T_max + bos
                else:
                    dg_h = dg_k + (i_b * H + i_h) * T_max
                    g_h = g + (i_b * H + i_h) * T_max
                b_dg_last = tl.zeros([1, ], dtype=tl.float32)

            if USE_G_GAMMA:
                b_gamma = tl.load(g_gamma + i_h)
                b_g = b_gamma * (tl.arange(0, BT) + 1)
                b_g_last = b_gamma * min(BT, T - i_t * BT)

            b_dq = tl.zeros([BT, BK], dtype=tl.float32)
            b_dk = tl.zeros([BT, BK], dtype=tl.float32)
            b_ds = tl.zeros([BT, BT], dtype=tl.float32)
            b_dw = tl.zeros([BT, BK], dtype=tl.float32) if USE_DW else None

            for i_v in range(tl.cdiv(V, BV)):
                p_v = tl.make_block_ptr(v_h, (T, V), (H * V, 1), (i_t * BT, i_v * BV), (BT, BV), (1, 0))
                p_do = tl.make_block_ptr(do_h, (T, V), (H * V, 1), (i_t * BT, i_v * BV), (BT, BV), (1, 0))
                p_h = tl.make_block_ptr(h_h, (V, K), (1, V), (i_v * BV, i_k * BK), (BV, BK), (0, 1))
                p_dh = tl.make_block_ptr(dh_h, (V, K), (1, V), (i_v * BV, i_k * BK), (BV, BK), (0, 1))

                b_v = tl.load(p_v, boundary_check=(0, 1))
                b_do = tl.load(p_do, boundary_check=(0, 1))
                b_h = tl.load(p_h, boundary_check=(0, 1))
                b_dh = tl.load(p_dh, boundary_check=(0, 1))

                if USE_G:
                    b_dg_last += (tl.sum(b_h * b_dh))

                b_ds += tl.dot(b_do, tl.trans(b_v))
                b_dq += tl.dot(b_do, b_h.to(b_do.dtype))
                b_dk += tl.dot(b_v, b_dh.to(b_v.dtype))

                if USE_DW:
                    p_dv = tl.make_block_ptr(dv_h, (T, V), (H * V, 1), (i_t * BT, i_v * BV), (BT, BV), (1, 0))
                    b_dv = tl.load(p_dv, boundary_check=(0, 1))
                    b_dw += tl.dot(b_dv.to(b_v.dtype), b_h.to(b_v.dtype))

            if USE_DW:
                p_dw = tl.make_block_ptr(dw_h, (T, K), (H * K, 1), (i_t * BT, i_k * BK), (BT, BK), (1, 0))
                tl.store(p_dw, -b_dw.to(p_dw.dtype.element_ty), boundary_check=(0, 1))

            tl.debug_barrier()

            p_q = tl.make_block_ptr(q_h, (T, K), (H * K, 1), (i_t * BT, i_k * BK), (BT, BK), (1, 0))
            p_k = tl.make_block_ptr(k_h, (T, K), (H * K, 1), (i_t * BT, i_k * BK), (BT, BK), (1, 0))
            b_q = tl.load(p_q, boundary_check=(0, 1))
            b_k = tl.load(p_k, boundary_check=(0, 1))

            p_dq = tl.make_block_ptr(dq_h, (T, K), (H * K, 1), (i_t * BT, i_k * BK), (BT, BK), (1, 0))
            p_dk = tl.make_block_ptr(dk_h, (T, K), (H * K, 1), (i_t * BT, i_k * BK), (BT, BK), (1, 0))

            o_t = i_t * BT + tl.arange(0, BT)
            m_t = o_t < T
            m_A = (o_t[:, None] >= o_t[None, :]) & (m_t[:, None] & m_t)

            if USE_G:
                b_dg = tl.zeros([BT, ], dtype=tl.float32)
                p_g = tl.make_block_ptr(g_h, (T,), (1,), (i_t * BT,), (BT,), (0,))
                b_g = tl.load(p_g, boundary_check=(0,))
                b_g_last = tl.load(g_h + (min(i_t * BT + BT, T) - 1) * 1)
                b_dg_last *= tl.exp(b_g_last)

                b_dq = b_dq * tl.exp(b_g)[:, None] * scale
                b_dg += tl.sum(b_dq * b_q, axis=1)

                b_dk = b_dk * tl.where(m_t, tl.exp(-b_g + b_g_last), 0)[:, None]
                b_dg -= tl.sum(b_k * b_dk, axis=1)
                b_dg_last += tl.sum(b_dk * b_k)

                if IS_VARLEN:
                    b_ds = tl.where(m_A, b_ds * exp(b_g[:, None] - b_g[None, :]), 0) * scale
                else:
                    p_gdiff = tl.make_block_ptr(gdiff + i_b * H * NT * BT * BT + i_h * NT * BT * BT + i_t * BT * BT,
                                                (BT, BT), (BT, 1), (0, 0), (BT, BT), (1, 0))
                    gdiff_ = tl.load(p_gdiff)
                    b_ds = b_ds * gdiff_ * scale

                b_ds2 = b_ds * tl.dot(b_q, tl.trans(b_k))
                b_dg += tl.sum(b_ds2, axis=1)
                b_dg -= tl.sum(b_ds2, axis=0)

                b_ds = b_ds.to(b_k.dtype)
                b_dq += tl.dot(b_ds, b_k)
                b_dk += tl.dot(tl.trans(b_ds), b_q)
                p_dg = tl.make_block_ptr(dg_h, (T,), (1,), (i_t * BT,), (BT,), (0,))

                last_index_local = min(BT, T - i_t * BT) - 1
                if last_index_local >= 0:
                    is_last_mask = tl.arange(0, BT) == last_index_local
                    b_dg = tl.where(is_last_mask, b_dg + b_dg_last, b_dg)
                else:
                    pass

                tl.store(p_dq, b_dq.to(p_dq.dtype.element_ty), boundary_check=(0, 1))
                tl.store(p_dk, b_dk.to(p_dk.dtype.element_ty), boundary_check=(0, 1))
                tl.store(p_dg, b_dg.to(p_dg.dtype.element_ty), boundary_check=(0,))

            elif USE_G_GAMMA:
                b_dq = b_dq * exp(b_g)[:, None] * scale
                b_dk = b_dk * tl.where(m_t, exp(-b_g + b_g_last), 0)[:, None]
                b_ds = tl.where(m_A, b_ds * exp(b_g[:, None] - b_g[None, :]), 0) * scale
                b_ds = b_ds.to(b_k.dtype)
                b_dq += tl.dot(b_ds, b_k)
                b_dk += tl.dot(tl.trans(b_ds), b_q)
                tl.store(p_dq, b_dq.to(p_dq.dtype.element_ty), boundary_check=(0, 1))
                tl.store(p_dk, b_dk.to(p_dk.dtype.element_ty), boundary_check=(0, 1))

            else:
                b_ds = tl.where(m_A, b_ds, 0)
                b_ds = b_ds.to(b_k.dtype)
                b_dq += tl.dot(b_ds, b_k)
                b_dk += tl.dot(tl.trans(b_ds), b_q) * scale
                b_dq *= scale
                tl.store(p_dq, b_dq.to(p_dq.dtype.element_ty), boundary_check=(0, 1))
                tl.store(p_dk, b_dk.to(p_dk.dtype.element_ty), boundary_check=(0, 1))


@triton.heuristics({
    'USE_G': lambda args: args['g'] is not None,
    'USE_G_GAMMA': lambda args: args['g_gamma'] is not None,
    'IS_VARLEN': lambda args: args['cu_seqlens'] is not None,
})
@triton.jit(do_not_specialize=['T'])
def chunk_bwd_kernel_dv_local(
    q,
    k,
    g,
    g_gamma,
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
    USE_G: tl.constexpr,
    USE_G_GAMMA: tl.constexpr,
    IS_VARLEN: tl.constexpr,
):
    i_t, i_b = tl.program_id(0), tl.program_id(1)
    T_max = T

    if IS_VARLEN:
        i_n, i_t = tl.load(chunk_indices + i_t * 2).to(tl.int32), tl.load(chunk_indices + i_t * 2 + 1).to(tl.int32)
        bos, eos = tl.load(cu_seqlens + i_n).to(tl.int32), tl.load(cu_seqlens + i_n + 1).to(tl.int32)
        T = eos - bos
    else:
        bos, eos = i_b * T, i_b * T + T

    for i_h in range(H):
        offset_kh = (bos * H + i_h) * K
        offset_vh = (bos * H + i_h) * V

        b_A = tl.zeros([BT, BT], dtype=tl.float32)
        for i_k in range(tl.cdiv(K, BK)):
            p_k = tl.make_block_ptr(k + offset_kh, (T, K), (H * K, 1), (i_t * BT, i_k * BK), (BT, BK), (1, 0))
            p_q = tl.make_block_ptr(q + offset_kh, (K, T), (1, H * K), (i_k * BK, i_t * BT), (BK, BT), (0, 1))
            b_q = tl.load(p_q, boundary_check=(0, 1))
            b_k = tl.load(p_k, boundary_check=(0, 1))
            b_A += tl.dot(b_k, b_q)

        if USE_G:
            if IS_VARLEN:
                offset_g = i_h * T_max + bos
            else:
                offset_g = i_b * H * T_max + i_h * T_max

            p_g = tl.make_block_ptr(g + offset_g, (T,), (1,), (i_t * BT,), (BT,), (0,))
            b_g = tl.load(p_g, boundary_check=(0,))

        if USE_G_GAMMA:
            b_gamma = tl.load(g_gamma + i_h)
            b_g = b_gamma * (tl.arange(0, BT) + 1)

        o_t = i_t * BT + tl.arange(0, BT)
        m_t = o_t < T
        m_A = (o_t[:, None] <= o_t[None, :]) & (m_t[:, None] & m_t)

        if USE_G:
            b_A = tl.where(m_A, b_A * tl.exp(b_g[None, :] - b_g[:, None]) * scale, 0).to(do.dtype.element_ty)
        else:
            b_A = tl.where(m_A, b_A * scale, 0).to(do.dtype.element_ty)

        for i_v in range(tl.cdiv(V, BV)):
            p_do = tl.make_block_ptr(do + offset_vh, (T, V), (H * V, 1), (i_t * BT, i_v * BV), (BT, BV), (1, 0))
            p_dv = tl.make_block_ptr(dv + offset_vh, (T, V), (H * V, 1), (i_t * BT, i_v * BV), (BT, BV), (1, 0))
            b_do = tl.load(p_do, boundary_check=(0, 1))
            b_dv = tl.dot(b_A.to(b_do.dtype), b_do)
            tl.store(p_dv, b_dv.to(p_dv.dtype.element_ty), boundary_check=(0, 1))


@triton.heuristics({
    'USE_G': lambda args: args['g'] is not None,
    'USE_G_GAMMA': lambda args: args['g_gamma'] is not None,
    'IS_VARLEN': lambda args: args['cu_seqlens'] is not None
})
@triton.jit(do_not_specialize=['T'])
def chunk_fwd_kernel_o(
    q,
    k,
    v,
    h,
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
416
417
418
419
420
421
422
423
424
425
426
427
    USE_G: tl.constexpr,
    USE_G_GAMMA: tl.constexpr,
    IS_VARLEN: tl.constexpr,
):
    T_max = T
    for i_v in range(tl.cdiv(V, BV)):
        for i_n in range(N):
            if IS_VARLEN:
                bos, eos = tl.load(cu_seqlens + i_n).to(tl.int32), tl.load(
                    cu_seqlens + i_n + 1
                ).to(tl.int32)
                T = eos - bos
                NT = tl.cdiv(T, BT)
                boh = tl.load(chunk_offsets + i_n).to(tl.int64)
            else:
                bos, eos = i_n * T, i_n * T + T
                NT = tl.cdiv(T, BT)
                boh = i_n * NT

            core_id = tl.program_id(0)
            total_cores = tl.num_programs(0)
            base_chunks_per_pid = NT // total_cores
            remainder = NT % total_cores

            if core_id < remainder:
                chunks_this_pid = base_chunks_per_pid + 1
                start_idx = core_id * chunks_this_pid
            else:
                chunks_this_pid = base_chunks_per_pid
                start_idx = core_id * base_chunks_per_pid + remainder

            # offset calculation
            for i_h in range(0, H):
                q_offset = (bos * Hg + i_h // (H // Hg)) * K
                k_offset = (bos * Hg + i_h // (H // Hg)) * K
                v_offset = (bos * H + i_h) * V
                o_offset = (bos * H + i_h) * V

                for i_t in range(start_idx, start_idx + chunks_this_pid):
                    i_tg = boh + i_t
                    h_base = h + (i_tg * H + i_h).to(tl.int64) * K * V
                    b_o = tl.zeros([BT, BV], dtype=tl.float32)
                    b_A = tl.zeros([BT, BT], dtype=tl.float32)
                    for i_k in range(tl.cdiv(K, BK)):
                        p_q = tl.make_block_ptr(
                            q + q_offset, (T, K), (Hg * K, 1), (i_t * BT, i_k * BK), (BT, BK), (1, 0)
                        )
                        p_k = tl.make_block_ptr(
                            k + k_offset, (K, T), (1, Hg * K), (i_k * BK, i_t * BT), (BK, BT), (0, 1)
                        )
                        p_h = tl.make_block_ptr(
                            h_base, (K, V), (V, 1), (i_k * BK, i_v * BV), (BK, BV), (1, 0)
                        )
                        b_q = tl.load(p_q, boundary_check=(0, 1))
                        b_k = tl.load(p_k, boundary_check=(0, 1))
                        b_h = tl.load(p_h, boundary_check=(0, 1))

                        # [BT, BK] @ [BK, BV] -> [BT, BV]
                        b_o += tl.dot(b_q, b_h)
                        # [BT, BK] @ [BK, BT] -> [BT, BT]
                        b_A += tl.dot(b_q, b_k)

                    if USE_G:
                        if IS_VARLEN:
                            p_g = tl.make_block_ptr(g + bos + i_h * T_max, (T,), (1,), (i_t * BT,), (BT,), (0,))
                        else:
                            p_g = tl.make_block_ptr(g + bos * H + i_h * T_max, (T,), (1,), (i_t * BT,), (BT,), (0,))
                        b_g = tl.load(p_g, boundary_check=(0,))
                        b_o = b_o * exp(b_g)[:, None]
                        b_A = b_A * exp(b_g[:, None] - b_g[None, :])
                    if USE_G_GAMMA:
                        b_gamma = tl.load(g_gamma + i_h)
                        b_g = b_gamma * (tl.arange(0, BT) + 1)

                    o_i = tl.arange(0, BT)
                    m_A = o_i[:, None] >= o_i[None, :]
                    b_A = tl.where(m_A, b_A, 0)

                    p_v = tl.make_block_ptr(
                        v + v_offset, (T, V), (H * V, 1), (i_t * BT, i_v * BV), (BT, BV), (1, 0)
                    )
                    p_o = tl.make_block_ptr(
                        o + o_offset, (T, V), (H * V, 1), (i_t * BT, i_v * BV), (BT, BV), (1, 0)
                    )
                    b_v = tl.load(p_v, boundary_check=(0, 1))

                    # to fix mma -> mma layout conversion
                    # already solved by triton v3.2 or higher
                    b_o = b_o * scale + tl.dot(b_A.to(b_v.dtype), b_v) * scale
                    tl.store(p_o, b_o.to(p_o.dtype.element_ty), boundary_check=(0, 1))


def chunk_bwd_dqkwg(
    q: torch.Tensor,
    k: torch.Tensor,
    v: torch.Tensor,
    do: torch.Tensor,
434
435
436
437
438
439
440
441
442
443
444
445
446
447
448
449
450
451
452
453
454
455
456
457
458
459
460
461
462
463
464
465
466
467
468
469
470
471
472
473
474
475
476
477
    cu_seqlens: Optional[torch.LongTensor] = None,
    chunk_size: int = 64,
    scale: float = 1.0,
) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
    B, T, H, K, V = *k.shape, v.shape[-1]
    BT = min(chunk_size, max(16, triton.next_power_of_2(T)))
    chunk_indices = prepare_chunk_indices(cu_seqlens, BT) if cu_seqlens is not None else None
    NT = triton.cdiv(T, BT) if cu_seqlens is None else len(chunk_indices)

    BK = 128 if cu_seqlens is None else 64
    BV = 64
    NK = triton.cdiv(K, BK)
    dq = torch.empty_like(q)
    dk = torch.empty_like(k)
    g = g.transpose(1, 2).contiguous()
    dg = torch.empty(NK, *g.shape, dtype=torch.float32, device=g.device) if g is not None else None
    dw = torch.empty_like(w) if w is not None else None
    grid = (NT, B)

    if cu_seqlens is None:
        if NT * BT == T:
            g_ = g.reshape(B, H, NT, BT)
            g_diff = g_[:, :, :, :, None] - g_[:, :, :, None, :]
            g_diff = g_diff.clamp(-60, 60).exp()
            g_diff[:, :, :] *= torch.tril(torch.ones(BT, BT), diagonal=0).to(g.device)
        else:
            diff = NT * BT - T
            g_ = torch.cat((g, torch.zeros(B, H, diff).to(g.device)), dim=-1).reshape(B, H, NT, BT)
            g_diff = g_[:, :, :, :, None] - g_[:, :, :, None, :]
            g_diff = g_diff.clamp(-60, 60).exp()
            g_diff[:, :, :] *= torch.tril(torch.ones(BT, BT), diagonal=0).to(g.device)
            bias = torch.arange(0, BT).to(g.device)
            o_t = (NT - 1) * BT + bias
            m_t = o_t < T
            m_A = (m_t[:, None] & m_t)
            g_diff[:, :, -1] *= m_A
    else:
        g_diff = None

    chunk_bwd_kernel_dqkwg[grid](
        q=q,
        k=k,
        v=v,
        h=h,
498
499
500
501
502
503
504
505
506
507
508
509
510
511
512
        BV=BV,
        gdiff=g_diff,
    )

    if dg is not None:
        dg = dg.sum(0)
        dg = dg.transpose(1, 2).contiguous()
    return dq, dk, dw, dg


def chunk_bwd_dv_local(
    q: torch.Tensor,
    k: torch.Tensor,
    do: torch.Tensor,
    g: Optional[torch.Tensor] = None,
514
515
516
517
518
519
520
521
522
523
524
525
526
527
528
529
530
531
532
533
    scale: float = None,
    cu_seqlens: Optional[torch.LongTensor] = None,
    chunk_size: int = 64
) -> torch.Tensor:
    B, T, H, K, V = *k.shape, do.shape[-1]
    BT = min(chunk_size, max(16, triton.next_power_of_2(T)))
    chunk_indices = prepare_chunk_indices(cu_seqlens, BT) if cu_seqlens is not None else None

    BK = 128
    BV = 128
    NT = triton.cdiv(T, BT) if cu_seqlens is None else len(chunk_indices)

    g = g.transpose(1, 2).contiguous()
    dv = torch.empty_like(do)
    grid = (NT, B)
    chunk_bwd_kernel_dv_local[grid](
        q=q,
        k=k,
        g=g,
        g_gamma=g_gamma,
543
544
545
546
547
548
549
550
551
552
553
554
        BT=BT,
        BK=BK,
        BV=BV,
    )
    return dv


def chunk_fwd_o(
    q: torch.Tensor,
    k: torch.Tensor,
    v: torch.Tensor,
    h: torch.Tensor,
557
558
559
560
561
562
563
564
565
566
567
568
569
570
571
572
573
574
575
576
577
578
579
580
581
582
583
584
585
586
587
    scale: Optional[float] = None,
    cu_seqlens: Optional[torch.LongTensor] = None,
    chunk_size: int = 64
) -> torch.Tensor:
    B, T, Hg, K, V = *q.shape, v.shape[-1]
    H = v.shape[-2]
    BT = min(chunk_size, max(16, triton.next_power_of_2(T)))
    chunk_indices = (
        prepare_chunk_indices(cu_seqlens, BT) if cu_seqlens is not None else None
    )
    NT = triton.cdiv(T, BT) if cu_seqlens is None else len(chunk_indices)
    if scale is None:
        scale = k.shape[-1] ** -0.5

    output = torch.empty_like(v)
    if cu_seqlens is None:
        N, chunk_offsets = B, None
    else:
        N, chunk_offsets = (
            len(cu_seqlens) - 1,
            prepare_chunk_offsets(cu_seqlens, BT),
        )

    g = g.transpose(1, 2).contiguous()
    h = h.contiguous()
    CV_kernel_num = 24
    chunk_fwd_kernel_o[(CV_kernel_num,)](
        q,
        k,
        v,
        h,
600
601
602
603
604
605
606
607
        BT=BT,
        BK=128,
        BV=128,
    )
    return output

bwd_chunk_dqkwg = chunk_bwd_dqkwg
bwd_chunk_dv_local = chunk_bwd_dv_local
hyper_parallel/components/functional/_triton/gated_delta_net/chunk_scaled_dot_kkt.py
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
# pylint: disable=unused-argument,invalid-name,missing-module-docstring
# pylint: disable=missing-function-docstring
# pylint: disable=forbidden-backend-import

from typing import Optional

import torch
import triton
import triton.language as tl

from .utils import prepare_chunk_indices


@triton.heuristics({
    'USE_G': lambda args: args['g'] is not None,
    'IS_VARLEN': lambda args: args['cu_seqlens'] is not None,
})
@triton.jit(do_not_specialize=['T', 'NT', 'TOTAL_TASKS'])
def chunk_scaled_dot_kkt_fwd_kernel(
    k,
    g,
    beta,
    A,
 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
    NT,
    B,
    TOTAL_TASKS,
):
    core_id = tl.program_id(0)
    num_blocks = tl.num_programs(0)
    T_max = T

    base_tasks_per_block = TOTAL_TASKS // num_blocks
    remainder_tasks = TOTAL_TASKS % num_blocks

    if core_id < remainder_tasks:
        tasks_this_core = base_tasks_per_block + 1
        start_idx = core_id * tasks_this_core
    else:
        tasks_this_core = base_tasks_per_block
        start_idx = core_id * base_tasks_per_block + remainder_tasks

    for idx in range(start_idx, start_idx + tasks_this_core):
        i_b = idx // NT
        local_idx = idx % NT

        if IS_VARLEN:
            i_n = tl.load(chunk_indices + local_idx * 2).to(tl.int32)
            i_t = tl.load(chunk_indices + local_idx * 2 + 1).to(tl.int32)
            bos = tl.load(cu_seqlens + i_n).to(tl.int32)
            eos = tl.load(cu_seqlens + i_n + 1).to(tl.int32)
            T_local = eos - bos
        else:
            bos, eos = 0, T
            i_t = local_idx
            T_local = T

        for i_h in range(H):
            k_batch_off = i_b * T_max * H * K
            beta_batch_off = i_b * H * T_max
            g_batch_off = i_b * H * T_max
            A_batch_off = i_b * T_max * H * BT

            p_beta = tl.make_block_ptr(beta + beta_batch_off + bos + i_h * T_max, (T_local,), (1,), (i_t * BT,), (BT,), (0,))
            b_beta = tl.load(p_beta, boundary_check=(0,))

            b_A = tl.zeros([BT, BT], dtype=tl.float32)
            for i_k in range(tl.cdiv(K, BK)):
                p_k = tl.make_block_ptr(k + k_batch_off + (bos * H + i_h) * K, (T_local, K), (H * K, 1), (i_t * BT, i_k * BK), (BT, BK), (1, 0))
                b_k = tl.load(p_k, boundary_check=(0, 1))
                dot_product = tl.dot(b_k, tl.trans(b_k))

                o_t = i_t * BT + tl.arange(0, BT)
                o_t = o_t.to(tl.float32)
                T_mask = (o_t < T_local).to(tl.float32)

                row_indices = tl.arange(0, BT)[:, None]
                col_indices = tl.arange(0, BT)[None, :]
                tril_mask = (row_indices > col_indices).to(tl.float32)
                tril_mask = tril_mask * T_mask[:, None]
                masked_dot = dot_product * tril_mask
                b_A += masked_dot

            if USE_G:
                p_g = tl.make_block_ptr(g + g_batch_off + bos + i_h * T_max, (T_local,), (1,), (i_t * BT,), (BT,), (0,))
                b_g = tl.load(p_g, boundary_check=(0,))
                b_g_diff = b_g[:, None] - b_g[None, :]
                b_g_diff = tl.minimum(tl.maximum(b_g_diff, -50.0), 50.0)
                b_A *= tl.exp(b_g_diff)
            b_A *= b_beta[:, None]

            p_A = tl.make_block_ptr(A + A_batch_off + (bos * H + i_h) * BT, (T_local, BT), (BT * H, 1), (i_t * BT, 0), (BT, BT), (1, 0))
            tl.store(p_A, b_A.to(p_A.dtype.element_ty), boundary_check=(0, 1))


@triton.heuristics({
    'IS_VARLEN': lambda args: args['cu_seqlens'] is not None
})
@triton.autotune(
    configs=[
        triton.Config({'BK': BK})
        for BK in [32, 64]
    ],
127
128
129
130
131
132
133
134
135
136
        for BK in [32, 64]
    ],
    key=["BC"]
)
@triton.jit(do_not_specialize=['T'])
def chunk_scaled_dot_kkt_fwd_kernel_intra_sub_inter(
    k,
    g,
    beta,
    A,
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
    BK: tl.constexpr,
    NC: tl.constexpr,
    IS_VARLEN: tl.constexpr,
):
    i_t, i_c, i_b = tl.program_id(0), tl.program_id(1), tl.program_id(2)
    i_i, i_j = i_c // NC, i_c % NC

    for i_h in range(H):
        if IS_VARLEN:
            i_n, i_t = tl.load(chunk_indices + i_t * 2).to(tl.int32), tl.load(chunk_indices + i_t * 2 + 1).to(tl.int32)
            bos, eos = tl.load(cu_seqlens + i_n).to(tl.int32), tl.load(cu_seqlens + i_n + 1).to(tl.int32)
            T_val = eos - bos
        else:
            bos, eos = i_b * T, i_b * T + T
            T_val = T

        should_compute = (i_t * BT + i_i * BC < T_val) and (i_i > i_j)

        if should_compute:
            k_ptr = k + (bos * H + i_h) * K
            g_ptr = g + (bos * H + i_h) * K
            A_ptr = A + (bos * H + i_h) * BT

            p_beta = tl.make_block_ptr(beta + bos * H + i_h, (T_val,), (H,), (i_t * BT + i_i * BC,), (BC,), (0,))
            b_beta = tl.load(p_beta, boundary_check=(0,))

            b_A = tl.zeros([BC, BC], dtype=tl.float32)
            for i_k in range(tl.cdiv(K, BK)):
                p_k = tl.make_block_ptr(k_ptr, (T_val, K), (H * K, 1), (i_t * BT + i_i * BC, i_k * BK), (BC, BK),
                                        (1, 0))
                p_g = tl.make_block_ptr(g_ptr, (T_val, K), (H * K, 1), (i_t * BT + i_i * BC, i_k * BK), (BC, BK),
                                        (1, 0))
                b_kt = tl.make_block_ptr(k_ptr, (K, T_val), (1, H * K), (i_k * BK, i_t * BT + i_j * BC), (BK, BC),
                                         (0, 1))
                p_gk = tl.make_block_ptr(g_ptr, (K, T_val), (1, H * K), (i_k * BK, i_t * BT + i_j * BC), (BK, BC),
                                         (0, 1))

                o_k = i_k * BK + tl.arange(0, BK)
                m_k = o_k < K
                b_gn = tl.load(g_ptr + (i_t * BT + i_i * BC) * H * K + o_k, mask=m_k, other=0)
                b_g = tl.load(p_g, boundary_check=(0, 1))
                b_k = tl.load(p_k, boundary_check=(0, 1)) * tl.exp(b_g - b_gn[None, :])
                b_gk = tl.load(p_gk, boundary_check=(0, 1))
                b_kt = tl.load(b_kt, boundary_check=(0, 1)) * tl.exp(b_gn[:, None] - b_gk)
                b_A += tl.dot(b_k, b_kt)
            b_A *= b_beta[:, None]

            p_A = tl.make_block_ptr(A_ptr, (T_val, BT), (H * BT, 1), (i_t * BT + i_i * BC, i_j * BC), (BC, BC), (1, 0))
            tl.store(p_A, b_A.to(A.dtype.element_ty), boundary_check=(0, 1))


@triton.heuristics({
    'IS_VARLEN': lambda args: args['cu_seqlens'] is not None
})
@triton.jit(do_not_specialize=['T'])
def chunk_scaled_dot_kkt_fwd_kernel_intra_sub_intra(
    k,
    g,
    beta,
    A,
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
    BC: tl.constexpr,
    BK: tl.constexpr,
    IS_VARLEN: tl.constexpr,
):
    i_t, i_i, i_b = tl.program_id(0), tl.program_id(1), tl.program_id(2)

    for i_h in range(H):
        if IS_VARLEN:
            i_n, i_t = tl.load(chunk_indices + i_t * 2).to(tl.int32), tl.load(chunk_indices + i_t * 2 + 1).to(tl.int32)
            bos, eos = tl.load(cu_seqlens + i_n).to(tl.int32), tl.load(cu_seqlens + i_n + 1).to(tl.int32)
            T_val = eos - bos
        else:
            bos, eos = i_b * T, i_b * T + T
            T_val = T

        should_compute = i_t * BT + i_i * BC < T_val

        if should_compute:
            o_i = tl.arange(0, BC)
            o_k = tl.arange(0, BK)
            m_k = o_k < K
            m_A = (i_t * BT + i_i * BC + o_i) < T_val
            o_A = (bos + i_t * BT + i_i * BC + o_i) * H * BT + i_h * BT + i_i * BC

            p_k = tl.make_block_ptr(k + (bos * H + i_h) * K, (T_val, K), (H * K, 1), (i_t * BT + i_i * BC, 0), (BC, BK),
                                    (1, 0))
            p_g = tl.make_block_ptr(g + (bos * H + i_h) * K, (T_val, K), (H * K, 1), (i_t * BT + i_i * BC, 0), (BC, BK),
                                    (1, 0))
            p_beta = beta + (bos + i_t * BT + i_i * BC + o_i) * H + i_h

            b_k = tl.load(p_k, boundary_check=(0, 1)) * tl.load(p_beta, mask=m_A, other=0)[:, None]
            b_g = tl.load(p_g, boundary_check=(0, 1))

            p_kt = k + (bos + i_t * BT + i_i * BC) * H * K + i_h * K + o_k
            p_gk = g + (bos + i_t * BT + i_i * BC) * H * K + i_h * K + o_k

            for j in range(0, min(BC, T_val - i_t * BT - i_i * BC)):
                b_kt = tl.load(p_kt, mask=m_k, other=0).to(tl.float32)
                b_gk = tl.load(p_gk, mask=m_k, other=0).to(tl.float32)
                b_A = tl.sum(b_k * b_kt[None, :] * tl.exp(b_g - b_gk[None, :]), 1)
                # 转化成f32
                o_i_tmp = o_i.to(tl.float32)
                b_A = tl.where(o_i_tmp > j, b_A, 0.)

                tl.store(A + o_A + j, b_A, mask=m_A)
                p_kt += H * K
                p_gk += H * K


def chunk_scaled_dot_kkt_fwd(
    k: torch.Tensor,
    g: Optional[torch.Tensor] = None,
    gk: Optional[torch.Tensor] = None,
    beta: Optional[torch.Tensor] = None,
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306

    Returns:
        beta * K * K^T of shape `[B, T, H, BT]` where `BT` is the chunk size.
    """
    B, T, H, K = k.shape
    BT = chunk_size
    chunk_indices = prepare_chunk_indices(cu_seqlens, BT) if cu_seqlens is not None else None
    NT = triton.cdiv(T, BT) if cu_seqlens is None else len(chunk_indices)
    beta = beta.transpose(1, 2).contiguous()
    g = g.transpose(1, 2).contiguous()
    BK = 128
    kernel_num = 24

    if gk is None:
        A = torch.empty(B, T, H, BT, device=k.device, dtype=output_dtype)
        chunk_scaled_dot_kkt_fwd_kernel[(kernel_num,)](
            k=k,
            g=g,
            beta=beta,
            A=A,
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
            NT=NT,
            B=B,
            TOTAL_TASKS=B * NT,
        )
        return A

    BC = min(16, BT)
    NC = triton.cdiv(BT, BC)
    BK = max(triton.next_power_of_2(K), 16)
    A = torch.zeros(B, T, H, BT, device=k.device, dtype=output_dtype)
    grid = (NT, NC * NC, B)
    chunk_scaled_dot_kkt_fwd_kernel_intra_sub_inter[grid](
        k=k,
        g=gk,
        beta=beta,
        A=A,
336
337
338
339
340
341
342
343
344
345
        BC=BC,
        NC=NC,
    )

    grid = (NT, NC, B)
    chunk_scaled_dot_kkt_fwd_kernel_intra_sub_intra[grid](
        k=k,
        g=gk,
        beta=beta,
        A=A,
351
352
353
354
355
        BT=BT,
        BC=BC,
        BK=BK,
    )
    return A
hyper_parallel/components/functional/_triton/gated_delta_net/cumsum.py
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
# pylint: disable=useless-return,unused-argument,no-else-return,invalid-name
# pylint: disable=missing-module-docstring,missing-function-docstring
# pylint: disable=forbidden-backend-import

from typing import Optional

import torch
import triton
import triton.language as tl

from .utils import prepare_chunk_indices


@triton.heuristics({
    'HAS_SCALE': lambda args: args['scale'] is not None,
    'IS_VARLEN': lambda args: args['cu_seqlens'] is not None
})
@triton.jit(do_not_specialize=['T'])
def chunk_local_cumsum_scalar_kernel(
    s,
    o,
    scale,
    cu_seqlens,
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
    IS_VARLEN: tl.constexpr,
    HEAD_FIRST: tl.constexpr,
    CHUNK_SIZE: tl.constexpr = 64,
):
    i_block, i_b = tl.program_id(0), tl.program_id(1)
    N_CHUNKS: tl.constexpr = BLOCK_T // CHUNK_SIZE

    if IS_VARLEN:
        i_s, i_block = tl.load(chunk_indices + i_block * 2).to(tl.int32), tl.load(
            chunk_indices + i_block * 2 + 1
        ).to(tl.int32)

        bos, eos = tl.load(cu_seqlens + i_s).to(tl.int32), tl.load(
            cu_seqlens + i_s + 1
        ).to(tl.int32)
        T = eos - bos
    else:
        bos, eos = i_b * T, i_b * T + T

    ptr_s = tl.make_block_ptr(
        s + bos * H, (T, H), (H, 1), (i_block * BLOCK_T, 0), (BLOCK_T, H), (1, 0)
    )
    ptr_o = tl.make_block_ptr(
        o + bos * H, (T, H), (H, 1), (i_block * BLOCK_T, 0), (BLOCK_T, H), (1, 0)
    )
    b_s = tl.load(ptr_s, boundary_check=(0,)).to(tl.float32)
    b_s = tl.reshape(b_s, (N_CHUNKS, CHUNK_SIZE, H))
    b_s = tl.trans(b_s, (1, 0, 2))
    b_o = tl.cumsum(b_s, axis=0)
    if REVERSE:
        b_z = tl.sum(b_s, axis=0)
        b_o = -b_o + b_z[None] + b_s
    if HAS_SCALE:
        b_o *= scale
    b_o = tl.trans(b_o, (1, 0, 2))
    b_o = tl.reshape(b_o, (BLOCK_T, H))

    tl.store(ptr_o, b_o.to(ptr_o.dtype.element_ty), boundary_check=(0,))
    return


def chunk_local_cumsum_scalar(
    g: torch.Tensor,
    chunk_size: int,
    reverse: bool = False,
    scale: float = None,
 95
 96
 97
 98
 99
100
101
102
103
104
105
    head_first: bool = False,
    output_dtype: Optional[torch.dtype] = torch.float
) -> torch.Tensor:

    B, T, H = g.shape
    if chunk_size != 2 ** (chunk_size.bit_length() - 1):
        raise ValueError(
            f"chunk_size must be a power of 2, chunk_size is {chunk_size}"
        )
    # We adjust the tiling strategy to prevent overflow in in backward passes and context parallel scenarios
    #  while maximizing UB utilization where possible.
106
107
108
109
110
111
112
113
114
115
116
117
118
119
    # The tiling strategy is as follows:
    # 1. BT must be greater than or equal to chunk_size.
    # 2. UB estimation varies directly with H.
    # 3. BT in reverse mode is smaller than in forward mode.
    BT = max(chunk_size, triton.next_power_of_2((1 << 11 if reverse else 1 << 12) // H))
    chunk_indices = prepare_chunk_indices(cu_seqlens, BT) if cu_seqlens is not None else None
    NT = triton.cdiv(T, BT) if cu_seqlens is None else len(chunk_indices)
    g_org, g = g, torch.empty_like(g, dtype=output_dtype or g.dtype)
    grid = (NT, B)
    chunk_local_cumsum_scalar_kernel[grid](
        s=g_org,
        o=g,
        scale=scale,
        cu_seqlens=cu_seqlens,
125
126
127
128
129
130
131
132
133
134
135
136
        HEAD_FIRST=head_first,
        REVERSE=reverse,
        CHUNK_SIZE=chunk_size,
    )
    return g


def chunk_local_cumsum(
    g: torch.Tensor,
    chunk_size: int,
    reverse: bool = False,
    scale: float = None,
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
    head_first: bool = False,
    output_dtype: Optional[torch.dtype] = torch.float,
    **kwargs
) -> torch.Tensor:
    if cu_seqlens is not None:
        if g.shape[0] != 1:
            raise ValueError(
                "Only batch size 1 is supported when cu_seqlens are provided, "
                f"current size is {g.shape[0]}"
            )
    if len(g.shape) == 3:
        return chunk_local_cumsum_scalar(
            g=g,
            chunk_size=chunk_size,
            reverse=reverse,
            scale=scale,
155
156
157
158
159
160
161
162
163
            head_first=head_first,
            output_dtype=output_dtype
        )
    else:
        raise ValueError(
            f"Unsupported input shape {g.shape}, "
            f"which should be (B, T, H, D) if `head_first=False` "
            f"or (B, H, T, D) otherwise"
        )
hyper_parallel/components/functional/_triton/gated_delta_net/solve_tril.py
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
# pylint: disable=import-outside-toplevel,unused-argument,unused-import
# pylint: disable=missing-module-docstring,missing-function-docstring
# pylint: disable=forbidden-backend-import

import os
from typing import Optional

import torch
import triton
import triton.language as tl

from .utils import prepare_chunk_indices, make_tensor_descriptor, input_guard


def _ensure_slice_ops() -> bool:
    """Probe and attach tl.extract_slice / insert_slice if missing; return success."""
    if hasattr(tl, "extract_slice") and hasattr(tl, "insert_slice"):
        return True
    try:
        from triton.language.extra.cann.extension import extract_slice, insert_slice
        tl.extract_slice = extract_slice
        tl.insert_slice = insert_slice
        return True
    except ImportError:
        return False

_TRITON_SLICE_AVAILABLE: bool = _ensure_slice_ops()
FLA_TRIL_PRECISION = os.environ.get('FLA_TRIL_PRECISION', 'ieee')


@triton.heuristics({"IS_VARLEN": lambda args: args["cu_seqlens"] is not None})
@triton.jit(do_not_specialize=["T"])
def solve_tril_16x16_loop_kernel_paral_v3(
        A_ptr,
        Ad_ptr,
        cu_seqlens,
        chunk_indices,
 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
        LARGE_BLOCK_T: tl.constexpr,
        NT: tl.constexpr,
        BH: tl.constexpr,
):
    worker_id = tl.program_id(0)
    total_tasks = NT * BH
    num_tasks = total_tasks // 48
    remainder = total_tasks - num_tasks * 48
    upper_bound = min(total_tasks, num_tasks * (worker_id + 1) + min(worker_id + 1, remainder))
    lower_bound = num_tasks * worker_id + min(worker_id, remainder)
    for task_id in range(lower_bound, upper_bound):
        i_t = task_id // BH
        i_bh = task_id % BH
        i_b, i_h = i_bh // H, i_bh % H
        if IS_VARLEN:
            i_n, i_t = tl.load(chunk_indices + i_t * 2).to(tl.int32), tl.load(
                chunk_indices + i_t * 2 + 1
            ).to(tl.int32)
            bos, eos = tl.load(cu_seqlens + i_n).to(tl.int32), tl.load(
                cu_seqlens + i_n + 1
            ).to(tl.int32)
            T = eos - bos
        else:
            bos, eos = i_b * T, i_b * T + T

        A = A_ptr + (bos * H + i_h) * BT
        Ad = Ad_ptr + (bos * H + i_h) * 16

        base_t = i_t * LARGE_BLOCK_T

        NTASKS: tl.constexpr = 2
        N_BLOCKS: tl.constexpr = LARGE_BLOCK_T // 16 // NTASKS

        for taskid in range(0, NTASKS):
            base_t += taskid * (LARGE_BLOCK_T // NTASKS)

            b_A = tl.zeros((N_BLOCKS, 16, 16), dtype=tl.float32)  # (N_BLOCKS, 16, 16)
            for blkid in range(0, N_BLOCKS):
                row_start_o = base_t + blkid * 16
                col_start_o = row_start_o % BT
                # using ptr with mask instead of tl.load(block_ptr)
                offs_rows_in_block = tl.arange(0, 16)
                offs_cols_in_block = tl.arange(0, 16)
                ptr_A_subrec16 = (
                        A
                        + row_start_o * H * BT
                        + col_start_o
                        + offs_rows_in_block[:, None] * H * BT
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
                        + col_start_o
                        + offs_rows_in_block[:, None] * H * BT
                        + offs_cols_in_block[None, :]
                )
                global_rows = row_start_o + offs_rows_in_block[:, None]
                global_cols = col_start_o + offs_cols_in_block[None, :]
                load_mask = (global_rows < T) & (global_cols < BT)
                b_A_subrec16 = tl.load(ptr_A_subrec16, mask=load_mask, other=0.0).to(
                    tl.float32
                )
                b_A = tl.insert_slice(
                    ful=b_A,
                    sub=b_A_subrec16[None, :, :],  # (1, 16, 16)
                    offsets=[blkid, 0, 0],
                    sizes=[1, 16, 16],
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
                    strides=[1, 1, 1],
                )

            # load multi 16x16
            local_ori_A = tl.trans(b_A, (1, 0, 2))
            local_ori_A = tl.reshape(local_ori_A, (16, 16 * N_BLOCKS))  # (16, N_BLOCKS*16)

            # change mask into matrix elementwise action
            tmp = tl.arange(0, 16).to(tl.float32)
            rows = tmp[:, None]
            cols = tmp[None, :]
            is_lower = (rows > cols).to(b_A.dtype)
            b_A = -b_A * is_lower

            for i in range(1, 16):
                nblks_vec16 = -tl.extract_slice(
                    local_ori_A, (i, 0), (1, 16 * N_BLOCKS), (16 * N_BLOCKS, 1)
                )
                b_a = tl.reshape(nblks_vec16, (N_BLOCKS, 16))

                dot_tmp = tl.trans(b_a[:, :, None] * b_A, (1, 0, 2))
                dot_product = tl.sum(dot_tmp, 0)
                b_a = b_a + dot_product  # (N_BLOCKS, 16)

                b_a_new_expanded = b_a[:, None, :]  # (N_BLOCKS, 1, 16)
                b_A = tl.insert_slice(
                    ful=b_A,
                    sub=b_a_new_expanded,
                    offsets=[0, i, 0],
                    sizes=[N_BLOCKS, 1, 16],
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
                    sizes=[N_BLOCKS, 1, 16],
                    strides=[1, 1, 1],
                )

            on_diagonal = rows == cols
            b_A = tl.where(on_diagonal, b_A + 1.0, b_A)

            b_A = tl.reshape(b_A, (N_BLOCKS * 16, 16))
            # using ptr with mask instead of tl.load(block_ptr)
            offs_rows_to_store = tl.arange(0, N_BLOCKS * 16)
            offs_cols_to_store = tl.arange(0, 16)
            p_Ai = (
                    Ad
                    + base_t * H * 16
                    + 0
                    + offs_rows_to_store[:, None] * H * 16
165
166
167
168
169
170
171
172
173
174
175
                    + 0
                    + offs_rows_to_store[:, None] * H * 16
                    + offs_cols_to_store[None, :]
            )
            global_store_rows = base_t + offs_rows_to_store[:, None]
            store_mask = global_store_rows < T
            tl.store(
                p_Ai,
                b_A.to(p_Ai.dtype.element_ty, fp_downcast_rounding="rtne"),
                mask=store_mask,
            )
174
175
176
177
178
179
180
181
182
183
184
                mask=store_mask,
            )


@triton.heuristics({"IS_VARLEN": lambda args: args["cu_seqlens"] is not None})
@triton.jit(do_not_specialize=["T", "NT"])
def merge_16x16_to_32x32_loop_inverse_kernel(
        A,
        Ad,
        Ai,
        cu_seqlens,
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
        BT: tl.constexpr,
        IS_VARLEN: tl.constexpr,
        BH: tl.constexpr,
):
    worker_id = tl.program_id(0)
    total_tasks = NT * BH
    num_tasks = total_tasks // 24
    remainder = total_tasks - num_tasks * 24
    upper_bound = min(total_tasks, num_tasks * (worker_id + 1) + min(worker_id + 1, remainder))
    lower_bound = num_tasks * worker_id + min(worker_id, remainder)
    for task_id in range(lower_bound, upper_bound):
        i_tt = task_id // BH
        i_bh = task_id % BH
        i_b, i_h = i_bh // H, i_bh % H
        if IS_VARLEN:
            i_n, i_t = tl.load(chunk_indices + i_tt * 2).to(tl.int32), tl.load(
                chunk_indices + i_tt * 2 + 1
            ).to(tl.int32)
            bos, eos = tl.load(cu_seqlens + i_n).to(tl.int32), tl.load(
                cu_seqlens + i_n + 1
            ).to(tl.int32)
            T = eos - bos
        else:
            bos, eos = i_b * T, i_b * T + T
            i_t = i_tt

        A_ptr = A + (bos * H + i_h) * BT
        Ad_ptr = Ad + (bos * H + i_h) * 16
        Ai_ptr = Ai + (bos * H + i_h) * 32

        p_A_21 = tl.make_block_ptr(
            A_ptr, (T, BT), (H * BT, 1), (i_t * 32 + 16, 0 + i_t % (BT // 32) * 32), (16, 16), (1, 0)
        )
        p_Ad_11 = tl.make_block_ptr(
            Ad_ptr, (T, 16), (H * 16, 1), (i_t * 32, 0), (16, 16), (1, 0)
        )
        p_Ad_22 = tl.make_block_ptr(
            Ad_ptr, (T, 16), (H * 16, 1), (i_t * 32 + 16, 0), (16, 16), (1, 0)
        )
        p_Ai_11 = tl.make_block_ptr(
            Ai_ptr, (T, 32), (H * 32, 1), (i_t * 32, 0), (16, 16), (1, 0)
        )
        p_Ai_22 = tl.make_block_ptr(
            Ai_ptr, (T, 32), (H * 32, 1), (i_t * 32 + 16, 16), (16, 16), (1, 0)
        )
        p_Ai_21 = tl.make_block_ptr(
            Ai_ptr, (T, 32), (H * 32, 1), (i_t * 32 + 16, 0), (16, 16), (1, 0)
        )

        A_21 = tl.load(p_A_21, boundary_check=(0, 1)).to(tl.float32)
        Ai_11 = tl.load(p_Ad_11, boundary_check=(0, 1)).to(tl.float32)
        Ai_22 = tl.load(p_Ad_22, boundary_check=(0, 1)).to(tl.float32)
        Ai_21 = -tl.dot(
            tl.dot(Ai_22, A_21, input_precision="ieee"), Ai_11, input_precision="ieee"
        )
        tl.store(
            p_Ai_11,
            Ai_11.to(p_Ai_11.dtype.element_ty, fp_downcast_rounding="rtne"),
            boundary_check=(0, 1),
        )
        tl.store(
            p_Ai_22,
            Ai_22.to(p_Ai_22.dtype.element_ty, fp_downcast_rounding="rtne"),
            boundary_check=(0, 1),
        )
        tl.store(
            p_Ai_21,
            Ai_21.to(p_Ai_21.dtype.element_ty, fp_downcast_rounding="rtne"),
            boundary_check=(0, 1),
        )
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
            boundary_check=(0, 1),
        )


@triton.heuristics(
    {
        "IS_VARLEN": lambda args: args["cu_seqlens"] is not None,
    }
)
@triton.jit(do_not_specialize=["T", "NT"])
def merge_32x32_to_64x64_loop_inverse_kernel(
        A,
        Ad,
        Ai,
        cu_seqlens,
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
        BT: tl.constexpr,
        IS_VARLEN: tl.constexpr,
        BH: tl.constexpr,
):
    worker_id = tl.program_id(0)
    total_tasks = NT * BH
    num_tasks = total_tasks // 24
    remainder = total_tasks - num_tasks * 24
    upper_bound = min(total_tasks, num_tasks * (worker_id + 1) + min(worker_id + 1, remainder))
    lower_bound = num_tasks * worker_id + min(worker_id, remainder)
    for task_id in range(lower_bound, upper_bound):
        i_tt = task_id // BH
        i_bh = task_id % BH
        i_b, i_h = i_bh // H, i_bh % H
        if IS_VARLEN:
            i_n, i_t = tl.load(chunk_indices + i_tt * 2).to(tl.int32), tl.load(
                chunk_indices + i_tt * 2 + 1
            ).to(tl.int32)
            bos, eos = tl.load(cu_seqlens + i_n).to(tl.int32), tl.load(
                cu_seqlens + i_n + 1
            ).to(tl.int32)
            T = eos - bos
        else:
            bos, eos = i_b * T, i_b * T + T
            i_t = i_tt

        A_ptr = A + (bos * H + i_h) * BT
        Ad_ptr = Ad + (bos * H + i_h) * 32
        Ai_ptr = Ai + (bos * H + i_h) * 64

        p_A_21 = tl.make_block_ptr(
            A_ptr, (T, BT), (H * BT, 1), (i_t * 64 + 32, 0 + i_t % (BT // 64) * 64), (32, 32), (1, 0)
        )

        p_Ad_11 = tl.make_block_ptr(
            Ad_ptr, (T, 32), (H * 32, 1), (i_t * 64, 0), (32, 32), (1, 0)
        )
        p_Ad_22 = tl.make_block_ptr(
            Ad_ptr, (T, 32), (H * 32, 1), (i_t * 64 + 32, 0), (32, 32), (1, 0)
        )

        p_Ai_11 = tl.make_block_ptr(
            Ai_ptr, (T, 64), (H * 64, 1), (i_t * 64, 0), (32, 32), (1, 0)
        )
        p_Ai_22 = tl.make_block_ptr(
            Ai_ptr, (T, 64), (H * 64, 1), (i_t * 64 + 32, 32), (32, 32), (1, 0)
        )
        p_Ai_21 = tl.make_block_ptr(
            Ai_ptr, (T, 64), (H * 64, 1), (i_t * 64 + 32, 0), (32, 32), (1, 0)
        )

        A_21 = tl.load(p_A_21, boundary_check=(0, 1)).to(tl.float32)
        Ai_11 = tl.load(p_Ad_11, boundary_check=(0, 1)).to(tl.float32)
        Ai_22 = tl.load(p_Ad_22, boundary_check=(0, 1)).to(tl.float32)
        Ai_21 = -tl.dot(
            tl.dot(Ai_22, A_21, input_precision="ieee"), Ai_11, input_precision="ieee"
        )
        tl.store(
            p_Ai_11,
            Ai_11.to(p_Ai_11.dtype.element_ty, fp_downcast_rounding="rtne"),
            boundary_check=(0, 1),
        )
        tl.store(
            p_Ai_22,
            Ai_22.to(p_Ai_22.dtype.element_ty, fp_downcast_rounding="rtne"),
            boundary_check=(0, 1),
        )
        tl.store(
            p_Ai_21,
            Ai_21.to(p_Ai_21.dtype.element_ty, fp_downcast_rounding="rtne"),
            boundary_check=(0, 1),
        )
346
347
348
349
350
351
352
353
354
355
356
            boundary_check=(0, 1),
        )


@triton.heuristics({"IS_VARLEN": lambda args: args["cu_seqlens"] is not None})
@triton.jit(do_not_specialize=['T'])
def solve_tril_64x64_kernel(
    A,
    Ai,
    cu_seqlens,
    chunk_indices,
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
    USE_TMA: tl.constexpr,
    IS_VARLEN: tl.constexpr,
    DOT_PRECISION: tl.constexpr
):
    i_t, i_bh = tl.program_id(0), tl.program_id(1)
    i_b, i_h = i_bh // H, i_bh % H
    if IS_VARLEN:
        i_n, i_t = tl.load(chunk_indices + i_t * 2).to(tl.int32), tl.load(chunk_indices + i_t * 2 + 1).to(tl.int32)
        bos, eos = tl.load(cu_seqlens + i_n).to(tl.int32), tl.load(cu_seqlens + i_n + 1).to(tl.int32)
        T = eos - bos
    else:
        bos, eos = i_b * T, i_b * T + T
    o_i = tl.arange(0, 64)
    m_I = o_i[:, None] == o_i[None, :]

    A = A + (bos * H + i_h) * BT
    Ai = Ai + (bos * H + i_h) * 64

    offset = (i_t * 64) % BT
    if not USE_TMA:
        p_A = tl.make_block_ptr(A, (T, BT), (H * BT, 1), (i_t * 64, offset), (64, 64), (1, 0))
        b_A = -tl.load(p_A, boundary_check=(0, 1)).to(tl.float32)
    else:
        desc = make_tensor_descriptor(A, [T, BT], [H * BT, 1], [64, 64])
        desc_o = make_tensor_descriptor(Ai, [T, 64], [H * 64, 1], [64, 64])
        b_A = -desc.load([i_t * 64, offset]).to(tl.float32)

    for i in range(2, min(64, T - i_t * 64)):
        b_a = -tl.load(A + (i_t * 64 + i) * H * BT + o_i + offset)
        b_a = b_a + tl.sum(b_a[:, None] * b_A, 0)
        b_A = tl.where((o_i == i)[:, None], b_a, b_A)
    b_A += m_I
    if not USE_TMA:
        p_Ai = tl.make_block_ptr(Ai, (T, 64), (H * 64, 1), (i_t * 64, 0), (64, 64), (1, 0))
        tl.store(p_Ai, b_A.to(p_Ai.dtype.element_ty, fp_downcast_rounding="rtne"), boundary_check=(0, 1))
    else:
        desc_o.store([i_t * 64, 0], b_A.to(desc_o.dtype, fp_downcast_rounding="rtne"))


def solve_tril_64(
        A: torch.Tensor,
        cu_seqlens: Optional[torch.Tensor] = None,
        output_dtype: torch.dtype = torch.float,
    ):
    B, T, H, BT = A.shape
    chunk_indices = prepare_chunk_indices(cu_seqlens, BT) if cu_seqlens is not None else None
    NT = len(chunk_indices) if cu_seqlens is not None else triton.cdiv(T, BT)

    Ai = torch.zeros_like(A, dtype=output_dtype)
    solve_tril_64x64_kernel[NT, B * H](
        A=A,
        Ai=Ai,
        cu_seqlens=cu_seqlens,
        chunk_indices=chunk_indices,
416
417
418
419
420
421
422
423
424
425
426
427
428
        BT=BT,
        USE_TMA=False,
        DOT_PRECISION=FLA_TRIL_PRECISION,
    )
    return Ai


@input_guard
def solve_tril(
    A: torch.Tensor,
    cu_seqlens: Optional[torch.Tensor] = None,
    output_dtype: torch.dtype = torch.float
) -> torch.Tensor:
441
442
443
444
445
446
447
448
449
450
451
452
453
454
455
456
457
458
459
460
461
462
463
464
465
466
467
468
469
470
471
472
473
474
475
476
477
478

    Returns:
        (I + A)^-1 with the same shape as A
    """
    output_dtype = A.dtype if output_dtype is None else output_dtype
    if not _TRITON_SLICE_AVAILABLE:
        if A.shape[-1] not in [64]:
            raise ValueError(
                f"A shape BT should in [64], but current is {A.shape[-1]}"
            )
        return solve_tril_64(A, cu_seqlens, output_dtype)
    if A.shape[-1] not in [16, 32, 64]:
        raise ValueError(
            f"A shape BT should in [16, 32, 64], but current is {A.shape[-1]}"
        )

    B, T, H, BT = A.shape
    # If BT matches the current processing level (final step), use output_dtype
    # (e.g. BF16) so the kernel can downcast internally, avoiding an extra
    # external cast that hurts performance. Otherwise, keep FP32 to preserve
    # precision for subsequent computation stages.
    Ad = torch.empty(
        B, T, H, 16, device=A.device, dtype=torch.float if BT != 16 else output_dtype
    )

    LARGE_BLOCK_T = 608 * 2

    chunk_indices = (
        prepare_chunk_indices(cu_seqlens, LARGE_BLOCK_T)
        if cu_seqlens is not None
        else None
    )
    NT = len(chunk_indices) if cu_seqlens is not None else triton.cdiv(T, LARGE_BLOCK_T)
    solve_tril_16x16_loop_kernel_paral_v3[(48,)](
        A,
        Ad,
        cu_seqlens=cu_seqlens,
        chunk_indices=chunk_indices,
483
484
485
486
487
488
489
490
491
492
493
494
495
496
497
498
499
500
501
502
503
        NT=NT,
        BH=B * H,
    )

    if BT == 16:
        return Ad

    # Same dtype logic as above: output_dtype for the final step, FP32 otherwise.
    Ai = torch.zeros(
        B, T, H, 32, device=A.device, dtype=torch.float if BT != 32 else output_dtype
    )

    chunk_indices = (
        prepare_chunk_indices(cu_seqlens, 32) if cu_seqlens is not None else None
    )
    NT = len(chunk_indices) if cu_seqlens is not None else triton.cdiv(T, 32)
    merge_16x16_to_32x32_loop_inverse_kernel[(24,)](
        A=A,
        Ad=Ad,
        Ai=Ai,
        cu_seqlens=cu_seqlens,
507
508
509
510
511
512
513
514
515
516
517
518
519
520
521
522
523
524
525
526
527
        BT=BT,
        NT=NT,
        BH=B * H,
    )
    if BT == 32:
        return Ai

    Ad = Ai
    # Same dtype logic as above: output_dtype for the final step, FP32 otherwise.
    Ai = torch.zeros(
        B, T, H, 64, device=A.device, dtype=torch.float if BT != 64 else output_dtype
    )
    chunk_indices = (
        prepare_chunk_indices(cu_seqlens, 64) if cu_seqlens is not None else None
    )
    NT = len(chunk_indices) if cu_seqlens is not None else triton.cdiv(T, 64)
    merge_32x32_to_64x64_loop_inverse_kernel[(24,)](
        A=A,
        Ad=Ad,
        Ai=Ai,
        cu_seqlens=cu_seqlens,
531
532
533
534
535
536
537
        BT=BT,
        NT=NT,
        BH=B * H,
    )
    if BT == 64:
        return Ai
    return Ai
hyper_parallel/components/functional/_triton/gated_delta_net/state_summary.py
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
# pylint: disable=missing-public-type-hints,invalid-name

"""Fixed-shape Triton-Ascend kernels for GDN state summaries."""

__all__ = ["gdn_packed_state_summary_kernel", "gdn_state_grad_ext_kernel"]

import triton
import triton.language as tl

from .utils import get_autotune_config


@triton.autotune(
    configs=get_autotune_config(
        multibuffer_list=(True, False),
        set_workspace_multibuffer_list=(2, 4),
        tile_mix_vector_loop_num_list=(2,),
36
37
38
39
40
41
42
43
44
45
        tile_mix_cube_loop_num_list=(2,),
    ),
    key=["H", "K", "V", "BT", "BV"],
)
@triton.jit(do_not_specialize=["T"])
def gdn_packed_state_summary_kernel(
    k,
    w,
    u,
    g,
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
    BV: tl.constexpr,
    NT: tl.constexpr,
):
    """Build the local affine state transition and extension in one buffer."""
    i_v = tl.program_id(0)
    i_bh = tl.program_id(1)
    i_b = i_bh // H
    i_h = i_bh % H

    stride_k = H * K
    stride_v = H * V
    k += (i_b * T * H + i_h) * K
    w += (i_b * T * H + i_h) * K
    u += (i_b * T * H + i_h) * V
    g += i_b * T * H + i_h
    packed_summary += i_bh * K * (V + K)

    col = tl.arange(0, BV)
    row1 = tl.arange(0, 64)
    row2 = 64 + tl.arange(0, 64)
    is_transition = i_v * BV >= V
    transition_col = i_v * BV - V + col
    b_h1 = tl.where(
        is_transition & (row1[:, None] == transition_col[None, :]), 1.0, 0.0
    ).to(tl.float32)
    b_h2 = tl.where(
        is_transition & (row2[:, None] == transition_col[None, :]), 1.0, 0.0
    ).to(tl.float32)

    for i_t in range(NT):
        p_w1 = tl.make_block_ptr(
            w, (T, K), (stride_k, 1), (i_t * BT, 0), (BT, 64), (1, 0)
        )
        p_w2 = tl.make_block_ptr(
            w, (T, K), (stride_k, 1), (i_t * BT, 64), (BT, 64), (1, 0)
        )
        b_w1 = tl.load(p_w1, boundary_check=(0, 1))
        b_w2 = tl.load(p_w2, boundary_check=(0, 1))
        b_v = tl.dot(b_w1, b_h1.to(b_w1.dtype))
        b_v += tl.dot(b_w2, b_h2.to(b_w2.dtype))

        p_u = tl.make_block_ptr(
            u,
            (T, V),
            (stride_v, 1),
            (i_t * BT, i_v * BV),
 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
            (i_t * BT, i_v * BV),
            (BT, BV),
            (1, 0),
        )
        b_v = tl.load(p_u, boundary_check=(0, 1)) - b_v

        last_idx = min((i_t + 1) * BT, T) - 1
        token = i_t * BT + tl.arange(0, BT)
        mask = token < T
        b_g_last = tl.load(g + last_idx * H).to(tl.float32)
        b_g = tl.load(g + token * H, mask=mask, other=0.0).to(tl.float32)
        b_v *= tl.where(mask, tl.exp(b_g_last - b_g), 0.0)[:, None]
        decay = tl.exp(b_g_last)
        b_h1 *= decay
        b_h2 *= decay
        b_v = b_v.to(k.dtype.element_ty)

        p_k1 = tl.make_block_ptr(
            k, (K, T), (1, stride_k), (0, i_t * BT), (64, BT), (0, 1)
        )
        p_k2 = tl.make_block_ptr(
            k, (K, T), (1, stride_k), (64, i_t * BT), (64, BT), (0, 1)
        )
        b_h1 += tl.dot(tl.load(p_k1, boundary_check=(0, 1)), b_v)
        b_h2 += tl.dot(tl.load(p_k2, boundary_check=(0, 1)), b_v)

    p_out1 = tl.make_block_ptr(
        packed_summary,
        (K, V + K),
        (V + K, 1),
        (0, i_v * BV),
127
128
129
130
131
132
133
134
135
        (0, i_v * BV),
        (64, BV),
        (1, 0),
    )
    p_out2 = tl.make_block_ptr(
        packed_summary,
        (K, V + K),
        (V + K, 1),
        (64, i_v * BV),
135
136
137
138
139
140
141
142
143
144
145
146
147
        (64, i_v * BV),
        (64, BV),
        (1, 0),
    )
    tl.store(p_out1, b_h1, boundary_check=(0, 1))
    tl.store(p_out2, b_h2, boundary_check=(0, 1))


@triton.autotune(
    configs=get_autotune_config(
        multibuffer_list=(True, False),
        set_workspace_multibuffer_list=(2, 4),
        tile_mix_vector_loop_num_list=(2,),
148
149
150
151
152
153
154
155
156
157
        tile_mix_cube_loop_num_list=(2,),
    ),
    key=["H", "K", "V", "BT"],
)
@triton.jit(do_not_specialize=["T"])
def gdn_state_grad_ext_kernel(
    q,
    k,
    w,
    g,
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
    BV: tl.constexpr,
    NT: tl.constexpr,
):
    """Build the local-loss contribution to the incoming state gradient."""
    i_v = tl.program_id(0)
    i_bh = tl.program_id(1)
    i_b = i_bh // H
    i_h = i_bh % H

    stride_k = H * K
    stride_v = H * V
    q += (i_b * T * H + i_h) * K
    k += (i_b * T * H + i_h) * K
    w += (i_b * T * H + i_h) * K
    g += i_b * T * H + i_h
    do += (i_b * T * H + i_h) * V
    dv += (i_b * T * H + i_h) * V
    grad_state_ext += i_bh * K * V

    b_dh1 = tl.zeros([64, BV], dtype=tl.float32)
    b_dh2 = tl.zeros([64, BV], dtype=tl.float32)

    for reverse_idx in range(NT):
        i_t = NT - 1 - reverse_idx
        last_idx = min((i_t + 1) * BT, T) - 1
        token = i_t * BT + tl.arange(0, BT)
        mask = token < T
        b_g_last = tl.load(g + last_idx * H).to(tl.float32)
        b_g = tl.load(g + token * H, mask=mask, other=0.0).to(tl.float32)

        p_k1 = tl.make_block_ptr(
            k, (T, K), (stride_k, 1), (i_t * BT, 0), (BT, 64), (1, 0)
        )
        p_k2 = tl.make_block_ptr(
            k, (T, K), (stride_k, 1), (i_t * BT, 64), (BT, 64), (1, 0)
        )
        b_k1 = tl.load(p_k1, boundary_check=(0, 1))
        b_k2 = tl.load(p_k2, boundary_check=(0, 1))
        b_dv = tl.dot(b_k1, b_dh1.to(b_k1.dtype))
        b_dv += tl.dot(b_k2, b_dh2.to(b_k2.dtype))
        b_dv *= tl.where(mask, tl.exp(b_g_last - b_g), 0.0)[:, None]

        p_dv = tl.make_block_ptr(
            dv,
            (T, V),
            (stride_v, 1),
            (i_t * BT, i_v * BV),
213
214
215
216
217
218
219
220
221
222
223
            (i_t * BT, i_v * BV),
            (BT, BV),
            (1, 0),
        )
        b_dv += tl.load(p_dv, boundary_check=(0, 1))

        p_do = tl.make_block_ptr(
            do,
            (T, V),
            (stride_v, 1),
            (i_t * BT, i_v * BV),
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
            (i_t * BT, i_v * BV),
            (BT, BV),
            (1, 0),
        )
        b_do = tl.load(p_do, boundary_check=(0, 1))
        decay = tl.exp(b_g_last)
        b_dh1 *= decay
        b_dh2 *= decay

        p_q1 = tl.make_block_ptr(
            q, (K, T), (1, stride_k), (0, i_t * BT), (64, BT), (0, 1)
        )
        p_q2 = tl.make_block_ptr(
            q, (K, T), (1, stride_k), (64, i_t * BT), (64, BT), (0, 1)
        )
        p_w1 = tl.make_block_ptr(
            w, (K, T), (1, stride_k), (0, i_t * BT), (64, BT), (0, 1)
        )
        p_w2 = tl.make_block_ptr(
            w, (K, T), (1, stride_k), (64, i_t * BT), (64, BT), (0, 1)
        )
        b_q1 = tl.load(p_q1, boundary_check=(0, 1))
        b_q2 = tl.load(p_q2, boundary_check=(0, 1))
        b_w1 = tl.load(p_w1, boundary_check=(0, 1))
        b_w2 = tl.load(p_w2, boundary_check=(0, 1))
        gate = tl.exp(b_g)[None, :]
        b_q1 *= gate
        b_q2 *= gate
        b_dh1 += tl.dot(b_q1, b_do.to(b_q1.dtype)) * scale
        b_dh1 -= tl.dot(b_w1, b_dv.to(b_w1.dtype))
        b_dh2 += tl.dot(b_q2, b_do.to(b_q2.dtype)) * scale
        b_dh2 -= tl.dot(b_w2, b_dv.to(b_w2.dtype))

    p_out1 = tl.make_block_ptr(
        grad_state_ext, (K, V), (V, 1), (0, i_v * BV), (64, BV), (1, 0)
    )
    p_out2 = tl.make_block_ptr(
        grad_state_ext, (K, V), (V, 1), (64, i_v * BV), (64, BV), (1, 0)
    )
    tl.store(p_out1, b_dh1, boundary_check=(0, 1))
    tl.store(p_out2, b_dh2, boundary_check=(0, 1))
hyper_parallel/components/functional/_triton/gated_delta_net/utils.py
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
# pylint: disable=invalid-name,missing-module-docstring,missing-function-docstring
# pylint: disable=missing-class-docstring,broad-exception-caught,protected-access
# pylint: disable=forbidden-backend-import

import itertools
import contextlib
import os
import functools
import warnings
import logging
from enum import Enum
from functools import lru_cache
from typing import Any, Callable, Optional
from packaging import version

import torch
import triton
import triton.language as tl
import triton.language.extra.libdevice as tldevice
import triton.runtime.driver as driver

logger = logging.getLogger(__name__)

FLA_CI_ENV = os.getenv("FLA_CI_ENV") == "1"


def tensor_cache(fn: Optional[Callable[..., torch.Tensor]] = None, *, maxsize: int = 1) -> Any:
    """
    A decorator that caches the most recent results of a function with tensor inputs.

    This decorator will store the outputs of the decorated function for the most recent
 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
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
    Returns:
        Callable[..., torch.Tensor]:
        A wrapped version of the input function with caching.
    """
    if maxsize < 1:
        raise ValueError("maxsize must be at least 1")

    def _is_match(a: Any, b: Any) -> bool:
        if isinstance(a, torch.Tensor) and isinstance(b, torch.Tensor):
            return a is b
        try:
            return a == b
        except Exception:
            return a is b

    def _make_wrapper(fn: Callable[..., torch.Tensor]) -> Callable[..., torch.Tensor]:
        cache: list = []

        @functools.wraps(fn)
        def wrapper(*args: Any, **kwargs: Any) -> Any:
            for i, (cached_args, cached_kwargs, cached_result) in enumerate(cache):
                if len(args) == len(cached_args) and len(kwargs) == len(cached_kwargs):
                    if all(_is_match(a, b) for a, b in zip(args, cached_args)) and all(
                        k in cached_kwargs and _is_match(v, cached_kwargs[k]) for k, v in kwargs.items()
                    ):
                        if i != 0:
                            cache.insert(0, cache.pop(i))
                        return cached_result

            result = fn(*args, **kwargs)
            cache.insert(0, (args, kwargs, result))
            if len(cache) > maxsize:
                cache.pop()
            return result

        return wrapper

    if fn is not None:
        return _make_wrapper(fn)
    return _make_wrapper


@tensor_cache
def prepare_lens(cu_seqlens: torch.LongTensor) -> torch.LongTensor:
    return cu_seqlens[1:] - cu_seqlens[:-1]


@tensor_cache(maxsize=3)
def prepare_chunk_indices(cu_seqlens: torch.LongTensor, chunk_size: int) -> torch.LongTensor:
    indices = torch.cat([torch.arange(n) for n in triton.cdiv(prepare_lens(cu_seqlens), chunk_size).tolist()])
    return torch.stack([indices.eq(0).cumsum(0) - 1, indices], 1).to(cu_seqlens)


def get_abs_err(x, y):
    return (x.detach() - y.detach()).flatten().abs().max().item()


def get_err_ratio(x, y):
    err = (x.detach() - y.detach()).flatten().square().mean().sqrt().item()
    base = (x.detach()).flatten().square().mean().sqrt().item()
    return err / (base + 1e-8)


def assert_close(prefix, ref, tri, ratio, warning=False, err_atol=1e-6):
    abs_atol = get_abs_err(ref, tri)
    msg = f"{prefix:>16} diff: {abs_atol:.6f} ratio: {get_err_ratio(ref, tri):.6f}"
    logger.info(msg)
    error_rate = get_err_ratio(ref, tri)
    if abs_atol <= err_atol:
        return
    allow_warning = warning or (FLA_CI_ENV and (error_rate < 0.01 or abs_atol <= 0.3))
    if allow_warning:
        if error_rate > ratio:
            warnings.warn(msg)
    else:
        if not error_rate < ratio:
            raise AssertionError(msg)


if hasattr(triton.language, '_experimental_make_tensor_descriptor'):
    # For Triton 3.3.x
    make_tensor_descriptor = triton.language._experimental_make_tensor_descriptor
elif hasattr(triton.language, 'make_tensor_descriptor'):
    # For Triton 3.4.x and later
    make_tensor_descriptor = triton.language.make_tensor_descriptor
else:
    """
    Fallback implementation when TMA is not supported.
    Returns None to indicate TMA descriptors are unavailable.
    Just make triton compiler happy.
    """
151
152
153
154
155
156
157
158
159
160
    Returns None to indicate TMA descriptors are unavailable.
    Just make triton compiler happy.
    """

    @triton.jit
    def make_tensor_descriptor(
        base,
        shape,
        strides,
        block_shape,
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
        strides,
        block_shape,
        _builder=None,
    ):
        return None


@lru_cache(maxsize=None)
def get_available_device() -> str:
    try:
        return triton.runtime.driver.active.get_current_target().backend
    except Exception:
        _cpu_device_warning()
        return 'cpu'


def map_triton_backend_to_torch_device() -> str:
    backend = get_available_device()  # 'cuda' | 'hip' | 'xpu' | 'cpu' | ...
    return {'cuda': 'cuda', 'hip': 'cuda', 'xpu': 'xpu'}.get(backend, backend)


device = get_available_device() if get_available_device() != 'hip' else 'cuda'
device_torch_lib = getattr(torch, device)
device_platform = get_available_device()
is_amd = device_platform == 'hip'
is_nvidia = device_platform == 'cuda'
is_nvidia_hopper = is_nvidia and (
    'NVIDIA H' in torch.cuda.get_device_name(0) or torch.cuda.get_device_capability()[0] >= 9
)

is_tf32_supported = is_nvidia and torch.cuda.get_device_capability(0)[0] >= 8
is_tma_supported = (
    (is_nvidia and torch.cuda.get_device_capability(0)[0] >= 9)
    and os.environ.get('FLA_NO_USE_TMA', '0') != '1'
    and (
        hasattr(triton.language, '_experimental_make_tensor_descriptor')
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
        or hasattr(triton.language, 'make_tensor_descriptor')
    )
)

if is_nvidia and not is_tf32_supported:
    # Make old card happy, since triton will use tf32 by default.
    # This is a workaround for old nvidia card.
    os.environ['TRITON_F32_DEFAULT'] = 'ieee'


@lru_cache(maxsize=None)
def check_pytorch_version(version_s: str = '2.4') -> bool:
    return version.parse(torch.__version__) >= version.parse(version_s)


if check_pytorch_version('2.4'):
    device = 'cuda' if device == 'cpu' else device
    autocast_custom_fwd = functools.partial(torch.amp.custom_fwd, device_type=device)
    autocast_custom_bwd = functools.partial(torch.amp.custom_bwd, device_type=device)

    def custom_device_ctx(index: int):
        return device_torch_lib.device(index)
else:
    if device != 'cuda':
        raise AssertionError('Only cuda device is supported for PyTorch version < 2.4.0.')
    autocast_custom_fwd = device_torch_lib.amp.custom_fwd
    autocast_custom_bwd = device_torch_lib.amp.custom_bwd

    def custom_device_ctx(index: int):
        return torch.cuda.device(index)


def input_guard(fn: Callable[..., torch.Tensor]) -> Callable[..., torch.Tensor]:
    """
    A decorator to make sure all input tensors are contiguous and set the device based on input tensors.
    """

    @functools.wraps(fn)
    def wrapper(*args, **kwargs):
        contiguous_args = (i if not isinstance(i, torch.Tensor) else i.contiguous() for i in args)
        contiguous_kwargs = {k: (v if not isinstance(v, torch.Tensor) else v.contiguous()) for k, v in kwargs.items()}

        tensor = None
        for arg in args:
            if isinstance(arg, torch.Tensor):
                tensor = arg
                break
        if tensor is None:
            for value in kwargs.values():
                if isinstance(value, torch.Tensor):
                    tensor = value
                    break

        if tensor is not None:
            ctx = custom_device_ctx(tensor.device.index)
        else:
            ctx = contextlib.nullcontext()

        with ctx:
            return fn(*contiguous_args, **contiguous_kwargs)

    return wrapper


def _cpu_device_warning():
    warnings.warn(('Triton is not supported on current platform, roll back to CPU.'), stacklevel=1)


@tensor_cache
def prepare_chunk_offsets(cu_seqlens: torch.LongTensor, chunk_size: int) -> torch.LongTensor:
    return torch.cat([cu_seqlens.new_tensor([0]), triton.cdiv(prepare_lens(cu_seqlens), chunk_size)]).cumsum(-1)


if os.environ.get('FLA_USE_FAST_OPS', '0') == '1':
    exp = tldevice.fast_expf
    exp2 = tldevice.exp2
    log = tldevice.fast_logf
    log2 = tldevice.fast_log2f
else:
    exp = tl.exp
    exp2 = tl.math.exp2
    log = tl.log
    log2 = tl.log2


def get_all_max_shared_mem():
    try:
        return [
            triton.runtime.driver.active.utils.get_device_properties(i)['max_shared_mem']
            for i in range(device_torch_lib.device_count())
        ]
    except Exception:
        _cpu_device_warning()
        return [-1]


class Backend(Enum):
    ADA = 101376  # RTX 4090
    AMPERE = 166912  # A100
    HOPPER = 232448  # H100
    DEFAULT = 102400  # Default

    @classmethod
    def get_shared_memory(cls, arch: str) -> int:
        try:
            return cls[arch.upper()].value
        except KeyError:
            return cls.DEFAULT.value


@lru_cache(maxsize=None)
def check_shared_mem(arch: str = "none", tensor_idx: int = 0) -> bool:
    try:
        device_shared_mem_list = get_all_max_shared_mem()
        max_shared_memory = device_shared_mem_list[tensor_idx]
        return max_shared_memory >= Backend.get_shared_memory(arch)
    except Exception:
        return False


def get_autotune_config(
    multibuffer_list: tuple = (False,),
    unit_flag_list: tuple = (False,),
    limit_auto_multi_buffer_only_for_local_buffer_list: tuple = (False,),
    limit_auto_multi_buffer_of_local_buffer_list: tuple = ("no-l0c",),
321
322
323
324
325
326
327
328
329
330
    enable_hivm_auto_cv_balance_list: tuple = (True,),
    tile_mix_vector_loop_num_list: tuple = (2, 4),
    tile_mix_cube_loop_num_list: tuple = (2, 4),
):
    configs = []
    for (
        multibuffer,
        unit_flag,
        limit_auto_multi_buffer_only_for_local_buffer,
        limit_auto_multi_buffer_of_local_buffer,
333
334
335
336
337
338
339
340
341
        list(unit_flag_list),
        list(limit_auto_multi_buffer_only_for_local_buffer_list),
        list(limit_auto_multi_buffer_of_local_buffer_list),
    ):
        base_config_dict = {
            'multibuffer': multibuffer,
            'unit_flag': unit_flag,
            'limit_auto_multi_buffer_only_for_local_buffer': limit_auto_multi_buffer_only_for_local_buffer,
            'limit_auto_multi_buffer_of_local_buffer': limit_auto_multi_buffer_of_local_buffer,
340
341
342
343
344
345
346
347
348
349
350
351
            'limit_auto_multi_buffer_only_for_local_buffer': limit_auto_multi_buffer_only_for_local_buffer,
            'limit_auto_multi_buffer_of_local_buffer': limit_auto_multi_buffer_of_local_buffer,
        }

        if limit_auto_multi_buffer_only_for_local_buffer:
            configs.append(triton.Config(base_config_dict))
        else:
            for (
                set_workspace_multibuffer,
                enable_hivm_auto_cv_balance,
                tile_mix_vector_loop,
                tile_mix_cube_loop,
354
355
356
357
358
359
360
361
362
363
                list(enable_hivm_auto_cv_balance_list),
                list(tile_mix_vector_loop_num_list),
                list(tile_mix_cube_loop_num_list),
            ):
                full_config_dict = base_config_dict.copy()
                full_config_dict.update(
                    {
                        'set_workspace_multibuffer': set_workspace_multibuffer,
                        'enable_hivm_auto_cv_balance': enable_hivm_auto_cv_balance,
                        'tile_mix_vector_loop': tile_mix_vector_loop,
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
                        'tile_mix_vector_loop': tile_mix_vector_loop,
                        'tile_mix_cube_loop': tile_mix_cube_loop,
                    }
                )
                configs.append(triton.Config(full_config_dict))
    return configs


def get_npu_properties():
    return driver.active.utils.get_device_properties(torch.npu.current_device())


@functools.cache
def get_vector_num() -> int:
    import torch_npu

    current_device = torch_npu.npu.current_device()
    properties = driver.active.utils.get_device_properties(current_device)
    return properties["num_vectorcore"]


@lru_cache
def is_arch35():
    try:
        import torch_npu

        return "Ascend910_95" in torch_npu.npu.get_device_name() or "Ascend950" in torch_npu.npu.get_device_name()
    except Exception:
        return False
hyper_parallel/components/functional/_triton/gated_delta_net/wy_fast.py
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
# pylint: disable=line-too-long,missing-public-type-hints,missing-public-docstring
# pylint: disable=invalid-name,missing-module-docstring,missing-function-docstring
# pylint: disable=forbidden-backend-import

from typing import Optional, Tuple

import torch
import triton
import triton.language as tl

from .utils import prepare_chunk_indices, exp


@triton.heuristics({
    'IS_VARLEN': lambda args: args['cu_seqlens'] is not None
})
@triton.jit(do_not_specialize=['T'])
def prepare_wy_repr_bwd_kernel(
        k,
        v,
        beta,
        g,
 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
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
        BK: tl.constexpr,
        BV: tl.constexpr,
        IS_VARLEN: tl.constexpr
):
    core_id = tl.program_id(0)
    total_cores = tl.num_programs(0)
    T_max = T

    base_chunks_per_pid = NT // total_cores
    remainder_chunks = NT % total_cores

    if core_id < remainder_chunks:
        chunks_this_pid = base_chunks_per_pid + 1
        start_idx = core_id * chunks_this_pid
    else:
        chunks_this_pid = base_chunks_per_pid
        start_idx = core_id * chunks_this_pid + remainder_chunks

    for idx in range(start_idx, start_idx + chunks_this_pid):
        for i_b in range(B):
            if IS_VARLEN:
                i_n, i_t = tl.load(chunk_indices + idx * 2).to(tl.int32), tl.load(chunk_indices + idx * 2 + 1).to(tl.int32)
                bos, eos = tl.load(cu_seqlens + i_n).to(tl.int32), tl.load(cu_seqlens + i_n + 1).to(tl.int32)
                T = eos - bos
            else:
                i_t = idx
                bos, eos = i_b * T, i_b * T + T

            o_t = i_t * BT + tl.arange(0, BT)
            m_t = o_t < T
            m_A = (o_t[:, None] > o_t[None, :]) & (m_t[:, None] & m_t)
            for i_h in range(0, H):
                if IS_VARLEN:
                    offset = bos + i_h * T_max
                else:
                    offset = bos * H + i_h * T_max

                p_beta = tl.make_block_ptr(beta + offset, (T,), (1,), (i_t * BT,), (BT,), (0,))
                p_g = tl.make_block_ptr(g + offset, (T,), (1,), (i_t * BT,), (BT,), (0,))
                p_A = tl.make_block_ptr(A + (bos * H + i_h) * BT, (BT, T), (1, H * BT), (0, i_t * BT), (BT, BT), (0, 1))

                b_A = tl.load(p_A, boundary_check=(0, 1))
                b_beta = tl.load(p_beta, boundary_check=(0,))
                b_g = tl.load(p_g, boundary_check=(0,))
                b_g_exp = tl.exp(b_g)

                b_dbeta = tl.zeros([BT], dtype=tl.float32)
                b_dA = tl.zeros([BT, BT], dtype=tl.float32)
                b_dg = tl.zeros([BT], dtype=tl.float32)

                for i_k in range(tl.cdiv(K, BK)):
                    p_k = tl.make_block_ptr(k + (bos * H + i_h) * K, (T, K), (H * K, 1), (i_t * BT, i_k * BK), (BT, BK), (1, 0))
                    p_dk = tl.make_block_ptr(dk + (bos * H + i_h) * K, (T, K), (H * K, 1), (i_t * BT, i_k * BK), (BT, BK), (1, 0))
                    p_dw = tl.make_block_ptr(dw + (bos * H + i_h) * K, (T, K), (H * K, 1), (i_t * BT, i_k * BK), (BT, BK), (1, 0))
                    b_k = tl.load(p_k, boundary_check=(0, 1))
                    b_k_beta_g = (b_k * b_beta[:, None] * b_g_exp[:, None]).to(b_k.dtype)
                    b_dw = tl.load(p_dw, boundary_check=(0, 1))
                    b_dA += tl.dot(b_dw, tl.trans(b_k_beta_g))
                    b_dk_beta_g = tl.dot(b_A, b_dw)
                    b_dk = b_dk_beta_g * b_beta[:, None] * b_g_exp[:, None]
                    b_dbeta += tl.sum(b_dk_beta_g * b_k * b_g_exp[:, None], 1)
                    b_dg += tl.sum(b_dk_beta_g * b_k * b_g_exp[:, None] * b_beta[:, None], 1)
                    tl.store(p_dk, b_dk.to(p_dk.dtype.element_ty), boundary_check=(0, 1))

                for i_v in range(tl.cdiv(V, BV)):
                    p_v = tl.make_block_ptr(v + (bos * H + i_h) * V, (T, V), (H * V, 1), (i_t * BT, i_v * BV), (BT, BV), (1, 0))
                    p_dv = tl.make_block_ptr(dv + (bos * H + i_h) * V, (T, V), (H * V, 1), (i_t * BT, i_v * BV), (BT, BV), (1, 0))
                    p_du = tl.make_block_ptr(du + (bos * H + i_h) * V, (T, V), (H * V, 1), (i_t * BT, i_v * BV), (BT, BV), (1, 0))
                    b_v = tl.load(p_v, boundary_check=(0, 1))
                    b_v_beta = (b_v * b_beta[:, None]).to(b_v.dtype)
                    b_du = tl.load(p_du, boundary_check=(0, 1))
                    b_dA += tl.dot(b_du, tl.trans(b_v_beta))
                    b_dv_beta = tl.dot(b_A, b_du)
                    b_dv = b_dv_beta * b_beta[:, None]
                    b_dbeta += tl.sum(b_dv_beta * b_v, 1)
                    tl.store(p_dv, b_dv.to(p_dv.dtype.element_ty), boundary_check=(0, 1))

                b_dA = tl.where(m_A, b_dA, 0)
                b_dA = tl.dot(b_dA.to(b_A.dtype), b_A)
                b_dA = tl.dot(b_A, b_dA.to(b_A.dtype))
                b_dA = tl.where(m_A, -b_dA * exp(b_g[:, None] - b_g[None, :]), 0)
                b_dA = b_dA.to(k.dtype.element_ty)
                b_A = tl.zeros([BT, BT], dtype=tl.float32)

                for i_k in range(tl.cdiv(K, BK)):
                    p_k = tl.make_block_ptr(k + (bos * H + i_h) * K, (T, K), (H * K, 1), (i_t * BT, i_k * BK), (BT, BK), (1, 0))
                    p_dk = tl.make_block_ptr(dk + (bos * H + i_h) * K, (T, K), (H * K, 1), (i_t * BT, i_k * BK), (BT, BK), (1, 0))
                    b_k = tl.load(p_k, boundary_check=(0, 1))
                    b_dk = tl.load(p_dk, boundary_check=(0, 1))
                    b_k_beta = (b_k * b_beta[:, None]).to(b_k.dtype)
                    b_A += tl.dot(b_k_beta, tl.trans(b_k))
                    b_dk_beta = tl.dot(b_dA, b_k)
                    b_dbeta += tl.sum(b_dk_beta * b_k, 1)
                    b_dk += tl.dot(tl.trans(b_dA), b_k_beta)
                    b_dk += b_dk_beta * b_beta[:, None]
                    tl.store(p_dk, b_dk.to(p_dk.dtype.element_ty), boundary_check=(0, 1))

                b_dA_A = b_dA * b_A
                b_dg += tl.sum(b_dA_A, axis=1) - tl.sum(b_dA_A, axis=0)
                p_dg = tl.make_block_ptr(dg + offset, (T,), (1,), (i_t * BT,), (BT,), (0,))
                p_dbeta = tl.make_block_ptr(dbeta + offset, (T,), (1,), (i_t * BT,), (BT,), (0,))
                tl.store(p_dg, b_dg.to(p_dg.dtype.element_ty), boundary_check=(0,))
                tl.store(p_dbeta, b_dbeta.to(p_dbeta.dtype.element_ty), boundary_check=(0,))


@triton.heuristics({
    'USE_G': lambda args: args['g'] is not None,
    'USE_GK': lambda args: args['gk'] is not None,
    'IS_VARLEN': lambda args: args['cu_seqlens'] is not None
})
@triton.jit(do_not_specialize=['T'])
def recompute_w_u_fwd_kernel(
        k,
        v,
        beta,
        w,
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
        USE_G: tl.constexpr,
        USE_GK: tl.constexpr,
        IS_VARLEN: tl.constexpr
):
    core_id = tl.program_id(0)
    total_cores = tl.num_programs(0)
    T_max = T_tmp

    base_chunks_per_pid = NT // total_cores
    remainder_chunks = NT % total_cores

    if core_id < remainder_chunks:
        chunks_this_pid = base_chunks_per_pid + 1
        start_idx = core_id * chunks_this_pid
    else:
        chunks_this_pid = base_chunks_per_pid
        start_idx = core_id * chunks_this_pid + remainder_chunks

    for idx in range(start_idx, start_idx + chunks_this_pid):
        for i_b in range(B):
            for i_h in range(0, H):

                if IS_VARLEN:
                    i_n, i_t = tl.load(chunk_indices + idx * 2).to(tl.int32), tl.load(chunk_indices + idx * 2 + 1).to(tl.int32)
                    bos, eos = tl.load(cu_seqlens + i_n).to(tl.int32), tl.load(cu_seqlens + i_n + 1).to(tl.int32)
                    offset = bos + i_h * T_max
                    T = eos - bos
                else:
                    T = T_tmp
                    i_t = idx
                    bos, eos = i_b * T, i_b * T + T
                    offset = bos * H + i_h * T_max

                p_beta = tl.make_block_ptr(beta + offset, (T,), (1,), (i_t * BT,), (BT,), (0,))
                b_beta = tl.load(p_beta, boundary_check=(0,))

                p_A = tl.make_block_ptr(A + (bos * H + i_h) * BT, (T, BT), (H * BT, 1), (i_t * BT, 0), (BT, BT), (1, 0))
                b_A = tl.load(p_A, boundary_check=(0, 1))

                for i_v in range(tl.cdiv(V, BV)):
                    p_v = tl.make_block_ptr(v + (bos * H + i_h) * V, (T, V), (H * V, 1), (i_t * BT, i_v * BV), (BT, BV), (1, 0))
                    p_u = tl.make_block_ptr(u + (bos * H + i_h) * V, (T, V), (H * V, 1), (i_t * BT, i_v * BV), (BT, BV), (1, 0))
                    b_v = tl.load(p_v, boundary_check=(0, 1))
                    b_vb = (b_v * b_beta[:, None]).to(b_v.dtype)
                    b_u = tl.dot(b_A, b_vb, allow_tf32=False)
                    tl.store(p_u, b_u.to(p_u.dtype.element_ty), boundary_check=(0, 1))

                if USE_G:
                    p_g = tl.make_block_ptr(g + offset, (T,), (1,), (i_t * BT,), (BT,), (0,))
                    b_g = tl.exp(tl.load(p_g, boundary_check=(0,)))

                for i_k in range(tl.cdiv(K, BK)):
                    p_k = tl.make_block_ptr(k + (bos * H + i_h) * K, (T, K), (H * K, 1), (i_t * BT, i_k * BK), (BT, BK), (1, 0))
                    p_w = tl.make_block_ptr(w + (bos * H + i_h) * K, (T, K), (H * K, 1), (i_t * BT, i_k * BK), (BT, BK), (1, 0))
                    b_k = tl.load(p_k, boundary_check=(0, 1))
                    b_kb = b_k * b_beta[:, None]
                    if USE_G:
                        b_kb *= b_g[:, None]
                    if USE_GK:
                        p_gk = tl.make_block_ptr(gk + (bos * H + i_h) * K, (T, K), (H * K, 1), (i_t * BT, i_k * BK), (BT, BK), (1, 0))
                        b_kb *= tl.exp(tl.load(p_gk, boundary_check=(0, 1)))
                    b_w = tl.dot(b_A, b_kb.to(b_k.dtype))
                    tl.store(p_w, b_w.to(p_w.dtype.element_ty), boundary_check=(0, 1))


def recompute_w_u_fwd(
        k: torch.Tensor,
        v: torch.Tensor,
        beta: torch.Tensor,
        A: torch.Tensor,
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
        g: Optional[torch.Tensor] = None,
        gk: Optional[torch.Tensor] = None,
        cu_seqlens: Optional[torch.LongTensor] = None,
) -> Tuple[torch.Tensor, torch.Tensor]:
    B, T, H, K, V = *k.shape, v.shape[-1]
    BT = A.shape[-1]
    BK = 128
    BV = 128

    chunk_indices = prepare_chunk_indices(cu_seqlens, BT) if cu_seqlens is not None else None
    NT = triton.cdiv(T, BT) if cu_seqlens is None else len(chunk_indices)
    g = g.transpose(1, 2).contiguous() if g is not None else None
    beta = beta.transpose(1, 2).contiguous()

    w = torch.empty_like(k)
    u = torch.empty_like(v)
    cv_kernel_num = 24
    recompute_w_u_fwd_kernel[(cv_kernel_num,)](
        k=k,
        v=v,
        beta=beta,
        w=w,
291
292
293
294
295
296
297
298
299
300
301
302
        BT=BT,
        BK=BK,
        BV=BV,
    )
    return w, u


def prepare_wy_repr_bwd(
        k: torch.Tensor,
        v: torch.Tensor,
        g: torch.Tensor,
        beta: torch.Tensor,
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
        du: torch.Tensor,
        cu_seqlens: Optional[torch.LongTensor],
        chunk_size: int = 64,
) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:
    B, T, H, K, V = *k.shape, v.shape[-1]
    BT = chunk_size
    chunk_indices = prepare_chunk_indices(cu_seqlens, BT) if cu_seqlens is not None else None
    NT = triton.cdiv(T, BT) if cu_seqlens is None else len(chunk_indices)
    BK = 128
    BV = 128
    beta = beta.transpose(1, 2).contiguous()
    g = g.transpose(1, 2).contiguous()

    dk = torch.empty_like(k)
    dv = torch.empty_like(v)
    dbeta = torch.empty_like(beta)
    dg = torch.empty_like(g)

    cv_kernel_num = 24
    prepare_wy_repr_bwd_kernel[(cv_kernel_num,)](
        k=k,
        v=v,
        beta=beta,
        g=g,
345
346
347
348
349
350
351
352
353
354
355
356
357
        BK=BK,
        BV=BV,
    )

    dbeta = dbeta.transpose(1, 2).contiguous()
    dg = dg.transpose(1, 2).contiguous()

    return dk, dv, dbeta, dg


bwd_prepare_wy_repr = prepare_wy_repr_bwd

fwd_recompute_w_u = recompute_w_u_fwd
hyper_parallel/components/functional/_triton/kimi_delta_attention/state_summary.py
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
# pylint: disable=invalid-name,missing-public-type-hints

"""Fixed-shape Triton-Ascend kernel for a packed KDA state summary."""

__all__ = [
    "kda_split_state_summary_kernel",
    "kda_state_grad_ext_kernel",
]

import triton
import triton.language as tl


@triton.jit
def _exp2(value):
    """Evaluate base-two exponentiation in FP32."""
    return tl.math.exp2(value.to(tl.float32))


@triton.jit(do_not_specialize=["T"])
def kda_split_state_summary_kernel(
    key,
    w,
    u,
    gate,
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
    OUTPUT_MODE: tl.constexpr = 0,
    W_HEAD_FIRST: tl.constexpr = False,
):
    """Build ``S_ext`` and ``M`` directly without a packed-output copy."""
    i_col = tl.program_id(0)
    i_bh = tl.program_id(1)
    i_b = i_bh // H
    i_h = i_bh % H + HEAD_OFFSET

    stride_k = H_TOTAL * K
    stride_w = K if W_HEAD_FIRST else stride_k
    stride_v = H_TOTAL * V
    key += (i_b * T * H_TOTAL + i_h) * K
    if W_HEAD_FIRST:
        w += (i_b * H_TOTAL + i_h) * T * K
    else:
        w += (i_b * T * H_TOTAL + i_h) * K
    u += (i_b * T * H_TOTAL + i_h) * V
    gate += (i_b * T * H_TOTAL + i_h) * K
    state_ext += i_bh * K * V
    transition += i_bh * K * K

    col = i_col * BV + tl.arange(0, BV)
    row1 = tl.arange(0, 64)
    row2 = 64 + tl.arange(0, 64)
    if OUTPUT_MODE == 1:
        is_transition = False
        transition_col = col
    elif OUTPUT_MODE == 2:
        is_transition = True
        transition_col = col
    else:
        is_transition = i_col * BV >= V
        transition_col = col - V
    state1 = tl.where(
        is_transition & (row1[:, None] == transition_col[None, :]),
        1.0,
        0.0,
    ).to(tl.float32)
    state2 = tl.where(
        is_transition & (row2[:, None] == transition_col[None, :]),
        1.0,
        0.0,
    ).to(tl.float32)
 94
 95
 96
 97
 98
 99
100
101
102
103
104
        1.0,
        0.0,
    ).to(tl.float32)

    num_chunks = tl.cdiv(T, BT)
    for chunk_idx in range(num_chunks):
        w1_ptr = tl.make_block_ptr(
            w,
            (T, K),
            (stride_w, 1),
            (chunk_idx * BT, 0),
104
105
106
107
108
109
110
111
112
            (chunk_idx * BT, 0),
            (BT, 64),
            (1, 0),
        )
        w2_ptr = tl.make_block_ptr(
            w,
            (T, K),
            (stride_w, 1),
            (chunk_idx * BT, 64),
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
            (chunk_idx * BT, 64),
            (BT, 64),
            (1, 0),
        )
        w1 = tl.load(w1_ptr, boundary_check=(0, 1))
        w2 = tl.load(w2_ptr, boundary_check=(0, 1))
        value_new = tl.dot(w1, state1.to(w1.dtype))
        value_new += tl.dot(w2, state2.to(w2.dtype))

        if OUTPUT_MODE == 2:
            value_new = -value_new
        else:
            u_col_offset = tl.where(is_transition, 0, i_col * BV)
            u_ptr = tl.make_block_ptr(
                u,
                (T, V),
                (stride_v, 1),
                (chunk_idx * BT, u_col_offset),
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
                (chunk_idx * BT, u_col_offset),
                (BT, BV),
                (1, 0),
            )
            u_value = tl.load(
                u_ptr,
                boundary_check=(0, 1),
                padding_option="zero",
            )
            value_new = tl.where(is_transition, 0.0, u_value) - value_new
        value_new = value_new.to(key.dtype.element_ty)

        last_idx = min((chunk_idx + 1) * BT, T) - 1
        gate_last_ptr = gate + last_idx * H_TOTAL * K
        gate_last1 = tl.load(gate_last_ptr + row1, mask=row1 < K, other=0.0)
        gate_last2 = tl.load(gate_last_ptr + row2, mask=row2 < K, other=0.0)
        state1 *= _exp2(gate_last1)[:, None]
        state2 *= _exp2(gate_last2)[:, None]

        key1_ptr = tl.make_block_ptr(
            key,
            (K, T),
            (1, stride_k),
            (0, chunk_idx * BT),
152
153
154
155
156
157
158
159
160
            (0, chunk_idx * BT),
            (64, BT),
            (0, 1),
        )
        key2_ptr = tl.make_block_ptr(
            key,
            (K, T),
            (1, stride_k),
            (64, chunk_idx * BT),
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
            (64, chunk_idx * BT),
            (64, BT),
            (0, 1),
        )
        key1 = tl.load(key1_ptr, boundary_check=(0, 1))
        key2 = tl.load(key2_ptr, boundary_check=(0, 1))
        state1 += tl.dot(key1, value_new)
        state2 += tl.dot(key2, value_new)

    state_mask = (~is_transition) & (col[None, :] < V)
    transition_mask = is_transition & (transition_col[None, :] < K)
    tl.store(
        state_ext + row1[:, None] * V + col[None, :],
        state1,
        mask=state_mask,
    )
    tl.store(
        state_ext + row2[:, None] * V + col[None, :],
        state2,
        mask=state_mask,
    )
    tl.store(
        transition + row1[:, None] * K + transition_col[None, :],
        state1,
        mask=transition_mask,
    )
    tl.store(
        transition + row2[:, None] * K + transition_col[None, :],
        state2,
        mask=transition_mask,
    )
189
190
191
192
193
194
195
196
197
198
        mask=transition_mask,
    )


@triton.jit(do_not_specialize=["T"])
def kda_state_grad_ext_kernel(
    query,
    key,
    w,
    gate,
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
    BT: tl.constexpr,
    BV: tl.constexpr,
):
    """Build the local-output contribution to the incoming state gradient."""
    i_col = tl.program_id(0)
    i_bh = tl.program_id(1)
    i_b = i_bh // H
    i_h = i_bh % H + HEAD_OFFSET

    stride_k = H_TOTAL * K
    stride_v = H_TOTAL * V
    query += (i_b * T * H_TOTAL + i_h) * K
    key += (i_b * T * H_TOTAL + i_h) * K
    w += (i_b * T * H_TOTAL + i_h) * K
    gate += (i_b * T * H_TOTAL + i_h) * K
    grad_output += (i_b * T * H_TOTAL + i_h) * V
    grad_value += (i_b * T * H_TOTAL + i_h) * V
    grad_state_ext += i_bh * K * V

    row1 = tl.arange(0, 64)
    row2 = 64 + tl.arange(0, 64)
    state1 = tl.zeros([64, BV], dtype=tl.float32)
    state2 = tl.zeros([64, BV], dtype=tl.float32)

    num_chunks = tl.cdiv(T, BT)
    for reverse_idx in range(num_chunks):
        chunk_idx = num_chunks - 1 - reverse_idx

        key1_ptr = tl.make_block_ptr(
            key,
            (T, K),
            (stride_k, 1),
            (chunk_idx * BT, 0),
241
242
243
244
245
246
247
248
249
            (chunk_idx * BT, 0),
            (BT, 64),
            (1, 0),
        )
        key2_ptr = tl.make_block_ptr(
            key,
            (T, K),
            (stride_k, 1),
            (chunk_idx * BT, 64),
249
250
251
252
253
254
255
256
257
258
259
260
261
262
            (chunk_idx * BT, 64),
            (BT, 64),
            (1, 0),
        )
        key1 = tl.load(key1_ptr, boundary_check=(0, 1))
        key2 = tl.load(key2_ptr, boundary_check=(0, 1))
        value_grad = tl.dot(key1, state1.to(key1.dtype))
        value_grad += tl.dot(key2, state2.to(key2.dtype))

        grad_value_ptr = tl.make_block_ptr(
            grad_value,
            (T, V),
            (stride_v, 1),
            (chunk_idx * BT, i_col * BV),
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
            (chunk_idx * BT, i_col * BV),
            (BT, BV),
            (1, 0),
        )
        value_grad += tl.load(grad_value_ptr, boundary_check=(0, 1))

        last_idx = min((chunk_idx + 1) * BT, T) - 1
        gate_last_ptr = gate + last_idx * H_TOTAL * K
        gate_last1 = tl.load(
            gate_last_ptr + row1,
            mask=row1 < K,
            other=0.0,
        )
        gate_last2 = tl.load(
            gate_last_ptr + row2,
            mask=row2 < K,
            other=0.0,
        )
        state1 *= _exp2(gate_last1)[:, None]
        state2 *= _exp2(gate_last2)[:, None]

        output_grad_ptr = tl.make_block_ptr(
            grad_output,
            (T, V),
            (stride_v, 1),
            (chunk_idx * BT, i_col * BV),
287
288
289
290
291
292
293
294
295
296
297
            (chunk_idx * BT, i_col * BV),
            (BT, BV),
            (1, 0),
        )
        output_grad = tl.load(output_grad_ptr, boundary_check=(0, 1))

        query1_ptr = tl.make_block_ptr(
            query,
            (K, T),
            (1, stride_k),
            (0, chunk_idx * BT),
297
298
299
300
301
302
303
304
305
            (0, chunk_idx * BT),
            (64, BT),
            (0, 1),
        )
        query2_ptr = tl.make_block_ptr(
            query,
            (K, T),
            (1, stride_k),
            (64, chunk_idx * BT),
305
306
307
308
309
310
311
312
313
            (64, chunk_idx * BT),
            (64, BT),
            (0, 1),
        )
        w1_ptr = tl.make_block_ptr(
            w,
            (K, T),
            (1, stride_k),
            (0, chunk_idx * BT),
313
314
315
316
317
318
319
320
321
            (0, chunk_idx * BT),
            (64, BT),
            (0, 1),
        )
        w2_ptr = tl.make_block_ptr(
            w,
            (K, T),
            (1, stride_k),
            (64, chunk_idx * BT),
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
            (64, chunk_idx * BT),
            (64, BT),
            (0, 1),
        )
        query1 = tl.load(query1_ptr, boundary_check=(0, 1))
        query2 = tl.load(query2_ptr, boundary_check=(0, 1))
        w1 = tl.load(w1_ptr, boundary_check=(0, 1))
        w2 = tl.load(w2_ptr, boundary_check=(0, 1))
        state1 += tl.dot(query1, output_grad.to(query1.dtype)) * scale
        state1 -= tl.dot(w1, value_grad.to(w1.dtype))
        state2 += tl.dot(query2, output_grad.to(query2.dtype)) * scale
        state2 -= tl.dot(w2, value_grad.to(w2.dtype))

    output1_ptr = tl.make_block_ptr(
        grad_state_ext,
        (K, V),
        (V, 1),
        (0, i_col * BV),
338
339
340
341
342
343
344
345
346
        (0, i_col * BV),
        (64, BV),
        (1, 0),
    )
    output2_ptr = tl.make_block_ptr(
        grad_state_ext,
        (K, V),
        (V, 1),
        (64, i_col * BV),
346
347
348
349
350
351
352
353
354
355
356
357
358
359
        (64, i_col * BV),
        (64, BV),
        (1, 0),
    )
    tl.store(
        output1_ptr,
        state1.to(output1_ptr.dtype.element_ty),
        boundary_check=(0, 1),
    )
    tl.store(
        output2_ptr,
        state2.to(output2_ptr.dtype.element_ty),
        boundary_check=(0, 1),
    )
hyper_parallel/components/functional/gated_delta_net.py
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
from typing import Optional

import torch

from ._triton.gated_delta_net.chunk_delta_h import chunk_gated_delta_rule_bwd_dhu, chunk_gated_delta_rule_fwd_h
from ._triton.gated_delta_net.chunk_o import chunk_bwd_dqkwg, chunk_bwd_dv_local, chunk_fwd_o
from ._triton.gated_delta_net.chunk_scaled_dot_kkt import chunk_scaled_dot_kkt_fwd
from ._triton.gated_delta_net.cumsum import chunk_local_cumsum
from ._triton.gated_delta_net.solve_tril import solve_tril
from ._triton.gated_delta_net.utils import autocast_custom_bwd, autocast_custom_fwd, input_guard
from ._triton.gated_delta_net.wy_fast import prepare_wy_repr_bwd, recompute_w_u_fwd


def _l2norm(x: torch.Tensor, eps: float = 1e-6) -> tuple[torch.Tensor, torch.Tensor]:
    inv_norm = torch.rsqrt((x * x).sum(dim=-1, keepdim=True) + eps)
49
50
51
52
53
54
55
56
57
        chunk_size: int = 64,
):
    """Compute forward intermediates that do not depend on the initial state."""
    g = chunk_local_cumsum(g, chunk_size=chunk_size, cu_seqlens=cu_seqlens, head_first=False)
    matrix_a = chunk_scaled_dot_kkt_fwd(
        k=k,
        g=g,
        beta=beta,
        cu_seqlens=cu_seqlens,
57
58
59
60
61
62
63
64
65
        cu_seqlens=cu_seqlens,
        chunk_size=chunk_size,
        output_dtype=torch.float32,
    )
    matrix_a = solve_tril(A=matrix_a, cu_seqlens=cu_seqlens, output_dtype=k.dtype)
    w, u = recompute_w_u_fwd(
        k=k,
        v=v,
        beta=beta,
66
67
68
69
70
71
72
73
74
        A=matrix_a,
        g=g,
        cu_seqlens=cu_seqlens,
    )
    return g, matrix_a, w, u


def chunk_gated_delta_rule_fwd_apply_state(
        k: torch.Tensor,
127
128
129
130
131
132
133
134
135
        output_final_state: bool,
        cu_seqlens: Optional[torch.LongTensor] = None,
        chunk_size: int = 64,
):
    g, matrix_a, w, u = chunk_gated_delta_rule_fwd_prepare(
        k=k,
        v=v,
        g=g,
        beta=beta,
145
146
147
148
149
150
151
152
153
        output_final_state=output_final_state,
        cu_seqlens=cu_seqlens,
        chunk_size=chunk_size,
    )
    output = chunk_gated_delta_rule_fwd_output(
        q=q,
        k=k,
        v_new=v_new,
        h=h,
155
156
157
158
159
160
161
162
163
        scale=scale,
        cu_seqlens=cu_seqlens,
        chunk_size=chunk_size,
    )
    return g, output, matrix_a, final_state


def chunk_gated_delta_rule_bwd_prepare(
        q: torch.Tensor,
370
371
372
373
374
375
376
377
378
    if use_qk_l2norm_in_kernel:
        q_norm, q_inv_norm = _l2norm(q)
        k_norm, k_inv_norm = _l2norm(k)

    g_cumsum, matrix_a, w, u = chunk_gated_delta_rule_fwd_prepare(
        k=k_norm,
        v=v,
        g=g,
        beta=beta,
703
704
705
706
707
708
709
710
711
            cu_seqlens: Optional[torch.LongTensor] = None,
            use_qk_l2norm_in_kernel: bool = False,
            chunk_size: int = 64,
    ):
        g, output, matrix_a, final_state = chunk_gated_delta_rule_fwd(
            q=q,
            k=k,
            v=v,
            g=g,
718
719
720
721
722
723
724
725
726
727
728
729
730
731
        )

        saved_initial_state = initial_state if initial_state is not None else q.new_empty(0)
        saved_cu_seqlens = cu_seqlens if cu_seqlens is not None else q.new_empty(0, dtype=torch.long)
        ctx.save_for_backward(q, k, v, g, beta, matrix_a, saved_initial_state, saved_cu_seqlens)
        ctx.has_initial_state = initial_state is not None
        ctx.has_cu_seqlens = cu_seqlens is not None
        ctx.scale = scale
        ctx.chunk_size = chunk_size
        return output.to(q.dtype), final_state

    @staticmethod
    @input_guard
    @autocast_custom_bwd
733
734
735
736
737
738
739
740
741
            ctx,
            do: torch.Tensor,
            dht: torch.Tensor
    ):
        q, k, v, g, beta, matrix_a, initial_state, cu_seqlens = ctx.saved_tensors
        if not ctx.has_initial_state:
            initial_state = None
        if not ctx.has_cu_seqlens:
            cu_seqlens = None
878
879
880
881
882
883
884
885
886
    if use_qk_l2norm_in_kernel:
        q, _ = _l2norm(q)
        k, _ = _l2norm(k)

    output, final_state = ChunkGatedDeltaRuleFunction.apply(
        q,
        k,
        v,
        g,
891
892
893
894
895
        cu_seqlens,
        False,
        chunk_size,
    )
    return output, final_state
hyper_parallel/components/functional/gated_delta_net_state_summary.py
15
16
17
18
19
20
21
22
23
"""Affine state-summary operations used by GDN State-P2P."""

# pylint: disable=forbidden-backend-import

__all__ = [
    "apply_gdn_state_gradient_summary",
    "apply_gdn_state_summary",
    "chunk_gated_delta_rule_state_gradient_summary_bwd",
    "chunk_gated_delta_rule_state_summary_fwd",
27
28
29
30
31
32
33
34

import torch
import triton

from ._triton.gated_delta_net.state_summary import (
    gdn_packed_state_summary_kernel,
    gdn_state_grad_ext_kernel,
)
56
57
58
59
60
61
62
63
64
    chunk_size: int = 64,
    block_size: int = 128,
) -> tuple[torch.Tensor, torch.Tensor]:
    """Return the local affine map ``state_out = M @ state_in + S``."""
    if (key.ndim, w.ndim, u.ndim, g.ndim) != (4, 4, 4, 3):
        raise ValueError("GDN state summary expects key/w/u [B,T,H,D] and g [B,T,H].")
    batch, seq_len, heads, key_dim = key.shape
    value_dim = u.shape[-1]
    _validate_fixed_summary_shape(key_dim, value_dim, chunk_size)
104
105
106
107
108
109
110
111
112
    transition = packed_summary[..., value_dim:].contiguous()
    return state_ext, transition


def _validate_gdn_gradient_summary_inputs(
    query: torch.Tensor,
    key: torch.Tensor,
    w: torch.Tensor,
    g: torch.Tensor,
123
124
125
126
127
128
129
130
131
            f"GDN state-gradient sequence length {seq_len} must be divisible by {chunk_size}."
        )
    qk_shape = (batch, seq_len, heads, key_dim)
    value_shape = (batch, seq_len, heads, value_dim)
    if (key.shape, w.shape, g.shape, grad_output.shape, dv.shape) != (
        qk_shape, qk_shape, qk_shape[:3], value_shape, value_shape
    ):
        raise ValueError(
            "Incompatible GDN state-gradient summary shapes: "
132
133
134
135
136
137
138
139
140
141
142
143
144
            f"query={tuple(query.shape)}, key={tuple(key.shape)}, "
            f"w={tuple(w.shape)}, g={tuple(g.shape)}, "
            f"grad_output={tuple(grad_output.shape)}, dv={tuple(dv.shape)}."
        )
    return batch, seq_len, heads, key_dim, value_dim


@torch.compiler.disable
def chunk_gated_delta_rule_state_gradient_summary_bwd(
    query: torch.Tensor,
    key: torch.Tensor,
    w: torch.Tensor,
    g: torch.Tensor,
148
149
150
151
152
153
154
155
156
    *,
    chunk_size: int = 64,
) -> torch.Tensor:
    """Return the local-loss contribution to the incoming state gradient."""
    batch, seq_len, heads, key_dim, value_dim = _validate_gdn_gradient_summary_inputs(
        query, key, w, g, grad_output, dv, chunk_size
    )

    query, key, w, g, grad_output, dv = (
hyper_parallel/components/functional/kimi_delta_attention.py
16
17
18
19
20
21
22
23
24
from __future__ import annotations

# pylint: disable=forbidden-backend-import

__all__ = ["fused_chunk_kda", "fused_chunk_kda_p2p"]

from typing import Any, Optional

import torch
32
33
34
35
36
37
38
39
40
    kda_state_summary_forward_from_prepared,
)


def _validate_local_shapes(
    query: torch.Tensor,
    key: torch.Tensor,
    value: torch.Tensor,
    gate: torch.Tensor,
40
41
42
43
44
45
46
47
48
    gate: torch.Tensor,
    beta: torch.Tensor,
) -> tuple[int, int, int, int]:
    """Validate local KDA tensor shapes and return dimensions used by later checks."""
    if (query.ndim, key.ndim, value.ndim, gate.ndim) != (4, 4, 4, 4):
        raise ValueError("Fused KDA expects rank-4 query/key/value/gate tensors.")
    if beta.ndim != 3:
        raise ValueError("Fused KDA expects a rank-3 beta tensor.")
    if query.shape != key.shape:
56
57
58
59
60
61
62
63
64
65
66
67
    if beta.shape != (batch, sequence_length, num_value_heads):
        raise ValueError("Fused KDA beta has an incompatible shape.")
    if num_value_heads % num_query_heads:
        raise ValueError("Fused KDA value heads must be divisible by query heads.")
    return sequence_length, key_dim, value_dim, num_value_heads


def _validate_local_inputs(
    query: torch.Tensor,
    key: torch.Tensor,
    value: torch.Tensor,
    gate: torch.Tensor,
72
73
74
75
76
77
78
79
80
    chunk_size: int,
    lower_bound: float,
) -> None:
    """Validate the fixed-shape local KDA backend before compiling kernels."""
    sequence_length, key_dim, value_dim, num_value_heads = _validate_local_shapes(
        query, key, value, gate, beta
    )
    if not (
        query.dtype == key.dtype == value.dtype == gate.dtype == beta.dtype
457
458
459
460
461
462
463
464
465
            chunk_size=ctx.chunk_size,
        )
        # These recomputed tensors have no consumers after dhu. Dropping local
        # references lets the allocator reuse their storage on this stream.
        del query_gated, key_gated, w, u, grad_value_local, _
        (
            grad_query,
            grad_key,
            grad_value,
481
482
483
484
485
486
487
488
489
            scale=ctx.scale,
            chunk_size=ctx.chunk_size,
        )
        # Release recomputed states before intra backward allocates its outputs.
        del states, grad_states, value_new
        grad_query, grad_key, grad_beta, grad_gate = ops.chunk_kda_bwd_intra(
            q=query,
            k=key,
            g=gate,
hyper_parallel/components/functional/kimi_delta_attention_fla_adapter.py
16
17
18
19
20
21
22
23
24
from __future__ import annotations

# pylint: disable=forbidden-backend-import

__all__ = [
    "FLAKDAStagedOps",
    "get_fla_kda_staged_ops",
    "is_fla_triton_kda_available",
    "run_fla_chunk_kda",
hyper_parallel/components/functional/kimi_delta_attention_state_summary.py
15
16
17
18
19
20
21
22
23
"""Affine KDA state summary built from prepared WY intermediates."""

# pylint: disable=forbidden-backend-import

__all__ = [
    "apply_kda_state_gradient_summary",
    "apply_kda_state_summary",
    "kda_state_gradient_summary_from_prepared",
    "kda_state_summary_forward_from_prepared",
36
37
38
39
40
41
42
43
44
    gate: torch.Tensor,
    chunk_size: int,
) -> tuple[int, int, int, int, int]:
    """Validate the fixed Kimi K3 summary contract and return dimensions."""
    if (key.ndim, w.ndim, u.ndim, gate.ndim) != (4, 4, 4, 4):
        raise ValueError("KDA prepared summary expects rank-4 key/w/u/gate tensors.")
    batch, sequence_length, heads, key_dim = key.shape
    value_dim = u.shape[-1]
    if key_dim != 128 or value_dim != 128 or chunk_size != 64:
149
150
151
152
153
154
155
156
157
    chunk_size: int,
    block_size: int = 64,
) -> torch.Tensor:
    """Launch the reverse-wavefront Triton-Ascend summary kernel."""
    from ._triton.kimi_delta_attention.state_summary import (  # pylint: disable=import-outside-toplevel
        kda_state_grad_ext_kernel,
    )

    batch, sequence_length, heads, key_dim, value_dim = (
210
211
212
213
214
215
216
217
218
    *,
    chunk_size: int = 64,
) -> tuple[torch.Tensor, torch.Tensor]:
    """Build ``S_ext`` and ``M`` in compile-time-separated BV=128 modes."""
    from ._triton.kimi_delta_attention.state_summary import (  # pylint: disable=import-outside-toplevel
        kda_split_state_summary_kernel,
    )

    batch, sequence_length, heads, key_dim, value_dim = (
hyper_parallel/components/modules/gated_delta_net.py
139
140
141
142
143
144
145
146
147
148
149
150
151


def _parse_version(version_text: str) -> tuple[int, int, int]:
    """Return a three-component numeric version tuple."""
    parts = []
    for part in version_text.split("+")[0].split(".")[:3]:
        match = re.match(r"\d+", part)
        if match is not None:
            parts.append(int(match.group()))
    return tuple((parts + [0, 0, 0])[:3])


def _is_triton_gdn_input_supported(
155
156
157
158
159
160
161
162
163
    g: Optional[torch.Tensor],
    beta: Optional[torch.Tensor],
) -> bool:
    """Check the fixed Qwen3.5 GDN contract validated by this backend."""
    if any(tensor is None for tensor in (key, value, g, beta)):
        return False
    if not (
        query.device.type == "npu"
        and query.dtype == key.dtype == value.dtype == beta.dtype == torch.bfloat16
hyper_parallel/components/modules/kimi_delta_attention.py
57
58
59
60
61
62
63
64
65
66
67
    dt_bias: torch.Tensor,
    chunk_size: int,
) -> bool:
    """Check the fixed dense shapes accepted by the Triton KDA path."""
    batch_size, sequence_length, num_query_heads, key_dim = query.shape
    num_value_heads, value_dim = value.shape[2:]
    return all((
        value.shape[:2] == (batch_size, sequence_length),
        gate.shape == (batch_size, sequence_length, num_value_heads, key_dim),
        beta.shape == (batch_size, sequence_length, num_value_heads),
        num_value_heads % num_query_heads == 0,
102
103
104
105
106
107
108
109
110
111
112
113
    )
    if not all(basic_contract):
        return False

    shape_supported = _has_triton_kda_shapes(
        query, value, gate, beta, a_log, dt_bias, chunk_size
    )
    lower_bound_supported = -5.0 <= lower_bound < 0
    if not (shape_supported and lower_bound_supported):
        return False
    return all(tensor.device == query.device for tensor in operands)

hyper_parallel/distributed/context_parallel/kimi_delta_attention.py
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
    convolution: nn.Conv1d,
    halo: Optional[torch.Tensor] = None,
) -> torch.Tensor:
    """Apply one local ShortConv, optionally with a preceding-rank halo."""
    if halo is None:
        conv_input = tensor.transpose(1, 2)
        padding = convolution.padding
    else:
        conv_input = torch.cat((halo, tensor), dim=1).transpose(1, 2)
        padding = 0
    output = F.conv1d(  # pylint: disable=not-callable
        input=conv_input,
        weight=convolution.weight,
        bias=convolution.bias,
        stride=convolution.stride,
228
229
230
231
232
233
234
235
236
237
238
        padding=padding,
        dilation=convolution.dilation,
        groups=convolution.groups,
    )
    if halo is None:
        output = output[:, :, : tensor.shape[1]]
    return F.silu(output).transpose(1, 2)


def _causal_short_convs_with_cp_halo(
    projected: tuple[torch.Tensor, torch.Tensor, torch.Tensor],
248
249
250
251
252
253
254
255
256
    if len(halo_widths) != 1:
        raise ValueError("KDA P2P requires Q/K/V ShortConv halo widths to match.")
    halo_width = halo_widths.pop()
    if halo_width == 0 or cp_size == 1:
        return tuple(
            _run_causal_short_conv(tensor, convolution)
            for tensor, convolution in zip(projected, convolutions)
        )
    if projected[0].shape[1] < halo_width:
270
271
272
273
274
275
276
277
        cp_rank,
        cp_size,
    )
    halos = torch.split(packed_halo, channel_sizes, dim=-1)
    return tuple(
        _run_causal_short_conv(tensor, convolution, halo)
        for tensor, halo, convolution in zip(projected, halos, convolutions)
    )
566
567
568
569
570
571
572
573
574
        a_log: torch.Tensor,
        dt_bias: torch.Tensor,
    ) -> None:
        """Validate the projected-tensor boundary and Ulysses divisibility."""
        if (query.dim(), key.dim(), value.dim(), gate.dim()) != (4, 4, 4, 4):
            raise ValueError("query, key, value, and gate must be rank-4 tensors.")
        if beta.dim() != 3:
            raise ValueError("beta must be a rank-3 tensor.")
        if query.shape != key.shape: