Diff Coverage

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

Source File Diff Coverage (%) Missing Lines
hyper_parallel/components/functional/_gdn_triton/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-277,279-280,282,284-289,309-311,314,321,325-326,353-360,362-364,366-372,374-382,384-387,389-391,393-404,406-417,419-422,424-429,431-433,435,438-443,445-451,453-459,461-467,469-472,474,476-518,520-531,534,550,552-553,555-558,560,562-564,566,568,570,592
hyper_parallel/components/functional/_gdn_triton/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/_gdn_triton/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/_gdn_triton/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/_gdn_triton/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/_gdn_triton/state_summary.py 0.0% 24-25,27,30,39-40,55-58,60-66,68-73,76,80-81,84,87-90,92,100,102-111,113,116,119-120,122,130,138-139,142,151-152,170-173,175-183,185-186,188-194,196,199,202-206,208,216,218,226-229,231,234,237,240,243-253,255,258,261-262,265
hyper_parallel/components/functional/_gdn_triton/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-135,137,140,142-143,145,147,153-154,161,164-170,173-175,178-183,187-188,197,200,203-205,208-211,213-214,216-218,220-221,224,229-232,234-243,245-246,248,250-251,253,256-257,260-262,265-269,271-274,277-279,283-285,288-292,294-299,302-309,312,322-323,334,341-342,344,355-356,364-365,368-369,372-374,376-378,381-384,386-388
hyper_parallel/components/functional/_gdn_triton/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/_kda_triton/state_summary.py 0.0% 22-23,26-27,29,32-33,52-55,57-62,64-68,70-78,80-82,87,93-95,103,111-114,116-117,119-120,128,133-134,136-141,143,151,159-162,164-166,171,176,181,188-189,208-211,213-221,223-226,228-230,232,240,248-251,253,261,263-265,270,275-276,278,286,288,296,304,312,320-327,329,337,345,350,357
hyper_parallel/components/functional/gated_delta_net.py 0.0% 24-25,27,29-35,38-40,43,52-53,61-62,70,73,84,96,107,119,131,139,149,159,162,176,184,194,203,206,220,235,253,267,278-282,289,292,306,319,332,349,352-354,365-366,368-372,374,381,394-396,406,417-419,429,440-442,454,474,483,492,505-507,520,534-536,549,563-565,587,603-609,614-615,618-620,639,651,663,684,687,690-693,707,720-727,729-732,737-742,756,759-760,838-839,843-844,847-848,852-853,857-858,864-866,870-871,875-876,878-880,882,895
hyper_parallel/components/functional/gated_delta_net_state_summary.py 0.0% 17,19-20,22,28,33-34,40-41,51-59,62-63,69-70,78,94-96,99-100,112-116,119-121,128,135,138,146,163,166,172-174,177,183-185,191
hyper_parallel/components/functional/kimi_delta_attention.py 0.0% 16,18,20-21,23-24,32,45-61,65-69,72-73,77-85,89-90,93,108,119-120,123,138,149,162,165,168-169,189,191,193,197-200,209,215-218,223,229,240-243,250-251,253-256,261-262,268,277,286-287,289-292,307-316,318-319,324,326,328,343-344,346-349,359,365-366,374,382,391,400-402,412-417,422-423,429,441,463,478-482,488,495,500,507-511,513,533-534,554,567-568,588
hyper_parallel/components/functional/kimi_delta_attention_fla_adapter.py 58.2% 63-67,70-71,75,79-80,84-87,98,104-108,111-112,119,133-136,150,158-159,166,173-175,216,225-229,234
hyper_parallel/components/functional/kimi_delta_attention_state_summary.py 0.0% 17,19-20,23,31-36,40-41,45-46,50-51,55-57,60,70-72,75-78,82-83,87-90,95-96,100-101,106-107,110,113,122,131,144,148,159,163,171,192,195-196,205,209,212,215,223,231,249,267,270-271,280,289-290,307,320,326-328,331,337-339,345
hyper_parallel/components/modules/gated_delta_net.py 0.0% 247
hyper_parallel/components/modules/kimi_delta_attention.py 0.0% 114,559
hyper_parallel/core/context_parallel/async_context_parallel.py 100%  
hyper_parallel/core/context_parallel/async_dsa_context_parallel.py 100%  
hyper_parallel/core/context_parallel/context_parallel.py 50.0% 876,895-896
hyper_parallel/core/context_parallel/dsa_context_parallel.py 80.0% 96,706
hyper_parallel/core/context_parallel/utils.py 31.0% 35,40-44,49-56,61-66,75-82,87-93,100-103,112-118,123-126,130-137,146-150,156-158,163-165,171-173,198-200,205-207,212-214,221,228,235,238,251-259,264-266,271
hyper_parallel/distributed/context_parallel/gated_delta_net.py 0.0% 594,599,700,705
hyper_parallel/distributed/context_parallel/kimi_delta_attention.py 0.0% 751
hyper_parallel/components/functional/_gdn_triton/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=line-too-long,missing-public-type-hints,missing-public-docstring
# pylint: disable=used-before-assignment,unsupported-binary-operation,unused-argument
# pylint: disable=invalid-name,missing-module-docstring,missing-function-docstring

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
        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)
    assert K <= 256, "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,
305
306
307
308
309
310
311
312
313
314
315
316
317
318
        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,
317
318
319
320
321
322
323
324
325
326
327
328
329
330
    '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,
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
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
        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,
546
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
    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
    assert K <= 256, "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,
588
589
590
591
592
        V=V,
        BT=BT,
        BV=BV,
    )
    return dh, dh0, dv2
hyper_parallel/components/functional/_gdn_triton/chunk_o.py
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
# pylint: disable=missing-public-type-hints,missing-public-docstring,disallowed-name
# pylint: disable=invalid-name,missing-module-docstring,missing-function-docstring
# pylint: disable=unused-variable,too-many-nested-blocks

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

    o = 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 o

bwd_chunk_dqkwg = chunk_bwd_dqkwg
bwd_chunk_dv_local = chunk_bwd_dv_local
hyper_parallel/components/functional/_gdn_triton/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=line-too-long,missing-public-type-hints,missing-public-docstring
# pylint: disable=unused-argument,invalid-name,missing-module-docstring
# pylint: disable=missing-function-docstring

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/_gdn_triton/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=missing-public-type-hints,missing-public-docstring,disallowed-name
# pylint: disable=useless-return,unused-argument,no-else-return,invalid-name
# pylint: disable=missing-module-docstring,missing-function-docstring

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/_gdn_triton/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=missing-public-type-hints,missing-public-docstring,invalid-name
# pylint: disable=import-outside-toplevel,unused-argument,unused-import
# pylint: disable=missing-module-docstring,missing-function-docstring

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/_gdn_triton/state_summary.py
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
# pylint: disable=missing-public-type-hints,invalid-name

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

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,),
35
36
37
38
39
40
41
42
43
44
        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,
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
    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),
 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
            (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),
126
127
128
129
130
131
132
133
134
        (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),
134
135
136
137
138
139
140
141
142
143
144
145
146
        (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,),
147
148
149
150
151
152
153
154
155
156
        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,
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
    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),
212
213
214
215
216
217
218
219
220
221
222
            (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),
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
            (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))


__all__ = ["gdn_packed_state_summary_kernel", "gdn_state_grad_ext_kernel"]
hyper_parallel/components/functional/_gdn_triton/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=unused-argument,no-member,import-outside-toplevel
# pylint: disable=invalid-name,missing-module-docstring,missing-function-docstring
# pylint: disable=missing-class-docstring,broad-exception-caught,protected-access

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
    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
    if warning or (FLA_CI_ENV and (error_rate < 0.01 or abs_atol <= 0.3)):
        if error_rate > ratio:
            warnings.warn(msg)
    else:
        assert error_rate < ratio, 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.
    """
149
150
151
152
153
154
155
156
157
158
    Returns None to indicate TMA descriptors are unavailable.
    Just make triton compiler happy.
    """

    @triton.jit
    def make_tensor_descriptor(
        base,
        shape,
        strides,
        block_shape,
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
        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')
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
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
        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:
    assert device == 'cuda', '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",),
318
319
320
321
322
323
324
325
326
327
    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,
330
331
332
333
334
335
336
337
338
        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,
337
338
339
340
341
342
343
344
345
346
347
348
            '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,
351
352
353
354
355
356
357
358
359
360
                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,
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
                        '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/_gdn_triton/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

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/_kda_triton/state_summary.py
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
# pylint: disable=invalid-name,missing-public-type-hints

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

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,
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
    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)
89
90
91
92
93
94
95
96
97
98
99
        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),
 99
100
101
102
103
104
105
106
107
            (chunk_idx * BT, 0),
            (BT, 64),
            (1, 0),
        )
        w2_ptr = tl.make_block_ptr(
            w,
            (T, K),
            (stride_w, 1),
            (chunk_idx * BT, 64),
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
            (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),
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
                (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),
147
148
149
150
151
152
153
154
155
            (0, chunk_idx * BT),
            (64, BT),
            (0, 1),
        )
        key2_ptr = tl.make_block_ptr(
            key,
            (K, T),
            (1, stride_k),
            (64, chunk_idx * BT),
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
            (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,
    )
184
185
186
187
188
189
190
191
192
193
        mask=transition_mask,
    )


@triton.jit(do_not_specialize=["T"])
def kda_state_grad_ext_kernel(
    query,
    key,
    w,
    gate,
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
    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),
236
237
238
239
240
241
242
243
244
            (chunk_idx * BT, 0),
            (BT, 64),
            (1, 0),
        )
        key2_ptr = tl.make_block_ptr(
            key,
            (T, K),
            (stride_k, 1),
            (chunk_idx * BT, 64),
244
245
246
247
248
249
250
251
252
253
254
255
256
257
            (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),
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
            (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),
282
283
284
285
286
287
288
289
290
291
292
            (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),
292
293
294
295
296
297
298
299
300
            (0, chunk_idx * BT),
            (64, BT),
            (0, 1),
        )
        query2_ptr = tl.make_block_ptr(
            query,
            (K, T),
            (1, stride_k),
            (64, chunk_idx * BT),
300
301
302
303
304
305
306
307
308
            (64, chunk_idx * BT),
            (64, BT),
            (0, 1),
        )
        w1_ptr = tl.make_block_ptr(
            w,
            (K, T),
            (1, stride_k),
            (0, chunk_idx * BT),
308
309
310
311
312
313
314
315
316
            (0, chunk_idx * BT),
            (64, BT),
            (0, 1),
        )
        w2_ptr = tl.make_block_ptr(
            w,
            (K, T),
            (1, stride_k),
            (64, chunk_idx * BT),
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
            (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),
333
334
335
336
337
338
339
340
341
        (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),
341
342
343
344
345
346
347
348
349
350
351
352
353
354
        (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),
    )
