| |
| |
|
|
| """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, |
| key, |
| value, |
| beta_broadcast, |
| g_cumsum, |
| g_last, |
| state_in, |
| ): |
| """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 |
|
|
| 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) |
| |
| 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) |
|
|
| |
| 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) |
|
|
| |
| |
| |
| |
|
|
| 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, |
| ) |
|
|
| |
| 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) |
|
|
| |
| 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_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_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 = 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 = 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, |
| ) |
|
|
| |
| |
| |
| |
| |
| |
| 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 = 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 = 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 = 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) |
|
|
| |
| |
| |
| 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) |
|
|
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| 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 = 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) |
|
|
| |
| |
| |
| |
| |
| 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) |
|
|
| |
| |
| |
| nisa.dma_copy(dst=N_out, src=P_acc) |
|
|
| |
| |
| |
| 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 = 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) |
|
|
| |
| |
| |
| nisa.dma_copy(dst=w_out, src=k_cumdecay) |
|
|
| |
| |
| |
| |
| |
| 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) |
|
|
| |
| |
| |
| 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) |
|
|
| |
| |
| |
| 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) |
|
|
| |
| |
| |
| 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_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) |
|
|
| |
| |
| |
| |
| 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 = 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 |
|
|