"""Chunk-parallel Kimi Delta Attention forward for SM120 (RTX PRO 6000 Blackwell). Written from the math in reference.py as three Triton kernels: 1. _prep_wu_kernel -- per (chunk, batch*head): builds the strictly-lower intra-chunk matrix L[c,i] = -beta_c * , forms the unit lower-triangular system N = I - L, and solves N W = diag(beta) E by 16x16-blocked forward substitution, emitting w = M @ (e^g k), u = M @ v, M = N^{-1} diag(beta). The inverse of each 16x16 diagonal block is computed with an exact row-recurrence; block combinations are tensor-core dots. The chunk recurrence is affine in the state, so it is solved at two levels: chunks are grouped into segments of W; each segment's net transition S_out = A @ S_in + b is built exactly (A as a dense K x K matrix), a short outer scan walks the segments, and per-chunk outputs are then finished in a fully parallel pass corrected by the checkpointed state at each segment start. Sequential depth drops from NT chunk steps to W + NT/W. 1. _prep_wu_kernel -- per (chunk, batch*head): builds the strictly-lower intra-chunk matrix L[c,i] = -beta_c * , forms the unit lower-triangular system N = I - L, and solves N W = diag(beta) E by 16x16-blocked forward substitution, emitting w = M @ (e^g k), u = M @ v, M = N^{-1} diag(beta). Also emits decay-normalized k/q (kg, qg) and per-chunk decay stats. 2. _seg_a_kernel -- per (segment, batch*head): builds A = diag(e^{Lam}) - sum_n kgR_n^T (w_n D_n) and per-chunk prefix/suffix decay factors. 3. _seg_b_kernel -- per (segment, batch*head, V-tile): runs the local chunk chain from a zero seed, emitting output partials o_loc and b (V-slice). 4. _outer_scan_kernel -- per (batch*head, V-tile): scans segment transitions, leaving checkpoint states at every segment start. 5. _corr_out_kernel -- per (chunk, batch*head, V-tile), fully parallel: o = o_loc + (q e^g D_n) @ cp - tril((scale q e^g)(k e^-g)^T) @ ((w D_n) @ cp) """ from __future__ import annotations import torch import torch.nn as nn import triton import triton.language as tl OP_TYPE = "linear_attention" SUPPORTED_PRECISIONS = ["bf16"] HARDWARE_REQUIRED = ["RTX_PRO_6000", "H100", "B200"] @triton.jit def _inv_unit_lower16(z): """Return y with (I - z) @ y = I, for strictly-lower 16x16 fp32 z.""" o = tl.arange(0, 16) y = (o[:, None] == o[None, :]).to(tl.float32) for i in tl.static_range(1, 16): sel = o == i zi = tl.where(sel[:, None] & (o[None, :] < i), z, 0.0) zv = tl.sum(zi, axis=0) row = sel.to(tl.float32) + tl.sum(zv[:, None] * y, axis=0) y = tl.where(sel[:, None], row[None, :], y) return y @triton.jit def _prep_wu_kernel( q_ptr, k_ptr, v_ptr, g_ptr, beta_ptr, w_ptr, u_ptr, kg_ptr, qg_ptr, gl_ptr, eg_ptr, ei_ptr, scale, T, H: tl.constexpr, K: tl.constexpr, V: tl.constexpr, BT: tl.constexpr, BC: tl.constexpr, NSUB: tl.constexpr, ): i_t = tl.program_id(0) i_bh = tl.program_id(1) i_b = i_bh // H i_h = i_bh % H o_c = tl.arange(0, BC) o_k = tl.arange(0, K) o_v = tl.arange(0, V) t0 = i_t * BT # --- load the four token sub-blocks of k / g / beta ----------------------- off0 = (i_b * T + t0 + 0 * BC + o_c) * H + i_h off1 = (i_b * T + t0 + 1 * BC + o_c) * H + i_h off2 = (i_b * T + t0 + 2 * BC + o_c) * H + i_h off3 = (i_b * T + t0 + 3 * BC + o_c) * H + i_h bk0 = tl.load(k_ptr + off0[:, None] * K + o_k[None, :]).to(tl.float32) bk1 = tl.load(k_ptr + off1[:, None] * K + o_k[None, :]).to(tl.float32) bk2 = tl.load(k_ptr + off2[:, None] * K + o_k[None, :]).to(tl.float32) bk3 = tl.load(k_ptr + off3[:, None] * K + o_k[None, :]).to(tl.float32) bg0 = tl.load(g_ptr + off0[:, None] * K + o_k[None, :]) bg1 = tl.load(g_ptr + off1[:, None] * K + o_k[None, :]) bg2 = tl.load(g_ptr + off2[:, None] * K + o_k[None, :]) bg3 = tl.load(g_ptr + off3[:, None] * K + o_k[None, :]) bt0 = tl.load(beta_ptr + off0).to(tl.float32) bt1 = tl.load(beta_ptr + off1).to(tl.float32) bt2 = tl.load(beta_ptr + off2).to(tl.float32) bt3 = tl.load(beta_ptr + off3).to(tl.float32) # Per-chunk decay normalizers shared by the state and output passes. # g is a per-token log-decay: form the in-chunk cumulative sum first. # (Computed as triangular matmuls -- tl.cumsum is miscompiled on this # Triton/Blackwell combination when several scans share a kernel.) ltri = (o_c[:, None] >= o_c[None, :]).to(tl.float32) acc = tl.zeros((K,), dtype=tl.float32) ga0 = tl.dot(ltri, bg0, input_precision="ieee") + acc[None, :] acc += tl.sum(bg0, 0) ga1 = tl.dot(ltri, bg1, input_precision="ieee") + acc[None, :] acc += tl.sum(bg1, 0) ga2 = tl.dot(ltri, bg2, input_precision="ieee") + acc[None, :] acc += tl.sum(bg2, 0) ga3 = tl.dot(ltri, bg3, input_precision="ieee") + acc[None, :] gl = acc + tl.sum(bg3, 0) pgl = gl_ptr + (i_bh * (T // BT) + i_t) * K + o_k tl.store(pgl, gl) bq0 = tl.load(q_ptr + off0[:, None] * K + o_k[None, :]).to(tl.float32) bq1 = tl.load(q_ptr + off1[:, None] * K + o_k[None, :]).to(tl.float32) bq2 = tl.load(q_ptr + off2[:, None] * K + o_k[None, :]).to(tl.float32) bq3 = tl.load(q_ptr + off3[:, None] * K + o_k[None, :]).to(tl.float32) kn0 = (bk0 * tl.exp(ga0)).to(tl.bfloat16) kn1 = (bk1 * tl.exp(ga1)).to(tl.bfloat16) kn2 = (bk2 * tl.exp(ga2)).to(tl.bfloat16) kn3 = (bk3 * tl.exp(ga3)).to(tl.bfloat16) kd0 = (bk0 * tl.exp(-ga0)).to(tl.bfloat16) kd1 = (bk1 * tl.exp(-ga1)).to(tl.bfloat16) kd2 = (bk2 * tl.exp(-ga2)).to(tl.bfloat16) kd3 = (bk3 * tl.exp(-ga3)).to(tl.bfloat16) # --- strictly-lower block rows of L (beta-scaled rows) ------------------- l10 = -tl.dot(kn1, tl.trans(kd0)) * bt1[:, None] l20 = -tl.dot(kn2, tl.trans(kd0)) * bt2[:, None] l21 = -tl.dot(kn2, tl.trans(kd1)) * bt2[:, None] l30 = -tl.dot(kn3, tl.trans(kd0)) * bt3[:, None] l31 = -tl.dot(kn3, tl.trans(kd1)) * bt3[:, None] l32 = -tl.dot(kn3, tl.trans(kd2)) * bt3[:, None] # --- inverses of the diagonal units N_ss = I - L_ss ---------------------- m_l = o_c[:, None] > o_c[None, :] y0 = _inv_unit_lower16(tl.where(m_l, -(tl.dot(kn0, tl.trans(kd0)) * bt0[:, None]), 0.0)) y1 = _inv_unit_lower16(tl.where(m_l, -(tl.dot(kn1, tl.trans(kd1)) * bt1[:, None]), 0.0)) y2 = _inv_unit_lower16(tl.where(m_l, -(tl.dot(kn2, tl.trans(kd2)) * bt2[:, None]), 0.0)) y3 = _inv_unit_lower16(tl.where(m_l, -(tl.dot(kn3, tl.trans(kd3)) * bt3[:, None]), 0.0)) # --- forward substitution over row-blocks for w = M @ (e^g k) ------------- wa0 = kn0.to(tl.float32) * bt0[:, None] w0 = tl.dot(y0, wa0, input_precision="tf32") wa1 = kn1.to(tl.float32) * bt1[:, None] + tl.dot(l10, w0, input_precision="tf32") w1 = tl.dot(y1, wa1, input_precision="tf32") wa2 = (kn2.to(tl.float32) * bt2[:, None] + tl.dot(l20, w0, input_precision="tf32") + tl.dot(l21, w1, input_precision="tf32")) w2 = tl.dot(y2, wa2, input_precision="tf32") wa3 = (kn3.to(tl.float32) * bt3[:, None] + tl.dot(l30, w0, input_precision="tf32") + tl.dot(l31, w1, input_precision="tf32") + tl.dot(l32, w2, input_precision="tf32")) w3 = tl.dot(y3, wa3, input_precision="tf32") tl.store(w_ptr + off0[:, None] * K + o_k[None, :], w0.to(tl.bfloat16)) tl.store(w_ptr + off1[:, None] * K + o_k[None, :], w1.to(tl.bfloat16)) tl.store(w_ptr + off2[:, None] * K + o_k[None, :], w2.to(tl.bfloat16)) tl.store(w_ptr + off3[:, None] * K + o_k[None, :], w3.to(tl.bfloat16)) # --- forward substitution over row-blocks for u = M @ v ------------------- bv0 = tl.load(v_ptr + off0[:, None] * V + o_v[None, :]).to(tl.float32) bv1 = tl.load(v_ptr + off1[:, None] * V + o_v[None, :]).to(tl.float32) bv2 = tl.load(v_ptr + off2[:, None] * V + o_v[None, :]).to(tl.float32) bv3 = tl.load(v_ptr + off3[:, None] * V + o_v[None, :]).to(tl.float32) ua0 = bv0 * bt0[:, None] u0 = tl.dot(y0, ua0, input_precision="tf32") ua1 = bv1 * bt1[:, None] + tl.dot(l10, u0, input_precision="tf32") u1 = tl.dot(y1, ua1, input_precision="tf32") ua2 = (bv2 * bt2[:, None] + tl.dot(l20, u0, input_precision="tf32") + tl.dot(l21, u1, input_precision="tf32")) u2 = tl.dot(y2, ua2, input_precision="tf32") ua3 = (bv3 * bt3[:, None] + tl.dot(l30, u0, input_precision="tf32") + tl.dot(l31, u1, input_precision="tf32") + tl.dot(l32, u2, input_precision="tf32")) u3 = tl.dot(y3, ua3, input_precision="tf32") tl.store(u_ptr + off0[:, None] * V + o_v[None, :], u0.to(tl.bfloat16)) tl.store(u_ptr + off1[:, None] * V + o_v[None, :], u1.to(tl.bfloat16)) tl.store(u_ptr + off2[:, None] * V + o_v[None, :], u2.to(tl.bfloat16)) tl.store(u_ptr + off3[:, None] * V + o_v[None, :], u3.to(tl.bfloat16)) # Decay-normalized k (k e^{g_last - g}) and decayed queries for the # state / output passes; both consume these instead of fp32 g. egl = tl.exp(gl) tl.store(eg_ptr + (i_bh * (T // BT) + i_t) * K + o_k, egl) tl.store(ei_ptr + (i_bh * (T // BT) + i_t) * K + o_k, tl.exp(-gl)) kg0 = (bk0 * egl[None, :] * tl.exp(-ga0)).to(tl.bfloat16) kg1 = (bk1 * egl[None, :] * tl.exp(-ga1)).to(tl.bfloat16) kg2 = (bk2 * egl[None, :] * tl.exp(-ga2)).to(tl.bfloat16) kg3 = (bk3 * egl[None, :] * tl.exp(-ga3)).to(tl.bfloat16) tl.store(kg_ptr + off0[:, None] * K + o_k[None, :], kg0) tl.store(kg_ptr + off1[:, None] * K + o_k[None, :], kg1) tl.store(kg_ptr + off2[:, None] * K + o_k[None, :], kg2) tl.store(kg_ptr + off3[:, None] * K + o_k[None, :], kg3) qg0 = (bq0 * (scale * tl.exp(ga0))).to(tl.bfloat16) qg1 = (bq1 * (scale * tl.exp(ga1))).to(tl.bfloat16) qg2 = (bq2 * (scale * tl.exp(ga2))).to(tl.bfloat16) qg3 = (bq3 * (scale * tl.exp(ga3))).to(tl.bfloat16) tl.store(qg_ptr + off0[:, None] * K + o_k[None, :], qg0) tl.store(qg_ptr + off1[:, None] * K + o_k[None, :], qg1) tl.store(qg_ptr + off2[:, None] * K + o_k[None, :], qg2) tl.store(qg_ptr + off3[:, None] * K + o_k[None, :], qg3) @triton.jit def _seg_a_kernel( w_ptr, kg_ptr, gl_ptr, a_ptr, dn_ptr, sf_ptr, T, H: tl.constexpr, K: tl.constexpr, BT: tl.constexpr, W: tl.constexpr, ): """Per (segment, batch*head): build the segment's affine state transition A = diag(e^{Lam}) - sum_n kgR_n^T (w_n D_n) as an exact K x K matrix, and emit per-chunk decay prefixes dn[n] and suffix factors sf[n].""" i_s = tl.program_id(0) i_bh = tl.program_id(1) NT = T // BT o_k = tl.arange(0, K) eye = (o_k[:, None] == o_k[None, :]).to(tl.float32) base = i_bh * NT lam_end = tl.zeros((K,), dtype=tl.float32) for r in range(0, W): lam_end += tl.load(gl_ptr + (base + i_s * W + r) * K + o_k) lam = tl.zeros((K,), dtype=tl.float32) A = eye * tl.exp(lam_end)[None, :] for r in range(0, W): n = i_s * W + r rows = n * BT + tl.arange(0, BT) off = ((i_bh // H) * T + rows) * H + (i_bh % H) gl = tl.load(gl_ptr + (base + n) * K + o_k) tl.store(dn_ptr + (base + n) * K + o_k, tl.exp(lam)) sfac = tl.exp(lam_end - lam - gl) tl.store(sf_ptr + (base + n) * K + o_k, sfac) b_w = tl.load(w_ptr + off[:, None] * K + o_k[None, :]).to(tl.float32) b_kg = tl.load(kg_ptr + off[:, None] * K + o_k[None, :]).to(tl.float32) wd = b_w * tl.exp(lam)[None, :] kgr = b_kg * sfac[None, :] A -= tl.dot(tl.trans(kgr), wd, input_precision="tf32") lam += gl tl.store(a_ptr + (i_bh * (T // (BT * W)) + i_s) * K * K + o_k[:, None] * K + o_k[None, :], A.to(tl.float16)) @triton.jit def _seg_b_kernel( w_ptr, u_ptr, kg_ptr, qg_ptr, eg_ptr, ei_ptr, gl_ptr, oloc_ptr, b_ptr, T, H: tl.constexpr, K: tl.constexpr, V: tl.constexpr, BT: tl.constexpr, BV: tl.constexpr, W: tl.constexpr, ): """Per (segment, batch*head, V-tile): run the chunk chain from a zero seed. Emits the local output partials o_loc and the segment offset term b = sum_n kgR_n^T u~_n (this V-slice).""" i_s = tl.program_id(0) i_bh = tl.program_id(1) i_v = tl.program_id(2) i_b = i_bh // H i_h = i_bh % H NT = T // BT o_c = tl.arange(0, BT) o_k = tl.arange(0, K) o_v = i_v * BV + tl.arange(0, BV) base = i_bh * NT lam_end = tl.zeros((K,), dtype=tl.float32) for r in range(0, W): lam_end += tl.load(gl_ptr + (base + i_s * W + r) * K + o_k) S = tl.zeros((K, BV), dtype=tl.float32) bacc = tl.zeros((K, BV), dtype=tl.float32) lam = tl.zeros((K,), dtype=tl.float32) for r in range(0, W): n = i_s * W + r rows = n * BT + o_c off = (i_b * T + rows) * H + i_h b_w = tl.load(w_ptr + off[:, None] * K + o_k[None, :]) b_u = tl.load(u_ptr + off[:, None] * V + o_v[None, :]).to(tl.float32) b_kg = tl.load(kg_ptr + off[:, None] * K + o_k[None, :]) b_qg = tl.load(qg_ptr + off[:, None] * K + o_k[None, :]) egn = tl.load(eg_ptr + (base + n) * K + o_k) ein = tl.load(ei_ptr + (base + n) * K + o_k) gl = tl.load(gl_ptr + (base + n) * K + o_k) # segment-local corrected values and output partials vp = (b_u - tl.dot(b_w, S.to(tl.bfloat16))).to(tl.bfloat16) kd = (b_kg.to(tl.float32) * ein[None, :]).to(tl.bfloat16) m_l = o_c[:, None] >= o_c[None, :] aq = tl.dot(b_qg, tl.trans(kd)) aq = tl.where(m_l, aq, 0.0).to(tl.bfloat16) ol = tl.dot(b_qg, S.to(tl.bfloat16)) + tl.dot(aq, vp) tl.store(oloc_ptr + off[:, None] * V + o_v[None, :], ol.to(tl.bfloat16)) sfac = tl.exp(lam_end - lam - gl) bacc += tl.dot(tl.trans((b_kg.to(tl.float32) * sfac[None, :]).to(tl.bfloat16)), vp) S = S * egn[:, None] + tl.dot(tl.trans(b_kg), vp) lam += gl tl.store(b_ptr + (i_bh * (NT // W) + i_s) * K * V + o_k[:, None] * V + o_v[None, :], bacc.to(tl.float16)) @triton.jit def _outer_scan_kernel( a_ptr, b_ptr, cp_ptr, H: tl.constexpr, K: tl.constexpr, V: tl.constexpr, BV: tl.constexpr, NSEG: tl.constexpr, ): """Per (batch*head, V-tile): scan the segment transitions, leaving the state at the START of every segment in cp.""" i_bh = tl.program_id(0) i_v = tl.program_id(1) o_k = tl.arange(0, K) o_v = i_v * BV + tl.arange(0, BV) S = tl.zeros((K, BV), dtype=tl.float32) for s in range(0, NSEG): tl.store(cp_ptr + (i_bh * NSEG + s) * K * V + o_k[:, None] * V + o_v[None, :], S.to(tl.float16)) a = tl.load(a_ptr + (i_bh * NSEG + s) * K * K + o_k[:, None] * K + o_k[None, :]) bb = tl.load(b_ptr + (i_bh * NSEG + s) * K * V + o_k[:, None] * V + o_v[None, :]).to(tl.float32) S = tl.dot(a, S.to(tl.float16)) + bb @triton.jit def _corr_out_kernel( w_ptr, kg_ptr, qg_ptr, ei_ptr, dn_ptr, cp_ptr, oloc_ptr, o_ptr, T, H: tl.constexpr, K: tl.constexpr, V: tl.constexpr, BT: tl.constexpr, BV: tl.constexpr, W: tl.constexpr, ): """Per (chunk, batch*head, V-tile), fully parallel: add the checkpointed state's contribution to the local partials and store the final output.""" i_n = tl.program_id(0) i_bh = tl.program_id(1) i_v = tl.program_id(2) i_b = i_bh // H i_h = i_bh % H NT = T // BT NSEG = NT // W o_c = tl.arange(0, BT) o_k = tl.arange(0, K) o_v = i_v * BV + tl.arange(0, BV) rows = i_n * BT + o_c off = (i_b * T + rows) * H + i_h b_w = tl.load(w_ptr + off[:, None] * K + o_k[None, :]) b_qg = tl.load(qg_ptr + off[:, None] * K + o_k[None, :]) b_kg = tl.load(kg_ptr + off[:, None] * K + o_k[None, :]) ein = tl.load(ei_ptr + (i_bh * NT + i_n) * K + o_k) dn = tl.load(dn_ptr + (i_bh * NT + i_n) * K + o_k) kd = (b_kg.to(tl.float32) * ein[None, :]).to(tl.float16) cq = (b_qg.to(tl.float32) * dn[None, :]).to(tl.float16) cw = (b_w.to(tl.float32) * dn[None, :]).to(tl.float16) cp = tl.load(cp_ptr + (i_bh * NSEG + i_n // W) * K * V + o_k[:, None] * V + o_v[None, :]) ol = tl.load(oloc_ptr + off[:, None] * V + o_v[None, :]).to(tl.float32) m_l = o_c[:, None] >= o_c[None, :] aq = tl.dot(b_qg.to(tl.float16), tl.trans(kd)) aq = tl.where(m_l, aq, 0.0).to(tl.float16) t_inter = tl.dot(cq, cp) t_vp = tl.dot(cw, cp).to(tl.float16) acc = ol + t_inter - tl.dot(aq, t_vp) tl.store(o_ptr + off[:, None] * V + o_v[None, :], acc.to(tl.bfloat16)) class Model(nn.Module): """KDA forward (chunk form). No learned parameters; all inputs are activations.""" def __init__(self, B: int, T: int, H: int, K: int, V: int, chunk_size: int = 64): super().__init__() assert chunk_size == 64, "this kernel is specialized for chunk_size=64" self.B, self.T, self.H, self.K, self.V = B, T, H, K, V self.chunk_size = chunk_size self.scale = float(K) ** -0.5 self._ws = None # lazily-built workspace cache, keyed by (B,T,H,K,V,device) self.register_buffer("_dummy", torch.zeros(1), persistent=False) def _workspace(self, B, T, H, K, V, dev): key = (B, T, H, K, V, dev) ws = self._ws if ws is not None and ws[0] == key: return ws[1] NT = T // self.chunk_size W = self._pick_w(NT) NSEG = NT // W ws = ( torch.empty(B * T * H * K, device=dev, dtype=torch.bfloat16), # w torch.empty(B * T * H * V, device=dev, dtype=torch.bfloat16), # u torch.empty(B * T * H * K, device=dev, dtype=torch.bfloat16), # kg torch.empty(B * T * H * K, device=dev, dtype=torch.bfloat16), # qg torch.empty(B * H * NT * K, device=dev, dtype=torch.float32), # gl torch.empty(B * H * NT * K, device=dev, dtype=torch.float32), # eg torch.empty(B * H * NT * K, device=dev, dtype=torch.float32), # ei torch.empty(B * T * H * V, device=dev, dtype=torch.bfloat16), # oloc torch.empty(B * H * NSEG * K * K, device=dev, dtype=torch.float16), # A torch.empty(B * H * NSEG * K * V, device=dev, dtype=torch.float16), # b torch.empty(B * H * NSEG * K * V, device=dev, dtype=torch.float16), # cp torch.empty(B * H * NT * K, device=dev, dtype=torch.float32), # dn torch.empty(B * H * NT * K, device=dev, dtype=torch.float32), # sf ) self._ws = (key, ws) return ws def _pick_w(self, NT): best, best_cost = 1, NT + 1 for w in (2, 4, 8, 16): if NT % w == 0: cost = w + NT // w if cost < best_cost or (cost == best_cost and w < best): best, best_cost = w, cost return best def forward( self, q: torch.Tensor, k: torch.Tensor, v: torch.Tensor, g: torch.Tensor, beta: torch.Tensor, ) -> torch.Tensor: B, T, H, K = q.shape V = v.shape[-1] BT = self.chunk_size NT = T // BT dev = q.device if not q.is_contiguous(): q = q.contiguous() if not k.is_contiguous(): k = k.contiguous() if not v.is_contiguous(): v = v.contiguous() if not g.is_contiguous(): g = g.contiguous() if not beta.is_contiguous(): beta = beta.contiguous() w, u, kg, qg, gl, eg, ei, oloc, mat_a, vec_b, cp, dn, sf = \ self._workspace(B, T, H, K, V, dev) o = torch.empty_like(v) W = self._pick_w(NT) NSEG = NT // W BC = 16 _prep_wu_kernel[(NT, B * H)]( q, k, v, g, beta, w, u, kg, qg, gl, eg, ei, self.scale, T, H=H, K=K, V=V, BT=BT, BC=BC, NSUB=BT // BC, num_warps=4, num_stages=1, ) _seg_a_kernel[(NSEG, B * H)]( w, kg, gl, mat_a, dn, sf, T, H=H, K=K, BT=BT, W=W, num_warps=4, num_stages=1, ) BV = 16 _seg_b_kernel[(NSEG, B * H, V // BV)]( w, u, kg, qg, eg, ei, gl, oloc, vec_b, T, H=H, K=K, V=V, BT=BT, BV=BV, W=W, num_warps=4, num_stages=1, ) _outer_scan_kernel[(B * H, V // BV)]( mat_a, vec_b, cp, H=H, K=K, V=V, BV=BV, NSEG=NSEG, num_warps=4, num_stages=1, ) _corr_out_kernel[(NT, B * H, V // BV)]( w, kg, qg, ei, dn, cp, oloc, o, T, H=H, K=K, V=V, BT=BT, BV=BV, W=W, num_warps=4, num_stages=1, ) return o