353
354
355
356
357
358
359
360
        boundary_check=(0, 1),
    )


__all__ = [
    "kda_split_state_summary_kernel",
    "kda_state_grad_ext_kernel",
]
hyper_parallel/components/functional/gated_delta_net.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
# pylint: disable=non-google-docstring,disallowed-name,unused-argument,invalid-name
# pylint: disable=missing-module-docstring,missing-function-docstring
# pylint: disable=abstract-method,arguments-differ

import warnings
from typing import Optional

import torch

from ._gdn_triton.chunk_delta_h import chunk_gated_delta_rule_bwd_dhu, chunk_gated_delta_rule_fwd_h
from ._gdn_triton.chunk_o import chunk_bwd_dqkwg, chunk_bwd_dv_local, chunk_fwd_o
from ._gdn_triton.chunk_scaled_dot_kkt import chunk_scaled_dot_kkt_fwd
from ._gdn_triton.cumsum import chunk_local_cumsum
from ._gdn_triton.solve_tril import solve_tril
from ._gdn_triton.utils import autocast_custom_bwd, autocast_custom_fwd, input_guard
from ._gdn_triton.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)
    return (x * inv_norm).to(x.dtype), inv_norm


def chunk_gated_delta_rule_fwd_prepare(
        k: torch.Tensor,
        v: torch.Tensor,
        g: torch.Tensor,
        beta: torch.Tensor,
48
49
50
51
52
53
54
55
56
57
        cu_seqlens: Optional[torch.LongTensor] = None,
        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)
    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
66
        cu_seqlens=cu_seqlens,
        chunk_size=chunk_size,
        output_dtype=torch.float32,
    )
    A = solve_tril(A=A, cu_seqlens=cu_seqlens, output_dtype=k.dtype)
    w, u = recompute_w_u_fwd(
        k=k,
        v=v,
        beta=beta,
        A=A,
66
67
68
69
70
71
72
73
74
75
76
77
        A=A,
        g=g,
        cu_seqlens=cu_seqlens,
    )
    return g, A, w, u


def chunk_gated_delta_rule_fwd_apply_state(
        k: torch.Tensor,
        g: torch.Tensor,
        w: torch.Tensor,
        u: torch.Tensor,
80
81
82
83
84
85
86
87
88
        cu_seqlens: Optional[torch.LongTensor] = None,
        chunk_size: int = 64,
):
    """Apply the recurrent initial state and return local state intermediates."""
    return chunk_gated_delta_rule_fwd_h(
        k=k,
        w=w,
        u=u,
        g=g,
 92
 93
 94
 95
 96
 97
 98
 99
100
        cu_seqlens=cu_seqlens,
    )


def chunk_gated_delta_rule_fwd_output(
        q: torch.Tensor,
        k: torch.Tensor,
        v_new: torch.Tensor,
        h: torch.Tensor,
103
104
105
106
107
108
109
110
111
        cu_seqlens: Optional[torch.LongTensor] = None,
        chunk_size: int = 64,
):
    """Compute local outputs after recurrent states have been applied."""
    return chunk_fwd_o(
        q=q,
        k=k,
        v=v_new,
        h=h,
115
116
117
118
119
120
121
122
123
        chunk_size=chunk_size,
    )


def chunk_gated_delta_rule_fwd(
        q: torch.Tensor,
        k: torch.Tensor,
        v: torch.Tensor,
        g: 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, A, w, u = chunk_gated_delta_rule_fwd_prepare(
        k=k,
        v=v,
        g=g,
        beta=beta,
135
136
137
138
139
140
141
142
143
        beta=beta,
        cu_seqlens=cu_seqlens,
        chunk_size=chunk_size,
    )
    h, v_new, final_state = chunk_gated_delta_rule_fwd_apply_state(
        k=k,
        g=g,
        w=w,
        u=u,
145
146
147
148
149
150
151
152
153
        output_final_state=output_final_state,
        cu_seqlens=cu_seqlens,
        chunk_size=chunk_size,
    )
    o = 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
164
165
166
        scale=scale,
        cu_seqlens=cu_seqlens,
        chunk_size=chunk_size,
    )
    return g, o, A, final_state


def chunk_gated_delta_rule_bwd_prepare(
        q: torch.Tensor,
        k: torch.Tensor,
        v: torch.Tensor,
        g: torch.Tensor,
172
173
174
175
176
177
178
179
180
        cu_seqlens: Optional[torch.LongTensor] = None,
        chunk_size: int = 64,
):
    """Compute backward intermediates that do not depend on final-state grad."""
    w, u = recompute_w_u_fwd(
        k=k,
        v=v,
        beta=beta,
        A=A,
180
181
182
183
184
185
186
187
188
        A=A,
        g=g,
        cu_seqlens=cu_seqlens,
    )
    h, v_new, _ = chunk_gated_delta_rule_fwd_apply_state(
        k=k,
        g=g,
        w=w,
        u=u,
190
191
192
193
194
195
196
197
198
        output_final_state=False,
        cu_seqlens=cu_seqlens,
        chunk_size=chunk_size,
    )
    dv = chunk_bwd_dv_local(
        q=q,
        k=k,
        g=g,
        do=do,
199
200
201
202
203
204
205
206
207
208
209
210
        scale=scale,
        cu_seqlens=cu_seqlens,
        chunk_size=chunk_size,
    )
    return w, h, v_new, dv


def chunk_gated_delta_rule_bwd_state(
        q: torch.Tensor,
        k: torch.Tensor,
        w: torch.Tensor,
        g: torch.Tensor,
216
217
218
219
220
221
222
223
224
        cu_seqlens: Optional[torch.LongTensor] = None,
        chunk_size: int = 64,
):
    """Apply the final-state gradient and produce the initial-state gradient."""
    return chunk_gated_delta_rule_bwd_dhu(
        q=q,
        k=k,
        w=w,
        g=g,
231
232
233
234
235
236
237
238
239
        chunk_size=chunk_size,
    )


def chunk_gated_delta_rule_bwd_finish(
        q: torch.Tensor,
        k: torch.Tensor,
        v: torch.Tensor,
        g: torch.Tensor,
249
250
251
252
253
254
255
256
257
        cu_seqlens: Optional[torch.LongTensor] = None,
        chunk_size: int = 64,
):
    """Finish local tensor gradients after the state-gradient handoff."""
    dq, dk, dw, dg = chunk_bwd_dqkwg(
        q=q,
        k=k,
        v=v_new,
        w=w,
263
264
265
266
267
268
269
270
271
        chunk_size=chunk_size,
        scale=scale,
        cu_seqlens=cu_seqlens,
    )
    dk2, dv, db, dg2 = prepare_wy_repr_bwd(
        k=k,
        v=v,
        beta=beta,
        g=g,
274
275
276
277
278
279
280
281
282
283
284
285
286
        du=dv,
        cu_seqlens=cu_seqlens,
        chunk_size=chunk_size,
    )
    dk.add_(dk2)
    dg.add_(dg2)
    if dg.dtype != torch.float32:
        raise ValueError(f"dg current type is {dg.dtype} , should be float32")
    dg = chunk_local_cumsum(
        dg,
        chunk_size=chunk_size,
        reverse=True,
        cu_seqlens=cu_seqlens,
285
286
287
288
289
290
291
292
293
294
295
296
        reverse=True,
        cu_seqlens=cu_seqlens,
        head_first=False,
    )
    return dq, dk, dv, db, dg


def chunk_gated_delta_rule_bwd(
        q: torch.Tensor,
        k: torch.Tensor,
        v: torch.Tensor,
        g: torch.Tensor,
302
303
304
305
306
307
308
309
310
        dht: torch.Tensor,
        cu_seqlens: Optional[torch.LongTensor] = None,
        chunk_size: int = 64,
):
    w, h, v_new, dv = chunk_gated_delta_rule_bwd_prepare(
        q=q,
        k=k,
        v=v,
        g=g,
315
316
317
318
319
320
321
322
323
        do=do,
        cu_seqlens=cu_seqlens,
        chunk_size=chunk_size,
    )
    dh, dh0, dv = chunk_gated_delta_rule_bwd_state(
        q=q,
        k=k,
        w=w,
        g=g,
328
329
330
331
332
333
334
335
336
        scale=scale,
        cu_seqlens=cu_seqlens,
        chunk_size=chunk_size,
    )
    dq, dk, dv, db, dg = chunk_gated_delta_rule_bwd_finish(
        q=q,
        k=k,
        v=v,
        g=g,
345
346
347
348
349
350
351
352
353
354
355
356
357
358
        scale=scale,
        cu_seqlens=cu_seqlens,
        chunk_size=chunk_size,
    )
    return dq, dk, dv, db, dg, dh0


@torch.compiler.disable
@input_guard
def chunk_gated_delta_rule_fwd_prepare_saved(
        q: torch.Tensor,
        k: torch.Tensor,
        v: torch.Tensor,
        g: torch.Tensor,
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
        use_qk_l2norm_in_kernel: bool = False,
        chunk_size: int = 64,
):
    """Prepare fused GDN forward tensors without consuming the initial state."""
    if scale is None:
        scale = k.shape[-1] ** -0.5

    q_norm, q_inv_norm = q, q.new_empty(0)
    k_norm, k_inv_norm = k, k.new_empty(0)
    if use_qk_l2norm_in_kernel:
        q_norm, q_inv_norm = _l2norm(q)
        k_norm, k_inv_norm = _l2norm(k)

    g_cumsum, A, w, u = chunk_gated_delta_rule_fwd_prepare(
        k=k_norm,
        v=v,
        g=g,
        beta=beta,
377
378
379
380
381
382
383
384
385
        g=g,
        beta=beta,
        chunk_size=chunk_size,
    )
    return (
        q_norm,
        k_norm,
        q_inv_norm,
        k_inv_norm,
390
391
392
393
394
395
396
397
398
399
400
        scale,
    )


@torch.compiler.disable
@input_guard
def chunk_gated_delta_rule_fwd_apply_state_saved(
        k_norm: torch.Tensor,
        g_cumsum: torch.Tensor,
        w: torch.Tensor,
        u: torch.Tensor,
402
403
404
405
406
407
408
409
410
        output_final_state: bool = True,
        chunk_size: int = 64,
):
    """Apply an initial state to fused prepared forward tensors."""
    return chunk_gated_delta_rule_fwd_apply_state(
        k=k_norm,
        g=g_cumsum,
        w=w,
        u=u,
413
414
415
416
417
418
419
420
421
422
423
        chunk_size=chunk_size,
    )


