# Copyright Amazon.com, Inc. or its affiliates. All Rights Reserved. # SPDX-License-Identifier: Apache-2.0 """NKI per-chunk KDA prefill kernel -- backward-ready variant. Extends `kda_chunk_step` with two additional HBM outputs that the training backward needs: - `N_out (128, 128)` -- the Neumann matrix `P_acc = (I - A)^{-1}` (lower triangular + I on diagonal). - `w_out (128, 128)` -- `k_cumdecay = P_acc @ (k * beta * exp(gc))`. The compute is otherwise identical to `kda_chunk_step` (same scalar-mean intra-chunk approximation and per-K state decay). See `kda_fused_chunked_fwd` for the fused multi-chunk single-launch version used by training. Input contract (same as `kda_chunk_step`): raw ``query``/``key`` (L2-normed by the caller, ``query`` scaled by ``1/sqrt(dk)``); all decay scaling is computed inside the kernel from ``g``. Requires NKI >= 0.4.0. """ import nki import nki.isa as nisa import nki.language as nl P_MAX = 128 @nki.jit def kda_chunk_step_v2( query, # (128, 128) float32 -- RAW: L2-normed q, scaled by 1/sqrt(dk) key, # (128, 128) float32 -- RAW: L2-normed k value, # (128, 128) float32 -- one chunk beta_broadcast, # (128, 128) float32 -- write gate broadcast across dim g_cumsum, # (128, 128) float32 -- per-dim cumsum of g within chunk g_last, # (128, 128) float32 -- g_cumsum[-1] per dim, broadcast to 128x128 state_in, # (128, 128) float32 -- recurrent state from previous chunk ): """Process one chunk of KDA. v2 variant: also emits N and w for backward. See module docstring for the wrapper contract and algorithm reference. Returns: output (128, 128) float32 -- per-token chunk output state_out (128, 128) float32 -- state after this chunk N_out (128, 128) float32 -- Neumann matrix P_acc = (I - A)^{-1} lower-tri w_out (128, 128) float32 -- k_cumdecay = P_acc @ (k*beta*exp(gc)) """ C, dim = query.shape # C = 128, dim = 128 output = nl.ndarray((P_MAX, dim), dtype=query.dtype, buffer=nl.shared_hbm) state_out = nl.ndarray((P_MAX, dim), dtype=nl.float32, buffer=nl.shared_hbm) # v2 save outputs N_out = nl.ndarray((P_MAX, P_MAX), dtype=nl.float32, buffer=nl.shared_hbm) w_out = nl.ndarray((P_MAX, dim), dtype=nl.float32, buffer=nl.shared_hbm) # Load all inputs into SBUF q_raw = nl.ndarray((P_MAX, dim), dtype=query.dtype, buffer=nl.sbuf) nisa.dma_copy(dst=q_raw, src=query) k_raw = nl.ndarray((P_MAX, dim), dtype=key.dtype, buffer=nl.sbuf) nisa.dma_copy(dst=k_raw, src=key) v_c = nl.ndarray((P_MAX, dim), dtype=value.dtype, buffer=nl.sbuf) nisa.dma_copy(dst=v_c, src=value) beta_c = nl.ndarray((P_MAX, dim), dtype=beta_broadcast.dtype, buffer=nl.sbuf) nisa.dma_copy(dst=beta_c, src=beta_broadcast) gc_c = nl.ndarray((P_MAX, dim), dtype=g_cumsum.dtype, buffer=nl.sbuf) nisa.dma_copy(dst=gc_c, src=g_cumsum) gl_c = nl.ndarray((P_MAX, dim), dtype=g_last.dtype, buffer=nl.sbuf) nisa.dma_copy(dst=gl_c, src=g_last) state = nl.ndarray((P_MAX, dim), dtype=nl.float32, buffer=nl.sbuf) nisa.dma_copy(dst=state, src=state_in) # ============================================================ # Generate masks internally to avoid layout-transformation issues. # Uses nisa.iota to generate row-col indices, then relu+sign for step function. # ============================================================ row_minus_col = nl.ndarray((P_MAX, P_MAX), dtype=nl.float32, buffer=nl.sbuf) nisa.iota(dst=row_minus_col, pattern=[[-1, P_MAX]], offset=0, channel_multiplier=1) ones_tile = nl.ndarray((P_MAX, P_MAX), dtype=nl.float32, buffer=nl.sbuf) nisa.iota(dst=ones_tile, pattern=[[0, P_MAX]], offset=1, channel_multiplier=0) eye = nl.ndarray((P_MAX, P_MAX), dtype=nl.float32, buffer=nl.sbuf) nisa.affine_select( dst=eye, pattern=[[-1, P_MAX]], offset=0, channel_multiplier=1, on_true_tile=ones_tile, on_false_value=0.0, ) # Lower triangular with diagonal (Lmask_d) rmc_shifted_d = nl.ndarray((P_MAX, P_MAX), dtype=nl.float32, buffer=nl.sbuf) nisa.tensor_scalar( dst=rmc_shifted_d, data=row_minus_col, op0=nl.add, operand0=0.5, engine=nisa.vector_engine, ) rmc_relu_d = nl.ndarray((P_MAX, P_MAX), dtype=nl.float32, buffer=nl.sbuf) nisa.activation(dst=rmc_relu_d, op=nl.relu, data=rmc_shifted_d, bias=None, scale=1.0) Lmask_d = nl.ndarray((P_MAX, P_MAX), dtype=nl.float32, buffer=nl.sbuf) nisa.activation(dst=Lmask_d, op=nl.sign, data=rmc_relu_d, bias=None, scale=1.0) # Strict lower triangular (Lmask) rmc_shifted_s = nl.ndarray((P_MAX, P_MAX), dtype=nl.float32, buffer=nl.sbuf) nisa.tensor_scalar( dst=rmc_shifted_s, data=row_minus_col, op0=nl.add, operand0=-0.5, engine=nisa.vector_engine, ) rmc_relu_s = nl.ndarray((P_MAX, P_MAX), dtype=nl.float32, buffer=nl.sbuf) nisa.activation(dst=rmc_relu_s, op=nl.relu, data=rmc_shifted_s, bias=None, scale=1.0) Lmask = nl.ndarray((P_MAX, P_MAX), dtype=nl.float32, buffer=nl.sbuf) nisa.activation(dst=Lmask, op=nl.sign, data=rmc_relu_s, bias=None, scale=1.0) # ============================================================ # gc_mean_per_row = gc.mean(axis=-1, keepdim=True) -- (128, 1) # # Compute per-row sum via nisa.tensor_reduce (reduces free dim), then # multiply by 1/dim to get the mean. # ============================================================ gc_row_sum = nl.ndarray((P_MAX, 1), dtype=nl.float32, buffer=nl.sbuf) nisa.tensor_reduce(dst=gc_row_sum, op=nl.add, data=gc_c, axis=(1,)) gc_mean = nl.ndarray((P_MAX, 1), dtype=nl.float32, buffer=nl.sbuf) nisa.tensor_scalar( dst=gc_mean, data=gc_row_sum, op0=nl.multiply, operand0=1.0 / float(dim), engine=nisa.vector_engine, ) # exp(gc_mean) and exp(-gc_mean), shape (128, 1) exp_pos_gc_mean = nl.ndarray((P_MAX, 1), dtype=nl.float32, buffer=nl.sbuf) nisa.activation(dst=exp_pos_gc_mean, op=nl.exp, data=gc_mean, bias=None, scale=1.0) exp_neg_gc_mean = nl.ndarray((P_MAX, 1), dtype=nl.float32, buffer=nl.sbuf) nisa.activation(dst=exp_neg_gc_mean, op=nl.exp, data=gc_mean, bias=None, scale=-1.0) # q_decay = q_raw * exp(gc_mean) per-row -- (128, dim) q_decay = nl.ndarray((P_MAX, dim), dtype=nl.float32, buffer=nl.sbuf) nisa.tensor_scalar( dst=q_decay, data=q_raw, op0=nl.multiply, operand0=exp_pos_gc_mean, engine=nisa.vector_engine, ) # k_decay = k_raw * exp(-gc_mean) per-row -- (128, dim) k_decay = nl.ndarray((P_MAX, dim), dtype=nl.float32, buffer=nl.sbuf) nisa.tensor_scalar( dst=k_decay, data=k_raw, op0=nl.multiply, operand0=exp_neg_gc_mean, engine=nisa.vector_engine, ) # ============================================================ # Two flavors of k_beta: # k_beta_QK = k_raw * exp(gc_mean) * beta for the A / QK construction # k_beta_KCD = k_raw * beta for k_cumdecay = P_acc @ (k_beta_KCD * exp_gc) # ============================================================ # k_beta_KCD first (simpler) k_beta_KCD = nl.ndarray((P_MAX, dim), dtype=nl.float32, buffer=nl.sbuf) nisa.tensor_tensor(dst=k_beta_KCD, data1=k_raw, data2=beta_c, op=nl.multiply) # k_scaled_up = k_raw * exp(+gc_mean) (per-row scalar broadcast) k_scaled_up = nl.ndarray((P_MAX, dim), dtype=nl.float32, buffer=nl.sbuf) nisa.tensor_scalar( dst=k_scaled_up, data=k_raw, op0=nl.multiply, operand0=exp_pos_gc_mean, engine=nisa.vector_engine, ) # k_beta_QK = k_scaled_up * beta k_beta_QK = nl.ndarray((P_MAX, dim), dtype=nl.float32, buffer=nl.sbuf) nisa.tensor_tensor(dst=k_beta_QK, data1=k_scaled_up, data2=beta_c, op=nl.multiply) # v_beta = v * beta v_beta = nl.ndarray((P_MAX, dim), dtype=nl.float32, buffer=nl.sbuf) nisa.tensor_tensor(dst=v_beta, data1=v_c, data2=beta_c, op=nl.multiply) # ============================================================ # Per-dimension exp(gc), exp(g_last - gc), exp(g_last) # ============================================================ exp_gc = nl.ndarray((P_MAX, dim), dtype=nl.float32, buffer=nl.sbuf) nisa.activation(dst=exp_gc, op=nl.exp, data=gc_c, bias=None, scale=1.0) gl_minus_gc = nl.ndarray((P_MAX, dim), dtype=nl.float32, buffer=nl.sbuf) nisa.tensor_tensor(dst=gl_minus_gc, data1=gl_c, data2=gc_c, op=nl.subtract) exp_gl_minus_gc = nl.ndarray((P_MAX, dim), dtype=nl.float32, buffer=nl.sbuf) nisa.activation( dst=exp_gl_minus_gc, op=nl.exp, data=gl_minus_gc, bias=None, scale=1.0 ) exp_gl = nl.ndarray((P_MAX, dim), dtype=nl.float32, buffer=nl.sbuf) nisa.activation(dst=exp_gl, op=nl.exp, data=gl_c, bias=None, scale=1.0) # ============================================================ # QK = k_beta_QK @ k_decay.T --> QK[i, j] = beta[i] * (k[i]·k[j]) * exp(gc_mean_i - gc_mean_j) # # Uses nc_transpose (previously moving=eye pattern). # nc_matmul semantic: nc_matmul(stationary=X, moving=Y) computes X.T @ Y. # step a: kbQK_T = k_beta_QK.T (via nc_transpose(dst, data=k_beta_QK)) # step b: k_decay_T = k_decay.T (via nc_transpose(dst, data=k_decay)) # step c: QK = kbQK_T.T @ k_decay_T (via nc_matmul(stationary=kbQK_T, moving=k_decay_T)) # = k_beta_QK @ k_decay.T # ============================================================ kbQK_T_psum = nl.ndarray((P_MAX, P_MAX), dtype=nl.float32, buffer=nl.psum) nisa.nc_transpose(dst=kbQK_T_psum, data=k_beta_QK) kbQK_T = nl.ndarray((P_MAX, P_MAX), dtype=nl.float32, buffer=nl.sbuf) nisa.tensor_copy(dst=kbQK_T, src=kbQK_T_psum) k_decay_T_psum = nl.ndarray((P_MAX, P_MAX), dtype=nl.float32, buffer=nl.psum) nisa.nc_transpose(dst=k_decay_T_psum, data=k_decay) k_decay_T = nl.ndarray((P_MAX, P_MAX), dtype=nl.float32, buffer=nl.sbuf) nisa.tensor_copy(dst=k_decay_T, src=k_decay_T_psum) QK_psum = nl.ndarray((P_MAX, P_MAX), dtype=nl.float32, buffer=nl.psum) nisa.nc_matmul(dst=QK_psum, stationary=kbQK_T, moving=k_decay_T) QK = nl.ndarray((P_MAX, P_MAX), dtype=nl.float32, buffer=nl.sbuf) nisa.tensor_copy(dst=QK, src=QK_psum) # ============================================================ # QK_decay = QK * lower_mask_diag, A = -QK_decay * strict_lower # ============================================================ QK_decay = nl.ndarray((P_MAX, P_MAX), dtype=nl.float32, buffer=nl.sbuf) nisa.tensor_tensor(dst=QK_decay, data1=QK, data2=Lmask_d, op=nl.multiply) neg_QK_decay = nl.ndarray((P_MAX, P_MAX), dtype=nl.float32, buffer=nl.sbuf) nisa.tensor_scalar( dst=neg_QK_decay, data=QK_decay, op0=nl.multiply, operand0=-1.0, engine=nisa.vector_engine, ) A = nl.ndarray((P_MAX, P_MAX), dtype=nl.float32, buffer=nl.sbuf) nisa.tensor_tensor(dst=A, data1=neg_QK_decay, data2=Lmask, op=nl.multiply) # ============================================================ # Neumann power-doubling: P_acc = (I+A)(I+A^2)(I+A^4)...(I+A^{64}) # After 6 rounds, P_acc = sum_{k=0}^{127} A^k, which is exact for 128x128 # strict lower triangular (nilpotent) A. # ============================================================ P_acc = nl.ndarray((P_MAX, P_MAX), dtype=nl.float32, buffer=nl.sbuf) nisa.tensor_tensor(dst=P_acc, data1=eye, data2=A, op=nl.add) A_pow = nl.ndarray((P_MAX, P_MAX), dtype=nl.float32, buffer=nl.sbuf) nisa.tensor_copy(dst=A_pow, src=A) for _round in nl.sequential_range(6): Ap_T_psum = nl.ndarray((P_MAX, P_MAX), dtype=nl.float32, buffer=nl.psum) nisa.nc_transpose(dst=Ap_T_psum, data=A_pow) Ap_T = nl.ndarray((P_MAX, P_MAX), dtype=nl.float32, buffer=nl.sbuf) nisa.tensor_copy(dst=Ap_T, src=Ap_T_psum) Ap_sq_psum = nl.ndarray((P_MAX, P_MAX), dtype=nl.float32, buffer=nl.psum) nisa.nc_matmul(dst=Ap_sq_psum, stationary=Ap_T, moving=A_pow) nisa.tensor_copy(dst=A_pow, src=Ap_sq_psum) IpA = nl.ndarray((P_MAX, P_MAX), dtype=nl.float32, buffer=nl.sbuf) nisa.tensor_tensor(dst=IpA, data1=eye, data2=A_pow, op=nl.add) IpA_T_psum = nl.ndarray((P_MAX, P_MAX), dtype=nl.float32, buffer=nl.psum) nisa.nc_transpose(dst=IpA_T_psum, data=IpA) IpA_T = nl.ndarray((P_MAX, P_MAX), dtype=nl.float32, buffer=nl.sbuf) nisa.tensor_copy(dst=IpA_T, src=IpA_T_psum) Pacc_psum = nl.ndarray((P_MAX, P_MAX), dtype=nl.float32, buffer=nl.psum) nisa.nc_matmul(dst=Pacc_psum, stationary=IpA_T, moving=P_acc) nisa.tensor_copy(dst=P_acc, src=Pacc_psum) # ============================================================ # v2 save: DMA P_acc to N_out HBM tensor for backward. # ============================================================ nisa.dma_copy(dst=N_out, src=P_acc) # ============================================================ # Apply N: value_corr = P_acc @ v_beta, k_cumdecay = P_acc @ (k_beta_KCD * exp_gc) # ============================================================ N_T_psum = nl.ndarray((P_MAX, P_MAX), dtype=nl.float32, buffer=nl.psum) nisa.nc_transpose(dst=N_T_psum, data=P_acc) N_T = nl.ndarray((P_MAX, P_MAX), dtype=nl.float32, buffer=nl.sbuf) nisa.tensor_copy(dst=N_T, src=N_T_psum) vc_psum = nl.ndarray((P_MAX, dim), dtype=nl.float32, buffer=nl.psum) nisa.nc_matmul(dst=vc_psum, stationary=N_T, moving=v_beta) value_corr = nl.ndarray((P_MAX, dim), dtype=nl.float32, buffer=nl.sbuf) nisa.tensor_copy(dst=value_corr, src=vc_psum) # kb_exp_gc = k_beta_KCD * exp_gc (uses RAW k * beta, NOT the pre-scaled versions) kb_exp_gc = nl.ndarray((P_MAX, dim), dtype=nl.float32, buffer=nl.sbuf) nisa.tensor_tensor(dst=kb_exp_gc, data1=k_beta_KCD, data2=exp_gc, op=nl.multiply) kcd_psum = nl.ndarray((P_MAX, dim), dtype=nl.float32, buffer=nl.psum) nisa.nc_matmul(dst=kcd_psum, stationary=N_T, moving=kb_exp_gc) k_cumdecay = nl.ndarray((P_MAX, dim), dtype=nl.float32, buffer=nl.sbuf) nisa.tensor_copy(dst=k_cumdecay, src=kcd_psum) # ============================================================ # v2 save: DMA k_cumdecay to w_out HBM tensor for backward. # ============================================================ nisa.dma_copy(dst=w_out, src=k_cumdecay) # ============================================================ # attn_intra = (q_decay @ k_decay.T) * lower_mask_diag # # This is where the scalar-mean-decay approximation lives. # ============================================================ q_decay_T_psum = nl.ndarray((P_MAX, P_MAX), dtype=nl.float32, buffer=nl.psum) nisa.nc_transpose(dst=q_decay_T_psum, data=q_decay) q_decay_T = nl.ndarray((P_MAX, P_MAX), dtype=nl.float32, buffer=nl.sbuf) nisa.tensor_copy(dst=q_decay_T, src=q_decay_T_psum) qk_psum = nl.ndarray((P_MAX, P_MAX), dtype=nl.float32, buffer=nl.psum) nisa.nc_matmul(dst=qk_psum, stationary=q_decay_T, moving=k_decay_T) qk_raw = nl.ndarray((P_MAX, P_MAX), dtype=nl.float32, buffer=nl.sbuf) nisa.tensor_copy(dst=qk_raw, src=qk_psum) attn_intra = nl.ndarray((P_MAX, P_MAX), dtype=nl.float32, buffer=nl.sbuf) nisa.tensor_tensor(dst=attn_intra, data1=qk_raw, data2=Lmask_d, op=nl.multiply) # ============================================================ # v_prime = k_cumdecay @ state (inter-chunk contribution) # ============================================================ kcd_T_psum = nl.ndarray((P_MAX, P_MAX), dtype=nl.float32, buffer=nl.psum) nisa.nc_transpose(dst=kcd_T_psum, data=k_cumdecay) kcd_T = nl.ndarray((P_MAX, P_MAX), dtype=nl.float32, buffer=nl.sbuf) nisa.tensor_copy(dst=kcd_T, src=kcd_T_psum) vp_psum = nl.ndarray((P_MAX, dim), dtype=nl.float32, buffer=nl.psum) nisa.nc_matmul(dst=vp_psum, stationary=kcd_T, moving=state) v_prime = nl.ndarray((P_MAX, dim), dtype=nl.float32, buffer=nl.sbuf) nisa.tensor_copy(dst=v_prime, src=vp_psum) v_new = nl.ndarray((P_MAX, dim), dtype=nl.float32, buffer=nl.sbuf) nisa.tensor_tensor(dst=v_new, data1=value_corr, data2=v_prime, op=nl.subtract) # ============================================================ # attn_inter = (q_raw * exp(gc)) @ state <-- RAW q, per-dim exp(gc) # ============================================================ q_exp = nl.ndarray((P_MAX, dim), dtype=nl.float32, buffer=nl.sbuf) nisa.tensor_tensor(dst=q_exp, data1=q_raw, data2=exp_gc, op=nl.multiply) qe_T_psum = nl.ndarray((P_MAX, P_MAX), dtype=nl.float32, buffer=nl.psum) nisa.nc_transpose(dst=qe_T_psum, data=q_exp) qe_T = nl.ndarray((P_MAX, P_MAX), dtype=nl.float32, buffer=nl.sbuf) nisa.tensor_copy(dst=qe_T, src=qe_T_psum) ai_psum = nl.ndarray((P_MAX, dim), dtype=nl.float32, buffer=nl.psum) nisa.nc_matmul(dst=ai_psum, stationary=qe_T, moving=state) attn_inter = nl.ndarray((P_MAX, dim), dtype=nl.float32, buffer=nl.sbuf) nisa.tensor_copy(dst=attn_inter, src=ai_psum) # ============================================================ # intra_out = attn_intra @ v_new # ============================================================ ai_T_psum = nl.ndarray((P_MAX, P_MAX), dtype=nl.float32, buffer=nl.psum) nisa.nc_transpose(dst=ai_T_psum, data=attn_intra) ai_T = nl.ndarray((P_MAX, P_MAX), dtype=nl.float32, buffer=nl.sbuf) nisa.tensor_copy(dst=ai_T, src=ai_T_psum) intra_psum = nl.ndarray((P_MAX, dim), dtype=nl.float32, buffer=nl.psum) nisa.nc_matmul(dst=intra_psum, stationary=ai_T, moving=v_new) intra_out = nl.ndarray((P_MAX, dim), dtype=nl.float32, buffer=nl.sbuf) nisa.tensor_copy(dst=intra_out, src=intra_psum) # ============================================================ # chunk_output = attn_inter + intra_out # ============================================================ chunk_out = nl.ndarray((P_MAX, dim), dtype=nl.float32, buffer=nl.sbuf) nisa.tensor_tensor(dst=chunk_out, data1=attn_inter, data2=intra_out, op=nl.add) nisa.dma_copy(dst=output, src=chunk_out) # ============================================================ # State update: state_new = exp(g_last) * state + k_state_decay^T @ v_new # k_state_decay = k_raw * exp(g_last - gc) <-- RAW k, per-dim exp # ============================================================ k_state_decay = nl.ndarray((P_MAX, dim), dtype=nl.float32, buffer=nl.sbuf) nisa.tensor_tensor( dst=k_state_decay, data1=k_raw, data2=exp_gl_minus_gc, op=nl.multiply ) kv_psum = nl.ndarray((P_MAX, dim), dtype=nl.float32, buffer=nl.psum) nisa.nc_matmul(dst=kv_psum, stationary=k_state_decay, moving=v_new) kv_outer = nl.ndarray((P_MAX, dim), dtype=nl.float32, buffer=nl.sbuf) nisa.tensor_copy(dst=kv_outer, src=kv_psum) # state_decayed = state * exp(g_last) (per-dim element-wise) state_decayed = nl.ndarray((P_MAX, dim), dtype=nl.float32, buffer=nl.sbuf) nisa.tensor_tensor(dst=state_decayed, data1=state, data2=exp_gl, op=nl.multiply) state_new = nl.ndarray((P_MAX, dim), dtype=nl.float32, buffer=nl.sbuf) nisa.tensor_tensor(dst=state_new, data1=state_decayed, data2=kv_outer, op=nl.add) nisa.dma_copy(dst=state_out, src=state_new) return output, state_out, N_out, w_out