@torch.compiler.disable
@input_guard
def chunk_gated_delta_rule_fwd_output_saved(
        q_norm: torch.Tensor,
        k_norm: torch.Tensor,
        g_cumsum: torch.Tensor,
        h: torch.Tensor,
425
426
427
428
429
430
431
432
433
        scale: float,
        chunk_size: int = 64,
):
    """Compute fused GDN output from prepared, state-applied tensors."""
    return chunk_gated_delta_rule_fwd_output(
        q=q_norm,
        k=k_norm,
        v_new=v_new,
        h=h,
436
437
438
439
440
441
442
443
444
445
446
        chunk_size=chunk_size,
    )


@torch.compiler.disable
@input_guard
def chunk_gated_delta_rule_fwd_saved(
        q: torch.Tensor,
        k: torch.Tensor,
        v: torch.Tensor,
        g: torch.Tensor,
450
451
452
453
454
455
456
457
458
        use_qk_l2norm_in_kernel: bool = False,
        chunk_size: int = 64,
):
    """Run fused GDN forward and return the tensors required by its backward."""
    (
        q_norm,
        k_norm,
        q_inv_norm,
        k_inv_norm,
470
471
472
473
474
475
476
477
478
        scale=scale,
        use_qk_l2norm_in_kernel=use_qk_l2norm_in_kernel,
        chunk_size=chunk_size,
    )
    h, v_new, final_state = chunk_gated_delta_rule_fwd_apply_state_saved(
        k_norm,
        g_cumsum,
        w,
        u,
479
480
481
482
483
484
485
486
487
        initial_state=initial_state,
        output_final_state=True,
        chunk_size=chunk_size,
    )
    output = chunk_gated_delta_rule_fwd_output_saved(
        q_norm,
        k_norm,
        g_cumsum,
        h,
488
489
490
491
492
493
494
495
496
        v_new,
        scale,
        chunk_size=chunk_size,
    )
    return (
        output.to(q.dtype),
        final_state,
        q_norm,
        k_norm,
501
502
503
504
505
506
507
508
509
510
511
        scale,
    )


@torch.compiler.disable
@input_guard
def chunk_gated_delta_rule_bwd_prepare_saved(
        q_norm: torch.Tensor,
        k_norm: torch.Tensor,
        v: torch.Tensor,
        g_cumsum: torch.Tensor,
516
517
518
519
520
521
522
523
524
        scale: float,
        chunk_size: int = 64,
):
    """Prepare fused backward tensors before the final-state grad arrives."""
    return chunk_gated_delta_rule_bwd_prepare(
        q=q_norm,
        k=k_norm,
        v=v,
        g=g_cumsum,
530
531
532
533
534
535
536
537
538
539
540
        chunk_size=chunk_size,
    )


@torch.compiler.disable
@input_guard
def chunk_gated_delta_rule_bwd_state_saved(
        q_norm: torch.Tensor,
        k_norm: torch.Tensor,
        g_cumsum: torch.Tensor,
        w: torch.Tensor,
545
546
547
548
549
550
551
552
553
        scale: float,
        chunk_size: int = 64,
):
    """Consume the final-state grad and produce the initial-state grad."""
    return chunk_gated_delta_rule_bwd_state(
        q=q_norm,
        k=k_norm,
        w=w,
        g=g_cumsum,
559
560
561
562
563
564
565
566
567
568
569
        chunk_size=chunk_size,
    )


@torch.compiler.disable
@input_guard
def chunk_gated_delta_rule_bwd_finish_saved(
        q: torch.Tensor,
        k: torch.Tensor,
        q_norm: torch.Tensor,
        k_norm: torch.Tensor,
583
584
585
586
587
588
589
590
591
        use_qk_l2norm_in_kernel: bool = False,
        chunk_size: int = 64,
):
    """Finish local fused gradients after the P2P state-gradient handoff."""
    dq, dk, dv, dbeta, dg = chunk_gated_delta_rule_bwd_finish(
        q=q_norm,
        k=k_norm,
        v=v,
        g=g_cumsum,
599
600
601
602
603
604
605
606
607
608
609
610
611
612
613
614
615
616
617
618
619
620
621
622
623
624
        dh=dh,
        scale=scale,
        chunk_size=chunk_size,
    )
    if use_qk_l2norm_in_kernel:
        with torch.enable_grad():
            q_leaf = q.detach().requires_grad_(True)
            k_leaf = k.detach().requires_grad_(True)
            q_recomputed, _ = _l2norm(q_leaf)
            k_recomputed, _ = _l2norm(k_leaf)
            dq, dk = torch.autograd.grad(
                (q_recomputed, k_recomputed),
                (q_leaf, k_leaf),
                grad_outputs=(dq, dk),
            )
    del q_inv_norm, k_inv_norm
    return dq, dk, dv, dg, dbeta


@torch.compiler.disable
@input_guard
def chunk_gated_delta_rule_bwd_saved(
        q: torch.Tensor,
        k: torch.Tensor,
        q_norm: torch.Tensor,
        k_norm: torch.Tensor,
635
636
637
638
639
640
641
642
643
        use_qk_l2norm_in_kernel: bool = False,
        chunk_size: int = 64,
):
    """Run fused GDN backward from a context saved by the forward helper."""
    w, h, v_new, dv = chunk_gated_delta_rule_bwd_prepare_saved(
        q_norm,
        k_norm,
        v,
        g_cumsum,
647
648
649
650
651
652
653
654
655
        grad_output,
        scale,
        chunk_size=chunk_size,
    )
    dh, dh0, dv = chunk_gated_delta_rule_bwd_state_saved(
        q_norm,
        k_norm,
        g_cumsum,
        w,
659
660
661
662
663
664
665
666
667
        dv,
        scale,
        chunk_size=chunk_size,
    )
    dq, dk, dv, dg, dbeta = chunk_gated_delta_rule_bwd_finish_saved(
        q,
        k,
        q_norm,
        k_norm,
680
681
682
683
684
685
686
687
688
689
690
691
692
693
694
695
696
697
        scale,
        use_qk_l2norm_in_kernel=use_qk_l2norm_in_kernel,
        chunk_size=chunk_size,
    )
    return dq, dk, dv, dg, dbeta, dh0


class ChunkGatedDeltaRuleFunction(torch.autograd.Function):
    """Autograd wrapper for the Triton-Ascend chunk Gated Delta Rule."""

    @staticmethod
    @input_guard
    @autocast_custom_fwd
    def forward(
            ctx,
            q: torch.Tensor,
            k: torch.Tensor,
            v: torch.Tensor,
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, o, A, final_state = chunk_gated_delta_rule_fwd(
            q=q,
            k=k,
            v=v,
            g=g,
716
717
718
719
720
721
722
723
724
725
726
727
728
729
730
731
732
733
734
735
736
737
738
739
740
741
742
743
744
745
746
            cu_seqlens=cu_seqlens,
            chunk_size=chunk_size,
        )

        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, 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 o.to(q.dtype), final_state

    @staticmethod
    @input_guard
    @autocast_custom_bwd
    def backward(
            ctx,
            do: torch.Tensor,
            dht: torch.Tensor
    ):
        q, k, v, g, beta, 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
        dq, dk, dv, db, dg, dh0 = chunk_gated_delta_rule_bwd(
            q=q,
            k=k,
            v=v,
            g=g,
752
753
754
755
756
757
758
759
760
761
762
763
764
            dht=dht,
            cu_seqlens=cu_seqlens,
            chunk_size=ctx.chunk_size,
        )
        return dq.to(q), dk.to(k), dv.to(v), dg.to(g), db.to(beta), None, dh0, None, None, None, None


@torch.compiler.disable
def chunk_gated_delta_rule(
        q: torch.Tensor,
        k: torch.Tensor,
        v: torch.Tensor,
        g: torch.Tensor,
834
835
836
837
838
839
840
841
842
843
844
845
846
847
848
849
850
851
852
853
854
855
856
857
858
859
860
861
862
            output_final_state=True,
            cu_seqlens=cu_seqlens
        )
    """
    if q.dtype != k.dtype or k.dtype != v.dtype:
        raise ValueError(
            f"q current type is {q.dtype}, k current type is {k.dtype}, "
            f"v current type is {v.dtype}, they should be equal"
        )
    if q.dtype == torch.float32:
        raise ValueError(
            "ChunkGatedDeltaRuleFunction does not support float32. Please use bfloat16."
        )
    if len(beta.shape) != 3:
        raise ValueError(
            f"beta current shape len is {len(beta.shape)}, beta must be of shape [B, T, H] if head_first=False, or [B, H, T] otherwise."
        )

    if head_first:
        warnings.warn(
            "head_first is deprecated and will be removed in a future version. "
            "Please use head_first=False for now instead."
        )
    if not head_first and q.shape[1] < q.shape[2]:
        warnings.warn(
            f"Input tensor shape suggests potential format mismatch: seq_len ({q.shape[1]}) < num_heads ({q.shape[2]}). "
            "This may indicate the inputs were passed in head-first format [B, H, T, ...] "
            "when head_first=False was specified. "
            "Please verify your input tensor format matches the expected shape [B, T, H, ...]."
860
861
862
863
864
865
866
867
868
869
870
871
872
873
874
875
876
877
878
879
880
881
882
883
884
885
886
            "This may indicate the inputs were passed in head-first format [B, H, T, ...] "
            "when head_first=False was specified. "
            "Please verify your input tensor format matches the expected shape [B, T, H, ...]."
        )
    if cu_seqlens is not None:
        if q.shape[0] != 1:
            raise ValueError(
                f"The batch size is expected to be 1 rather than {q.shape[0]} when using `cu_seqlens`."
                f"Please flatten variable-length inputs before processing."
            )
        if initial_state is not None and initial_state.shape[0] != len(cu_seqlens) - 1:
            raise ValueError(
                f"The number of initial states is expected to be equal to the number of input sequences, "
                f"i.e., {len(cu_seqlens) - 1} rather than {initial_state.shape[0]}."
            )
    if scale is None:
        scale = k.shape[-1] ** -0.5

    if use_qk_l2norm_in_kernel:
        q, _ = _l2norm(q)
        k, _ = _l2norm(k)

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

from typing import Optional

import torch
import triton

from ._gdn_triton.state_summary import (
    gdn_packed_state_summary_kernel,
    gdn_state_grad_ext_kernel,
)
24
25
26
27
28
29
30
31
32
33
34
35
36
37
    gdn_state_grad_ext_kernel,
)


def _validate_fixed_summary_shape(
    key_dim: int,
    value_dim: int,
    chunk_size: int,
) -> None:
    if key_dim != 128 or value_dim != 128 or chunk_size != 64:
        raise NotImplementedError(
            "Triton GDN state summary requires key_dim=value_dim=128 and "
            "chunk_size=64."
        )
36
37
38
39
40
41
42
43
44
45
            "chunk_size=64."
        )


@torch.compiler.disable
def chunk_gated_delta_rule_state_summary_fwd(
    key: torch.Tensor,
    w: torch.Tensor,
    u: torch.Tensor,
    g: torch.Tensor,
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
    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 != 4 or w.ndim != 4 or u.ndim != 4 or g.ndim != 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)
    if block_size not in (64, 128):
        raise ValueError(f"GDN state-summary block_size must be 64 or 128, got {block_size}.")
    if seq_len % chunk_size != 0:
        raise ValueError(
            f"GDN state-summary sequence length {seq_len} must be divisible by {chunk_size}."
        )
    if w.shape != key.shape or u.shape[:3] != key.shape[:3] or g.shape != key.shape[:3]:
        raise ValueError(
            "Incompatible GDN state-summary shapes: "
            f"key={tuple(key.shape)}, w={tuple(w.shape)}, "
            f"u={tuple(u.shape)}, g={tuple(g.shape)}."
        )
65
66
67
68
69
70
71
72
73
74
            f"key={tuple(key.shape)}, w={tuple(w.shape)}, "
            f"u={tuple(u.shape)}, g={tuple(g.shape)}."
        )

    key, w, u, g = (tensor.contiguous() for tensor in (key, w, u, g))
    packed_summary = torch.empty(
        batch,
        heads,
        key_dim,
        value_dim + key_dim,
74
75
76
77
78
79
80
81
82
        value_dim + key_dim,
        device=key.device,
        dtype=torch.float32,
    )
    gdn_packed_state_summary_kernel[
        (triton.cdiv(value_dim + key_dim, block_size), batch * heads)
    ](
        key,
        w,
 90
 91
 92
 93
 94
 95
 96
 97
 98
 99
100
101
102
103
104
        BT=chunk_size,
        BV=block_size,
        NT=seq_len // chunk_size,
    )
    state_ext = packed_summary[..., :value_dim].contiguous()
    transition = packed_summary[..., value_dim:].contiguous()
    return state_ext, transition


@torch.compiler.disable
def chunk_gated_delta_rule_state_gradient_summary_bwd(
    query: torch.Tensor,
    key: torch.Tensor,
    w: torch.Tensor,
    g: torch.Tensor,
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
    *,
    chunk_size: int = 64,
) -> torch.Tensor:
    """Return the local-loss contribution to the incoming state gradient."""
    batch, seq_len, heads, key_dim = query.shape
    value_dim = grad_output.shape[-1]
    _validate_fixed_summary_shape(key_dim, value_dim, chunk_size)
    if seq_len % chunk_size != 0:
        raise ValueError(
            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 != qk_shape
        or w.shape != qk_shape
        or g.shape != qk_shape[:3]
        or grad_output.shape != value_shape
124
125
126
127
128
129
130
131
132
        or g.shape != qk_shape[:3]
        or grad_output.shape != value_shape
        or dv.shape != value_shape
    ):
        raise ValueError(
            "Incompatible GDN state-gradient summary shapes: "
            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)}."
131
132
133
134
135
136
137
138
139
140
141
142
            f"w={tuple(w.shape)}, g={tuple(g.shape)}, "
            f"grad_output={tuple(grad_output.shape)}, dv={tuple(dv.shape)}."
        )

    query, key, w, g, grad_output, dv = (
        tensor.contiguous() for tensor in (query, key, w, g, grad_output, dv)
    )
    grad_state_ext = torch.empty(
        batch,
        heads,
        key_dim,
        value_dim,
142
143
144
145
146
147
148
149
150
        value_dim,
        device=query.device,
        dtype=torch.float32,
    )
    gdn_state_grad_ext_kernel[(1, batch * heads)](
        query,
        key,
        w,
        g,
159
160
161
162
163
164
165
166
167
168
169
170
        BT=chunk_size,
        BV=128,
        NT=seq_len // chunk_size,
    )
    return grad_state_ext


def apply_gdn_state_summary(
    state_ext: torch.Tensor,
    transition: torch.Tensor,
    initial_state: Optional[torch.Tensor],
) -> torch.Tensor:
168
169
170
171
172
173
174
175
176
177
178
179
180
181
    transition: torch.Tensor,
    initial_state: Optional[torch.Tensor],
) -> torch.Tensor:
    """Apply a local affine state summary in FP32."""
    if initial_state is None:
        return state_ext
    return torch.matmul(transition, initial_state.float()) + state_ext


def apply_gdn_state_gradient_summary(
    grad_state_ext: torch.Tensor,
    transition: torch.Tensor,
    grad_final_state: Optional[torch.Tensor],
) -> torch.Tensor:
179
180
181
182
183
184
185
186
187
188
    transition: torch.Tensor,
    grad_final_state: Optional[torch.Tensor],
) -> torch.Tensor:
    """Apply the adjoint affine summary to a gradient from the next rank."""
    if grad_final_state is None:
        return grad_state_ext
    return (
        torch.matmul(transition.transpose(-2, -1), grad_final_state.float())
        + grad_state_ext
    )
187
188
189
190
191
192
193
194
195
        + grad_state_ext
    )


__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",
hyper_parallel/components/functional/kimi_delta_attention.py
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
# See the License for the specific language governing permissions and
# limitations under the License.
# ============================================================================
"""Dense Triton-Ascend KDA execution boundaries for Torch training."""
from __future__ import annotations

from typing import Any, Optional

import torch
import torch.distributed as dist

from .kimi_delta_attention_fla_adapter import get_fla_kda_staged_ops, run_fla_chunk_kda
from .kimi_delta_attention_state_summary import (
    apply_kda_state_gradient_summary,
    apply_kda_state_summary,
    kda_state_gradient_summary_from_prepared,
    kda_state_summary_forward_from_prepared,
28
29
30
31
32
33
34
35
36
    kda_state_summary_forward_from_prepared,
)


def _validate_local_inputs(
    query: torch.Tensor,
    key: torch.Tensor,
    value: torch.Tensor,
    gate: torch.Tensor,
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
    chunk_size: int,
    lower_bound: float,
) -> None:
    """Validate the fixed-shape local KDA backend before compiling kernels."""
    if query.ndim != 4 or key.ndim != 4 or value.ndim != 4 or gate.ndim != 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:
        raise ValueError("Fused KDA query and key must have identical shapes.")
    batch, sequence_length, num_query_heads, key_dim = query.shape
    num_value_heads, value_dim = value.shape[2:]
    if value.shape[:2] != (batch, sequence_length):
        raise ValueError("Fused KDA value must match query batch and sequence dimensions.")
    if gate.shape != (batch, sequence_length, num_value_heads, key_dim):
        raise ValueError("Fused KDA gate has an incompatible shape.")
    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.")
    if not (
        query.dtype == key.dtype == value.dtype == gate.dtype == beta.dtype
        == torch.bfloat16
    ):
        raise TypeError("Fused KDA requires q/k/v/gate/beta tensors in bfloat16.")
    if a_log.dtype != torch.float32 or dt_bias.dtype != torch.float32:
        raise TypeError("Fused KDA requires a_log and dt_bias tensors in float32.")
    if key_dim != 128 or value_dim != 128 or chunk_size != 64:
        raise NotImplementedError(
            "Fused KDA currently requires key_dim=value_dim=128 and chunk_size=64."
        )
    if sequence_length % chunk_size:
        raise ValueError(
            f"Fused KDA sequence length {sequence_length} must be divisible "
            f"by chunk_size {chunk_size}."
        )
    if a_log.numel() != num_value_heads:
        raise ValueError("Fused KDA a_log must contain one value per value head.")
    if dt_bias.numel() != num_value_heads * key_dim:
        raise ValueError("Fused KDA dt_bias must contain one value per gate channel.")
    if not -5.0 <= lower_bound < 0:
        raise ValueError("Fused KDA lower_bound must lie in [-5, 0).")
    if not query.is_npu:
        raise RuntimeError("The fused KDA backend requires Ascend NPU tensors.")
    if any(
        tensor.device != query.device
        for tensor in (key, value, gate, beta, a_log, dt_bias)
    ):
        raise ValueError("Fused KDA inputs must reside on the same NPU device.")
    get_fla_kda_staged_ops()


def _validate_p2p_inputs(
    query: torch.Tensor,
    key: torch.Tensor,
    value: torch.Tensor,
    gate: torch.Tensor,
104
105
106
107
108
109
110
111
112
    cp_rank: int,
    cp_size: int,
) -> None:
    """Validate the fixed-shape KDA P2P backend before communication starts."""
    _validate_local_inputs(
        query,
        key,
        value,
        gate,
115
116
117
118
119
120
121
122
123
124
125
126
127
        dt_bias,
        chunk_size=chunk_size,
        lower_bound=lower_bound,
    )
    if cp_size <= 0 or not 0 <= cp_rank < cp_size:
        raise ValueError(f"Invalid KDA P2P rank {cp_rank} for cp_size {cp_size}.")


def fused_chunk_kda(
    query: torch.Tensor,
    key: torch.Tensor,
    value: torch.Tensor,
    gate: torch.Tensor,
134
135
136
137
138
139
140
141
142
    chunk_size: int = 64,
    safe_gate: bool = True,
) -> torch.Tensor:
    """Run the local Triton-Ascend KDA backend without CP communication."""
    _validate_local_inputs(
        query,
        key,
        value,
        gate,
145
146
147
148
149
150
151
152
153
        dt_bias,
        chunk_size=chunk_size,
        lower_bound=lower_bound,
    )
    output, _ = run_fla_chunk_kda(
        query,
        key,
        value,
        gate,
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
        lower_bound=lower_bound,
        chunk_size=chunk_size,
        safe_gate=safe_gate,
    )
    return output


class _KDAStateP2PFunction(torch.autograd.Function):
    """Run fused local KDA around an affine recurrent-state P2P wavefront."""

    @staticmethod
    def forward(  # pylint: disable=arguments-differ,too-many-arguments,too-many-locals
        ctx: Any,
        query: torch.Tensor,
        key: torch.Tensor,
        value: torch.Tensor,
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
        cp_rank: int,
        cp_size: int,
    ) -> torch.Tensor:
        """Prepare local KDA, propagate its state, and produce token outputs."""
        ops = get_fla_kda_staged_ops()

        rcp_ln2 = 1.4426950408889634

        query, key, value, gate_raw, beta_raw = (
            tensor.contiguous()
            for tensor in (query, key, value, gate_raw, beta_raw)
        )
        query, query_rstd = ops.l2norm_fwd(query)
        key, key_rstd = ops.l2norm_fwd(key)
        beta = ops.fused_beta_sigmoid(beta_raw)
        gate = ops.kda_gate_chunk_cumsum(
            g=gate_raw,
            A_log=a_log,
            dt_bias=dt_bias,
            scale=rcp_ln2,
205
206
207
208
209
210
211
212
213
            chunk_size=chunk_size,
            lower_bound=lower_bound,
        )

        state_shape = (
            query.shape[0],
            value.shape[2],
            query.shape[-1],
            value.shape[-1],
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
            value.shape[2],
            query.shape[-1],
            value.shape[-1],
        )
        initial_state = None
        recv_work = None
        if cp_rank > 0:
            initial_state = torch.empty(
                state_shape,
                device=query.device,
                dtype=torch.float32,
            )
            recv_work = dist.irecv(
                initial_state,
                src=prev_rank,
                group=cp_group,
            )
225
226
227
228
229
230
231
232
233
                src=prev_rank,
                group=cp_group,
            )

        w, u, _, kg, attention_qk, attention_kk = ops.chunk_kda_fwd_intra(
            q=query,
            k=key,
            v=value,
            gk=gate,
236
237
238
239
240
241
242
243
244
245
246
247
            chunk_size=chunk_size,
            safe_gate=safe_gate,
            disable_recompute=True,
        )
        state_ext = None
        transition = None
        if cp_rank < cp_size - 1:
            state_ext, transition = kda_state_summary_forward_from_prepared(
                kg,
                w,
                u,
                gate,
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
                u,
                gate,
                chunk_size=chunk_size,
            )
        if recv_work is not None:
            recv_work.wait()

        send_work = None
        send_state = None
        if cp_rank < cp_size - 1:
            final_state = apply_kda_state_summary(
                state_ext,
                transition,
                initial_state,
            )
            send_state = final_state.contiguous()
            send_work = dist.isend(
                send_state,
                dst=next_rank,
                group=cp_group,
            )
264
265
266
267
268
269
270
271
272
                dst=next_rank,
                group=cp_group,
            )

        states, value_new, _ = ops.chunk_gated_delta_rule_fwd_h(
            k=kg,
            w=w,
            u=u,
            gk=gate,
273
274
275
276
277
278
279
280
281
            initial_state=initial_state,
            output_final_state=False,
            chunk_size=chunk_size,
        )
        output = ops.chunk_gla_fwd_o_gk(
            q=query,
            v=value_new,
            g=gate,
            A=attention_qk,
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
            h=states,
            scale=scale,
            chunk_size=chunk_size,
        )
        if send_work is not None:
            send_work.wait()

        saved_initial_state = initial_state
        if saved_initial_state is None:
            saved_initial_state = query.new_empty(0, dtype=torch.float32)
        ctx.save_for_backward(
            query,
            query_rstd,
            key,
            key_rstd,
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
            attention_kk,
            transition if transition is not None else query.new_empty(0),
            saved_initial_state,
        )
        ctx.scale = scale
        ctx.lower_bound = lower_bound
        ctx.chunk_size = chunk_size
        ctx.safe_gate = safe_gate
        ctx.cp_group = cp_group
        ctx.prev_rank = prev_rank
        ctx.next_rank = next_rank
        ctx.cp_rank = cp_rank
        ctx.cp_size = cp_size
        return output.to(value.dtype)

    @staticmethod
    def backward(  # pylint: disable=arguments-differ,too-many-locals
        ctx: Any,
        grad_output: torch.Tensor,
    ) -> tuple[Any, ...]:
        """Reverse the state wavefront, then finish the fused local backward."""
        ops = get_fla_kda_staged_ops()

        rcp_ln2 = 1.4426950408889634

        (
            query,
            query_rstd,
            key,
            key_rstd,
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
            attention_kk,
            transition,
            saved_initial_state,
        ) = ctx.saved_tensors
        initial_state = saved_initial_state if saved_initial_state.numel() else None
        grad_output = grad_output.contiguous()

        grad_final_state = None
        recv_work = None
        if ctx.cp_rank < ctx.cp_size - 1:
            grad_final_state = torch.empty(
                (
                    query.shape[0],
                    value.shape[2],
                    query.shape[-1],
355
356
357
358
359
360
361
362
363
                ),
                device=query.device,
                dtype=torch.float32,
            )
            recv_work = dist.irecv(
                grad_final_state,
                src=ctx.next_rank,
                group=ctx.cp_group,
            )
361
362
363
364
365
366
367
368
369
370
                src=ctx.next_rank,
                group=ctx.cp_group,
            )

        beta = ops.fused_beta_sigmoid(beta_raw)
        gate = ops.kda_gate_chunk_cumsum(
            g=gate_raw,
            A_log=a_log,
            dt_bias=dt_bias,
            scale=rcp_ln2,
370
371
372
373
374
375
376
377
378
            scale=rcp_ln2,
            chunk_size=ctx.chunk_size,
            lower_bound=ctx.lower_bound,
        )
        w, u, query_gated, key_gated = ops.recompute_w_u_fwd(
            q=query,
            k=key,
            v=value,
            beta=beta,
378
379
380
381
382
383
384
385
386
            beta=beta,
            A=attention_kk,
            gk=gate,
        )
        states, value_new, _ = ops.chunk_gated_delta_rule_fwd_h(
            k=key_gated,
            w=w,
            u=u,
            gk=gate,
387
388
389
390
391
392
393
394
395
            initial_state=initial_state,
            output_final_state=False,
            chunk_size=ctx.chunk_size,
        )
        grad_attention_qk, grad_value_local = ops.chunk_kda_bwd_dav(
            q=query,
            k=key,
            v=value_new,
            do=grad_output,
396
397
398
399
400
401
402
403
404
405
406
            A=attention_qk,
            scale=ctx.scale,
            chunk_size=ctx.chunk_size,
        )
        grad_state_ext = None
        if ctx.cp_rank > 0:
            grad_state_ext = kda_state_gradient_summary_from_prepared(
                query_gated,
                key_gated,
                w,
                gate,
408
409
410
411
412
413
414
415
416
417
418
419
420
421
422
423
424
425
426
427
                grad_value_local,
                ctx.scale,
                chunk_size=ctx.chunk_size,
            )
        if recv_work is not None:
            recv_work.wait()
        send_work = None
        send_state_gradient = None
        if ctx.cp_rank > 0:
            grad_initial_state = apply_kda_state_gradient_summary(
                grad_state_ext,
                transition,
                grad_final_state,
            )
            send_state_gradient = grad_initial_state.contiguous()
            send_work = dist.isend(
                send_state_gradient,
                dst=ctx.prev_rank,
                group=ctx.cp_group,
            )
425
426
427
428
429
430
431
432
433
                dst=ctx.prev_rank,
                group=ctx.cp_group,
            )

        grad_states, _, grad_value = ops.chunk_gated_delta_rule_bwd_dhu(
            q=query_gated,
            k=key_gated,
            w=w,
            gk=gate,
437
438
439
440
441
442
443
444
445
            dv=grad_value_local,
            scale=ctx.scale,
            chunk_size=ctx.chunk_size,
        )
        (
            grad_query,
            grad_key,
            grad_value,
            grad_beta,
459
460
461
462
463
464
465
466
467
            dv=grad_value,
            scale=ctx.scale,
            chunk_size=ctx.chunk_size,
        )
        grad_query, grad_key, grad_beta, grad_gate = ops.chunk_kda_bwd_intra(
            q=query,
            k=key,
            g=gate,
            beta=beta,
474
475
476
477
478
479
480
481
482
483
484
485
486
            chunk_size=ctx.chunk_size,
            safe_gate=ctx.safe_gate,
        )

        num_query_heads = query.shape[2]
        num_value_heads = value.shape[2]
        if num_value_heads > num_query_heads:
            groups = num_value_heads // num_query_heads
            grad_query = grad_query.view(
                *grad_query.shape[:2],
                num_query_heads,
                groups,
                grad_query.shape[-1],
484
485
486
487
488
489
490
491
492
                num_query_heads,
                groups,
                grad_query.shape[-1],
            ).sum(dim=3)
            grad_key = grad_key.view(
                *grad_key.shape[:2],
                num_query_heads,
                groups,
                grad_key.shape[-1],
491
492
493
494
495
496
497
498
499
500
501
502
503
504
                groups,
                grad_key.shape[-1],
            ).sum(dim=3)

        grad_gate = ops.chunk_local_cumsum(
            grad_gate,
            chunk_size=ctx.chunk_size,
            reverse=True,
        )
        grad_gate, grad_a_log, grad_dt_bias = ops.kda_gate_bwd(
            g=gate_raw,
            A_log=a_log,
            dt_bias=dt_bias,
            dyg=grad_gate,
503
504
505
506
507
508
509
510
511
512
513
514
515
516
517
            dt_bias=dt_bias,
            dyg=grad_gate,
            lower_bound=ctx.lower_bound,
        )
        grad_beta = ops.fused_beta_sigmoid_bwd(beta_raw, grad_beta)
        grad_query = ops.l2norm_bwd(query, query_rstd, grad_query)
        grad_key = ops.l2norm_bwd(key, key_rstd, grad_key)
        if send_work is not None:
            send_work.wait()

        return (
            grad_query,
            grad_key,
            grad_value,
            grad_gate,
529
530
531
532
533
534
535
536
537
538
            None,
        )


@torch.compiler.disable
def fused_chunk_kda_p2p(
    query: torch.Tensor,
    key: torch.Tensor,
    value: torch.Tensor,
    gate: torch.Tensor,
550
551
552
553
554
555
556
557
558
    chunk_size: int = 64,
    safe_gate: bool = True,
) -> torch.Tensor:
    """Run one sequence-sharded KDA segment with state P2P."""
    _validate_p2p_inputs(
        query,
        key,
        value,
        gate,
563
564
565
566
567
568
569
570
571
572
        lower_bound=lower_bound,
        cp_rank=cp_rank,
        cp_size=cp_size,
    )
    effective_scale = query.shape[-1] ** -0.5 if scale is None else float(scale)
    return _KDAStateP2PFunction.apply(
        query,
        key,
        value,
        gate,
584
585
586
587
588
        cp_size,
    )


__all__ = ["fused_chunk_kda", "fused_chunk_kda_p2p"]
hyper_parallel/components/functional/kimi_delta_attention_fla_adapter.py
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


def _get_attribute(module_name: str, attribute: str) -> Any:
    """Import one required FLA symbol with a useful compatibility error."""
    try:
        module = importlib.import_module(module_name)
    except ModuleNotFoundError as exc:
        missing_name = getattr(exc, "name", None)
        is_fla_module = (
            missing_name is not None and str(missing_name).startswith("fla.")
        )
        if missing_name == module_name or is_fla_module:
            raise RuntimeError(
                "The installed FLA package is incompatible with Hyper KDA: "
                f"missing module {module_name}."
            ) from exc
        raise RuntimeError(
            f"FLA module {module_name} could not load runtime dependency "
            f"{missing_name!r}."
        ) from exc
    except ImportError as exc:
        raise RuntimeError(
            f"FLA module {module_name} failed to import; inspect the chained "
            "exception for the incompatible runtime library."
        ) from exc
    try:
        return getattr(module, attribute)
    except AttributeError as exc:
        raise RuntimeError(
            "The installed FLA package is incompatible with Hyper KDA: "
            f"missing {module_name}.{attribute}. Install an FLA revision with "
            "the Triton-Ascend KDA backend."
        ) from exc
 94
 95
 96
 97
 98
 99
100
101
102
def _parse_version(version: str) -> tuple[int, int, int]:
    """Parse the numeric prefix of an FLA semantic version."""
    match = re.match(r"^(\d+)\.(\d+)\.(\d+)", version)
    if match is None:
        raise RuntimeError(f"Unable to parse the installed FLA version {version!r}.")
    return tuple(int(part) for part in match.groups())


def _require_triton_ascend() -> None:
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115


def _require_triton_ascend() -> None:
    """Require the external FLA runtime to have a Triton Ascend backend."""
    try:
        importlib.import_module("triton")
        ascend_backend = importlib.util.find_spec("triton.backends.ascend")
    except (ImportError, ModuleNotFoundError) as exc:
        raise RuntimeError(
            "KDA backend='triton' requires Triton-Ascend in the runtime environment."
        ) from exc
    if ascend_backend is None:
        raise RuntimeError(
            "KDA backend='triton' found Triton, but its Ascend backend is missing."
        )

115
116
117
118
119
120
121
122
123


def _require_npu_backends() -> None:
    """Verify that FLA registered every backend used by the staged P2P path."""
    backend_specs = (
        (
            "fla.ops.kda.backends.triton_ascend",
            "TritonAscendKDABackend",
        ),
129
130
131
132
133
134
135
136
137
138
139
            "fla.ops.gla.backends.triton_ascend",
            "TritonAscendGLABackend",
        ),
    )
    for module_name, class_name in backend_specs:
        backend = _get_attribute(module_name, class_name)
        if not backend.is_available():
            raise RuntimeError(
                "KDA backend='triton' requires FLA's Triton-Ascend backend "
                f"{class_name}, but it is unavailable in this process."
            )
146
147
148
149
150
151
152
153
154
        fla = importlib.import_module("fla")
    except ModuleNotFoundError as exc:
        missing_name = getattr(exc, "name", None)
        if missing_name != "fla":
            raise RuntimeError(
                "The external FLA package is present but cannot load runtime "
                f"dependency {missing_name!r}."
            ) from exc
        raise RuntimeError(
154
155
156
157
158
159
160
161
162
        raise RuntimeError(
            "KDA backend='triton' requires the optional flash-linear-attention "
            "package with its Triton-Ascend KDA backend."
        ) from exc
    except ImportError as exc:
        raise RuntimeError(
            "The external FLA package is present but failed to import; inspect "
            "the chained exception for the incompatible runtime library."
        ) from exc
162
163
164
165
166
167
168
169
170
        ) from exc

    version = getattr(fla, "__version__", None)
    if not isinstance(version, str):
        raise RuntimeError("The installed FLA package does not expose __version__.")
    if _parse_version(version) < _MIN_FLA_VERSION:
        minimum = ".".join(str(part) for part in _MIN_FLA_VERSION)
        raise RuntimeError(
            f"KDA backend='triton' requires FLA >= {minimum}, got {version}."
169
170
171
172
173
174
175
176
177
178
179
        raise RuntimeError(
            f"KDA backend='triton' requires FLA >= {minimum}, got {version}."
        )

    _require_triton_ascend()
    _require_npu_backends()
    staged = FLAKDAStagedOps(
        l2norm_fwd=_get_attribute("fla.modules.l2norm", "l2norm_fwd"),
        l2norm_bwd=_get_attribute("fla.modules.l2norm", "l2norm_bwd"),
        fused_beta_sigmoid=_get_attribute(
            "fla.ops.common.gate", "fused_beta_sigmoid"
212
213
214
215
216
217
218
219
220
        chunk_local_cumsum=_get_attribute(
            "fla.ops.utils.cumsum", "chunk_local_cumsum"
        ),
    )
    return _FLAKDARuntime(
        version=version,
        chunk_kda=_get_attribute("fla.ops.kda", "chunk_kda"),
        staged=staged,
    )
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238


def is_fla_triton_kda_available() -> bool:
    """Return whether the optional external FLA KDA backend is usable."""
    try:
        _require_fla_kda_runtime()
    except RuntimeError:
        return False
    return True


def get_fla_kda_staged_ops() -> FLAKDAStagedOps:
    """Return FLA staged operators for Hyper's state-P2P autograd function."""
    return _require_fla_kda_runtime().staged


def run_fla_chunk_kda(
    query: torch.Tensor,
hyper_parallel/components/functional/kimi_delta_attention_state_summary.py
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
# limitations under the License.
# ============================================================================
"""Affine KDA state summary built from prepared WY intermediates."""

from typing import Optional

import torch
import triton


def _validate_prepared_summary_inputs(
    key: torch.Tensor,
    w: torch.Tensor,
    u: torch.Tensor,
    gate: torch.Tensor,
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
    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 != 4 or w.ndim != 4 or u.ndim != 4 or gate.ndim != 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:
        raise NotImplementedError(
            "Triton KDA state summary requires key_dim=value_dim=128 and "
            "chunk_size=64."
        )
    if sequence_length % chunk_size:
        raise ValueError(
            f"KDA summary sequence length {sequence_length} must be divisible "
            f"by chunk_size {chunk_size}."
        )
    if w.shape != key.shape or gate.shape != key.shape:
        raise ValueError(
            "KDA prepared key, w, and gate must have identical shapes, got "
            f"key={tuple(key.shape)}, w={tuple(w.shape)}, gate={tuple(gate.shape)}."
        )
    if u.shape[:3] != key.shape[:3]:
        raise ValueError(
            f"KDA prepared u prefix {tuple(u.shape[:3])} must match "
            f"key prefix {tuple(key.shape[:3])}."
        )
    if not key.is_npu:
        raise RuntimeError("Triton KDA state summary requires Ascend NPU tensors.")
    return batch, sequence_length, heads, key_dim, value_dim


def _validate_gradient_summary_inputs(
    query: torch.Tensor,
    key: torch.Tensor,
    w: torch.Tensor,
    gate: torch.Tensor,
 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
    grad_value: torch.Tensor,
    chunk_size: int,
) -> tuple[int, int, int, int, int]:
    """Validate prepared tensors used by the reverse state wavefront."""
    tensors = (query, key, w, gate, grad_output, grad_value)
    if any(tensor.ndim != 4 for tensor in tensors):
        raise ValueError(
            "KDA gradient summary expects rank-4 query/key/w/gate/do/dv tensors."
        )
    batch, sequence_length, heads, key_dim = query.shape
    value_dim = grad_output.shape[-1]
    if key_dim != 128 or value_dim != 128 or chunk_size != 64:
        raise NotImplementedError(
            "Triton KDA state-gradient summary requires key_dim=value_dim=128 "
            "and chunk_size=64."
        )
    if sequence_length % chunk_size:
        raise ValueError(
            f"KDA gradient-summary sequence length {sequence_length} must be "
            f"divisible by chunk_size {chunk_size}."
        )
    prepared_shape = (batch, sequence_length, heads, key_dim)
    value_shape = (batch, sequence_length, heads, value_dim)
    if key.shape != prepared_shape or w.shape != prepared_shape:
        raise ValueError(
            "KDA prepared query, key, and w must have identical shapes, got "
            f"query={tuple(query.shape)}, key={tuple(key.shape)}, "
            f"w={tuple(w.shape)}."
        )
    if gate.shape != prepared_shape:
        raise ValueError(
            f"KDA prepared gate shape {tuple(gate.shape)} must match "
            f"query shape {prepared_shape}."
        )
    if grad_output.shape != value_shape or grad_value.shape != value_shape:
        raise ValueError(
            "KDA grad_output and grad_value must have shape "
            f"{value_shape}, got do={tuple(grad_output.shape)}, "
            f"dv={tuple(grad_value.shape)}."
        )
    if not query.is_npu:
        raise RuntimeError(
            "Triton KDA state-gradient summary requires Ascend NPU tensors."
        )
    return batch, sequence_length, heads, key_dim, value_dim


def _launch_kda_state_summary(
    key: torch.Tensor,
    w: torch.Tensor,
    u: torch.Tensor,
    gate: torch.Tensor,
118
119
120
121
122
123
124
125
126
    *,
    chunk_size: int,
) -> tuple[torch.Tensor, torch.Tensor]:
    """Launch the validated mixed-tile summary used by the P2P path."""
    return _launch_kda_mixed_state_summary(
        key,
        w,
        u,
        gate,
127
128
129
130
131
132
133
134
135
        chunk_size=chunk_size,
    )


def _launch_kda_state_gradient_summary(
    query: torch.Tensor,
    key: torch.Tensor,
    w: torch.Tensor,
    gate: torch.Tensor,
140
141
142
143
144
145
146
147
148
149
150
151
152
    chunk_size: int,
    block_size: int = 64,
) -> torch.Tensor:
    """Launch the reverse-wavefront Triton-Ascend summary kernel."""
    from ._kda_triton.state_summary import (  # pylint: disable=import-outside-toplevel
        kda_state_grad_ext_kernel,
    )

    batch, sequence_length, heads, key_dim, value_dim = (
        _validate_gradient_summary_inputs(
            query,
            key,
            w,
155
156
157
158
159
160
161
162
163
164
165
166
167
            grad_value,
            chunk_size,
        )
    )
    query, key, w, gate, grad_output, grad_value = (
        tensor.contiguous()
        for tensor in (query, key, w, gate, grad_output, grad_value)
    )
    grad_state_ext = torch.empty(
        batch,
        heads,
        key_dim,
        value_dim,
167
168
169
170
171
172
173
174
175
        value_dim,
        dtype=torch.float32,
        device=query.device,
    )
    kda_state_grad_ext_kernel[
        (triton.cdiv(value_dim, block_size), batch * heads)
    ](
        query,
        key,
188
189
190
191
192
193
194
195
196
197
198
199
200
        BT=chunk_size,
        BV=block_size,
        num_warps=2,
    )
    return grad_state_ext


@torch.compiler.disable
def _launch_kda_mixed_state_summary(
    key: torch.Tensor,
    w: torch.Tensor,
    u: torch.Tensor,
    gate: torch.Tensor,
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
    *,
    chunk_size: int = 64,
) -> tuple[torch.Tensor, torch.Tensor]:
    """Build ``S_ext`` and ``M`` in compile-time-separated BV=128 modes."""
    from ._kda_triton.state_summary import (  # pylint: disable=import-outside-toplevel
        kda_split_state_summary_kernel,
    )

    batch, sequence_length, heads, key_dim, value_dim = (
        _validate_prepared_summary_inputs(key, w, u, gate, chunk_size)
    )
    key, w, u, gate = (
        tensor.contiguous() for tensor in (key, w, u, gate)
    )
    state_ext = torch.empty(
        batch,
        heads,
        key_dim,
        value_dim,
219
220
221
222
223
224
225
226
227
        value_dim,
        dtype=torch.float32,
        device=key.device,
    )
    transition = torch.empty(
        batch,
        heads,
        key_dim,
        key_dim,
227
228
229
230
231
232
233
234
235
        key_dim,
        dtype=torch.float32,
        device=key.device,
    )
    kda_split_state_summary_kernel[(1, batch * heads)](
        key,
        w,
        u,
        gate,
245
246
247
248
249
250
251
252
253
        BV=128,
        OUTPUT_MODE=1,
        num_warps=2,
    )
    kda_split_state_summary_kernel[(1, batch * heads)](
        key,
        w,
        u,
        gate,
263
264
265
266
267
268
269
270
271
272
273
274
275
        BV=128,
        OUTPUT_MODE=2,
        num_warps=1,
    )
    return state_ext, transition


@torch.compiler.disable
def kda_state_summary_forward_from_prepared(
    key: torch.Tensor,
    w: torch.Tensor,
    u: torch.Tensor,
    gate: torch.Tensor,
276
277
278
279
280
281
282
283
284
    *,
    chunk_size: int = 64,
) -> tuple[torch.Tensor, torch.Tensor]:
    """Build a fused summary for an outer P2P custom-autograd function."""
    return _launch_kda_state_summary(
        key,
        w,
        u,
        gate,
285
286
287
288
289
290
291
292
293
294
        chunk_size=chunk_size,
    )


@torch.compiler.disable
def kda_state_gradient_summary_from_prepared(
    query: torch.Tensor,
    key: torch.Tensor,
    w: torch.Tensor,
    gate: torch.Tensor,
303
304
305
306
307
308
309
310
311
    ``grad_value`` is read-only in this operator. Callers must build this
    summary before passing the same buffer to a local backward implementation
    that updates ``grad_value`` in place, such as FLA ``bwd_dhu``.
    """
    return _launch_kda_state_gradient_summary(
        query,
        key,
        w,
        gate,
316
317
318
319
320
321
322
323
324
        block_size=128,
    )


def apply_kda_state_summary(
    state_ext: torch.Tensor,
    transition: torch.Tensor,
    initial_state: Optional[torch.Tensor],
) -> torch.Tensor:
322
323
324
325
326
327
328
329
330
331
332
333
334
335
    transition: torch.Tensor,
    initial_state: Optional[torch.Tensor],
) -> torch.Tensor:
    """Apply a prepared KDA state summary in FP32."""
    if initial_state is None:
        return state_ext
    return torch.matmul(transition, initial_state.float()) + state_ext


def apply_kda_state_gradient_summary(
    grad_state_ext: torch.Tensor,
    transition: torch.Tensor,
    grad_final_state: Optional[torch.Tensor],
) -> torch.Tensor:
333
334
335
336
337
338
339
340
341
342
    transition: torch.Tensor,
    grad_final_state: Optional[torch.Tensor],
) -> torch.Tensor:
    """Apply the adjoint summary for the reverse P2P state wavefront."""
    if grad_final_state is None:
        return grad_state_ext
    return (
        torch.matmul(transition.transpose(-2, -1), grad_final_state.float())
        + grad_state_ext
    )
341
342
343
344
345
346
347
348
349
        + grad_state_ext
    )


__all__ = [
    "apply_kda_state_gradient_summary",
    "apply_kda_state_summary",
    "kda_state_gradient_summary_from_prepared",
    "kda_state_summary_forward_from_prepared",
hyper_parallel/components/modules/gated_delta_net.py
243
244
245
246
247
248
249
250
251
            "q/k/v/beta=bf16, g=fp32, head_k_dim=head_v_dim=128, "
            "chunk_size=64, and sequence length divisible by 64."
        )

    from hyper_parallel.components.functional.gated_delta_net import (  # pylint: disable=import-outside-toplevel
        chunk_gated_delta_rule as triton_chunk_gated_delta_rule,
    )

    return triton_chunk_gated_delta_rule(
hyper_parallel/components/modules/kimi_delta_attention.py
110
111
112
113
114
115
116
117
        chunk_size=chunk_size,
        lower_bound=lower_bound,
    ):
        return False
    from hyper_parallel.components.functional.kimi_delta_attention_fla_adapter import (  # pylint: disable=import-outside-toplevel
        is_fla_triton_kda_available,
    )
    return is_fla_triton_kda_available()
555
556
557
558
559
560
561
562
563
            "q/k/v/gate/beta=bf16, a_log/dt_bias=fp32, "
            "head_k_dim=head_v_dim=128, chunk_size=64, and a supported safe gate."
        )

    from hyper_parallel.components.functional.kimi_delta_attention import (  # pylint: disable=import-outside-toplevel
        fused_chunk_kda,
    )

    output = fused_chunk_kda(
hyper_parallel/core/context_parallel/context_parallel.py
872
873
874
875
876
877
878
879
880
        half = q.shape[seq_dim] // 2
        q_keep = q.narrow(seq_dim, 0, half)
        q_mine = q.narrow(seq_dim, half, half)

        q_peer = utils.p2p_exchange(q_mine, peer_rank)
        k_full = _gather_seq(new_args[k_idx], co_submesh, seq_dim).to_local()
        v_full = _gather_seq(new_args[v_idx], co_submesh, seq_dim).to_local()

        # K/V are Replicate; wrap once and reuse for both FA calls
891
892
893
894
895
896
897
            return out.to_local() if isinstance(out, DTensor) else out

        fa1_out = _fa(q_keep, split_id=2 * local_idx)
        fa2_out = _fa(q_peer, split_id=2 * target_idx + 1)
        fa2_our = utils.p2p_exchange(fa2_out, peer_rank)
        out = utils.cat([fa1_out, fa2_our], dim=seq_dim)
        return _finalize_colossal_output(out, output_layout, co_submesh, seq_dim, self.use_local_output)
hyper_parallel/core/context_parallel/dsa_context_parallel.py
 92
 93
 94
 95
 96
 97
 98
 99
100
                "DSA shared replicate backward requires a divisible sequence dimension, "
                f"got {output_shape[0]} and CP size {ctx.world_size}."
            )
        output_shape[0] //= ctx.world_size
        local_grad, work = utils.reduce_scatter_single(
            grad_front, output_shape, ctx.group, async_op=False
        )
        if work is not None:
            work.wait()
702
703
704
705
706
707
708
709
710
            return _finalize_output(value, use_local_output=True)

        if isinstance(value, DTensor):
            value = _dtensor_to_local_reducing_partial(value)
        if not utils.is_tensor(value):
            return value

        target_len = target_shape[self.seq_dim]
        if value.shape[self.seq_dim] == target_len:
hyper_parallel/core/context_parallel/utils.py
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70


def _normalize_dim(dim: int, ndim: int) -> int:
    """Normalize a possibly negative dimension index."""
    return dim + ndim if dim < 0 else dim


def _move_dim_to_front(tensor: Tensor, dim: int) -> Tensor:
    """Move ``dim`` to the front while keeping the other dimensions ordered."""
    dim = _normalize_dim(dim, tensor.dim())
    if dim == 0:
        return tensor.contiguous()
    perm = [dim] + [index for index in range(tensor.dim()) if index != dim]
    return tensor.permute(perm).contiguous()


def _move_dim_from_front(tensor: Tensor, dim: int) -> Tensor:
    """Move the leading dimension back to ``dim``."""
    dim = _normalize_dim(dim, tensor.dim())
    if dim == 0:
        return tensor.contiguous()
    perm = [dim] + [index for index in range(tensor.dim()) if index != dim]
    inverse = [0] * len(perm)
    for index, value in enumerate(perm):
        inverse[value] = index
    return tensor.permute(inverse).contiguous()


def _a2a_reconstruct(out_perm: Tensor, concat_dim: int) -> Tensor:
    """Reconstruct an all-to-all output from its leading-rank layout."""
    chunk_in_perm = concat_dim + 1
    recon_perm = list(range(1, chunk_in_perm)) + [0] + list(range(chunk_in_perm, out_perm.dim()))
    reconstructed = out_perm.permute(recon_perm).contiguous()
    shape = list(reconstructed.shape)
    merged = shape[concat_dim] * shape[concat_dim + 1]
    return reconstructed.reshape(shape[:concat_dim] + [merged] + shape[concat_dim + 2:])


class _AsyncA2AWait(torch.autograd.Function):
    """Wait for a pre-launched all-to-all and preserve its backward overlap."""
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

    @staticmethod
    def forward(ctx, tensor, work, out_perm, group, world_size, concat_dim, split_dim, handle_box):
        """Wait and reconstruct the forward all-to-all output."""
        ctx.group = group
        ctx.world_size = world_size
        ctx.concat_dim = concat_dim
        ctx.split_dim = split_dim
        ctx.handle_box = handle_box
        ctx.input_shape = tensor.shape
        work.wait()
        return _a2a_reconstruct(out_perm, concat_dim)

    @staticmethod
    def backward(ctx, grad_output):
        """Launch the reverse all-to-all when overlap was requested."""
        if ctx.handle_box is not None:
            grad_output = grad_output.contiguous()
            shape = list(grad_output.shape)
            seq_dim = ctx.concat_dim
            full_size = shape[seq_dim]
            ndim = len(shape) + 1
            grad_perm = grad_output.reshape(
                shape[:seq_dim]
                + [ctx.world_size, full_size // ctx.world_size]
                + shape[seq_dim + 1:]
            ).permute(
 96
 97
 98
 99
100
101
102
103
104
105
106
107
                + shape[seq_dim + 1:]
            ).permute(
                [seq_dim] + list(range(seq_dim)) + list(range(seq_dim + 1, ndim))
            ).contiguous()
            out_perm = torch.empty_like(grad_perm)
            work = dist.all_to_all_single(out_perm, grad_perm, group=ctx.group, async_op=True)
            ctx.handle_box.append((work, out_perm))
        return grad_output.new_zeros(ctx.input_shape), None, None, None, None, None, None, None


class _AsyncAllGatherWait(torch.autograd.Function):
    """Wait for a pre-launched all-gather and provide reduce-scatter backward."""
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

    @staticmethod
    def forward(ctx, tensor, work, out_perm, group, world_size, gather_dim, handle_box):
        """Wait and reconstruct the gathered tensor."""
        ctx.group = group
        ctx.world_size = world_size
        ctx.gather_dim = gather_dim
        ctx.handle_box = handle_box
        ctx.input_shape = tensor.shape
        work.wait()
        return _move_dim_from_front(out_perm, gather_dim)

    @staticmethod
    def backward(ctx, grad_output):
        """Reduce-scatter the gathered gradient."""
        grad_perm = _move_dim_to_front(grad_output.contiguous(), ctx.gather_dim)
        output_shape = list(grad_perm.shape)
        if output_shape[0] % ctx.world_size != 0:
            raise ValueError(
                "all_gather backward expected gathered dimension to be divisible by world_size, "
                f"got {output_shape[0]} and {ctx.world_size}."
            )
        output_shape[0] //= ctx.world_size
        output = torch.empty(output_shape, dtype=grad_perm.dtype, device=grad_perm.device)
        work = dist.reduce_scatter_tensor(output, grad_perm, group=ctx.group, async_op=True)
        if ctx.handle_box is not None:
            ctx.handle_box.append((work, output, ctx.gather_dim))
            return grad_output.new_zeros(ctx.input_shape), None, None, None, None, None, None
        work.wait()
        return _move_dim_from_front(output, ctx.gather_dim), None, None, None, None, None, None


class _P2PExchange(torch.autograd.Function):
    """Symmetrically exchange a tensor with one peer in forward and backward."""
142
143
144
145
146
147
148
149
150
151
152
153
154

    @staticmethod
    def forward(ctx, tensor: Tensor, peer_rank: int, group):
        """Exchange the forward tensor."""
        ctx.peer_rank = peer_rank
        ctx.group = group
        send_buffer = tensor.contiguous()
        receive_buffer = torch.empty_like(send_buffer)
        requests = dist.batch_isend_irecv(
            [
                dist.P2POp(dist.isend, send_buffer, peer_rank, group),
                dist.P2POp(dist.irecv, receive_buffer, peer_rank, group),
            ]
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
                dist.P2POp(dist.isend, send_buffer, peer_rank, group),
                dist.P2POp(dist.irecv, receive_buffer, peer_rank, group),
            ]
        )
        for request in requests:
            request.wait()
        return receive_buffer

    @staticmethod
    def backward(ctx, grad_output: Tensor):
        """Exchange the backward gradient."""
        send_buffer = grad_output.contiguous()
        receive_buffer = torch.empty_like(send_buffer)
        requests = dist.batch_isend_irecv(
            [
                dist.P2POp(dist.isend, send_buffer, ctx.peer_rank, ctx.group),
                dist.P2POp(dist.irecv, receive_buffer, ctx.peer_rank, ctx.group),
            ]
167
168
169
170
171
172
173
174
175
176
177
                dist.P2POp(dist.isend, send_buffer, ctx.peer_rank, ctx.group),
                dist.P2POp(dist.irecv, receive_buffer, ctx.peer_rank, ctx.group),
            ]
        )
        for request in requests:
            request.wait()
        return receive_buffer, None, None


def is_tensor(value) -> bool:
    """Return whether ``value`` is a PyTorch tensor."""
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


def all_to_all_single(input_tensor: Tensor, output_shape: Sequence[int], group, async_op: bool = False):
    """Run a fixed-shape all-to-all collective."""
    output = torch.empty(output_shape, device=input_tensor.device, dtype=input_tensor.dtype)
    work = dist.all_to_all_single(output, input_tensor, group=group, async_op=async_op)
    return output, work


def all_gather_single(input_tensor: Tensor, output_shape: Sequence[int], group, async_op: bool = False):
    """Run an all-gather-into-tensor collective."""
    output = torch.empty(output_shape, device=input_tensor.device, dtype=input_tensor.dtype)
    work = dist.all_gather_into_tensor(output, input_tensor, group=group, async_op=async_op)
    return output, work


def reduce_scatter_single(input_tensor: Tensor, output_shape: Sequence[int], group, async_op: bool = False):
    """Run a reduce-scatter-into-tensor collective."""
    output = torch.empty(output_shape, device=input_tensor.device, dtype=input_tensor.dtype)
    work = dist.reduce_scatter_tensor(output, input_tensor, group=group, async_op=async_op)
    return output, work


def differentiable_async_a2a_wait(
    tensor: Tensor, work, out_perm: Tensor, group, world_size: int, concat_dim: int, split_dim: int, handle_box=None
217
218
219
220
221
222
223
224
225
def differentiable_async_a2a_wait(
    tensor: Tensor, work, out_perm: Tensor, group, world_size: int, concat_dim: int, split_dim: int, handle_box=None
) -> Tensor:
    """Wait for an asynchronous all-to-all while preserving autograd."""
    return _AsyncA2AWait.apply(tensor, work, out_perm, group, world_size, concat_dim, split_dim, handle_box)


def differentiable_async_allgather_wait(
    tensor: Tensor, work, out_perm: Tensor, group, world_size: int, gather_dim: int, handle_box=None
224
225
226
227
228
229
230
231
232
def differentiable_async_allgather_wait(
    tensor: Tensor, work, out_perm: Tensor, group, world_size: int, gather_dim: int, handle_box=None
) -> Tensor:
    """Wait for an asynchronous all-gather while preserving autograd."""
    return _AsyncAllGatherWait.apply(tensor, work, out_perm, group, world_size, gather_dim, handle_box)


def differentiable_all_to_all_single(
    input_tensor: Tensor, input_splits: Sequence[int], output_splits: Sequence[int], group
231
232
233
234
235
236
237
238
239
240
241
242
def differentiable_all_to_all_single(
    input_tensor: Tensor, input_splits: Sequence[int], output_splits: Sequence[int], group
) -> Tensor:
    """Run a differentiable variable-split all-to-all collective."""
    output = torch.empty(
        sum(output_splits), *input_tensor.shape[1:], dtype=input_tensor.dtype, device=input_tensor.device
    )
    return dist_func.all_to_all_single(
        output,
        input_tensor,
        output_split_sizes=list(output_splits),
        input_split_sizes=list(input_splits),
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
def differentiable_all_gather_concat(
    tensor: Tensor, group, concat_size: int, concat_dim: int, rank_list=None
) -> Tensor:
    """Differentiably gather and concatenate tensors in mesh rank order."""
    del concat_size
    tensor = tensor.contiguous()
    output = list(dist_func.all_gather(tensor, group=group))
    if rank_list is not None:
        group_ranks = dist.get_process_group_ranks(group)
        if tuple(rank_list) != tuple(group_ranks):
            rank_to_index = {int(rank): index for index, rank in enumerate(group_ranks)}
            output = [output[rank_to_index[int(rank)]] for rank in rank_list]
    return torch.cat(output, dim=concat_dim)


def p2p_exchange(tensor: Tensor, peer_rank: int, group=None) -> Tensor:
    """Symmetrically exchange a tensor with ``peer_rank``."""
    if peer_rank == dist.get_rank(group):
        return tensor
    return _P2PExchange.apply(tensor, peer_rank, group)


def cat(tensors, dim: int = 0) -> Tensor:
    """Concatenate tensors along ``dim``."""
    return torch.cat(tensors, dim=dim)
hyper_parallel/distributed/context_parallel/gated_delta_net.py
590
591
592
593
594
595
596
597
598
599
600
601
602
        prev_rank: int,
        next_rank: int,
    ) -> torch.Tensor:
        """Run fused local GDN and forward its affine state across CP ranks."""
        from hyper_parallel.components.functional.gated_delta_net import (  # pylint: disable=import-outside-toplevel
            chunk_gated_delta_rule_fwd_apply_state_saved,
            chunk_gated_delta_rule_fwd_output_saved,
            chunk_gated_delta_rule_fwd_prepare_saved,
        )
        from hyper_parallel.components.functional.gated_delta_net_state_summary import (  # pylint: disable=import-outside-toplevel
            apply_gdn_state_summary,
            chunk_gated_delta_rule_state_summary_fwd,
        )
696
697
698
699
700
701
702
703
704
705
706
707
708

    @staticmethod
    def backward(ctx, grad_output: torch.Tensor):  # pylint: disable=too-many-locals
        """Backpropagate local GDN tensors and the state gradient wavefront."""
        from hyper_parallel.components.functional.gated_delta_net import (  # pylint: disable=import-outside-toplevel
            chunk_gated_delta_rule_bwd_finish_saved,
            chunk_gated_delta_rule_bwd_prepare_saved,
            chunk_gated_delta_rule_bwd_state_saved,
        )
        from hyper_parallel.components.functional.gated_delta_net_state_summary import (  # pylint: disable=import-outside-toplevel
            apply_gdn_state_gradient_summary,
            chunk_gated_delta_rule_state_gradient_summary_bwd,
        )
hyper_parallel/distributed/context_parallel/kimi_delta_attention.py
747
748
749
750
751
752
753
754
755
                cp_rank=self.cp_rank,
                cp_size=self.cp_size,
            )

        from hyper_parallel.components.functional.kimi_delta_attention import (  # pylint: disable=import-outside-toplevel
            fused_chunk_kda_p2p,
        )
        return fused_chunk_kda_p2p(
            query,