| """NKI kernel for BATCHED KDA decode: advance B*nv (request, head) items one token. |
| |
| . Companion to `kda_recurrent_fwd` in nki_kda.py. |
| |
| Motivation (an internal kernel comparison): single-(b,h) decode is |
| overhead-dominated. Serving B requests one call at a time pays the fixed per-launch |
| cost and the prologue B*nv times. This kernel flattens all (request, head) items |
| into one work list and processes them in a single invocation, amortizing launch |
| overhead and filling the partition axis. |
| |
| CONTRACT (matches this package's kda_recurrent_fwd -- inputs are PRE-PROCESSED): |
| query : [B, nv, dk] L2-normed and scaled by 1/sqrt(dk) |
| key : [B, nv, dk] L2-normed |
| value : [B, nv, dv] |
| g_in : [B, nv, dk] per-channel ACTIVATED log-decay (<= 0) |
| beta : [B, nv, dk] write gate broadcast across dk (per-item scalar) |
| state : [B, nv, dk, dv] recurrent state IN/OUT (advanced one token in place) |
| |
| This differs from the reference library's batched-decode kernel, which consumes RAW pre-activation |
| inputs and fuses L2-norm / gate activation / beta sigmoid inside. We keep the |
| pre-processed contract so this kernel is a drop-in batched sibling of |
| kda_recurrent_fwd and validates against the same fla-core reference. |
| |
| Per item, the recurrence is the key-fold body: |
| k_dec = k * exp_g # fold per-channel decay into the key |
| kv_mem = k_dec @ S # [dv], contract dk (undecayed S) |
| diff = v - kv_mem |
| S = Diag(exp_g) @ S + outer(k*beta, diff) # fused decay + update |
| out = q @ S # [dv] |
| |
| Written with PLAIN for-loops and indexed access only (NO list/generator |
| comprehensions, NO tuple-unpacking targets) so it compiles on the public NKI |
| frontend (0.5.0 and 0.6.0). |
| |
| nc_matmul(dst, stationary, moving) = stationary.T @ moving, contract partition dim. |
| """ |
|
|
| import nki |
| import nki.isa as nisa |
| import nki.language as nl |
|
|
| P_MAX = 128 |
|
|
|
|
| @nki.jit |
| def kda_decode_batch( |
| query: nl.ndarray, |
| key: nl.ndarray, |
| value: nl.ndarray, |
| g_in: nl.ndarray, |
| beta_in: nl.ndarray, |
| state_in: nl.ndarray, |
| ): |
| """Batched single-token KDA decode for all B*nv items in one call. |
| |
| Returns: |
| out : [B, nv, dv] this token's attention output |
| state_out : [B, nv, dk, dv] state advanced by one token |
| """ |
| B, nv, dk = query.shape |
| dv = value.shape[-1] |
| n_items = B * nv |
|
|
| out = nl.ndarray((B, nv, dv), dtype=query.dtype, buffer=nl.shared_hbm) |
| state_out = nl.ndarray((B, nv, dk, dv), dtype=nl.float32, buffer=nl.shared_hbm) |
|
|
| |
| |
| S_bufs = [ |
| nl.ndarray((dk, dv), dtype=nl.float32, buffer=nl.sbuf), |
| nl.ndarray((dk, dv), dtype=nl.float32, buffer=nl.sbuf), |
| ] |
| |
| nisa.dma_copy(dst=S_bufs[0], src=state_in[0, 0, 0:dk, 0:dv]) |
|
|
| for i in nl.static_range(n_items): |
| b = i // nv |
| h = i % nv |
|
|
| S_h = S_bufs[i % 2] |
| |
| if i + 1 < n_items: |
| nb = (i + 1) // nv |
| nh = (i + 1) % nv |
| nisa.dma_copy(dst=S_bufs[(i + 1) % 2], src=state_in[nb, nh, 0:dk, 0:dv]) |
|
|
| |
| q_t = nl.ndarray((dk, 1), dtype=query.dtype, buffer=nl.sbuf) |
| nisa.dma_copy(dst=q_t, src=query[b, h, 0:dk].reshape((dk, 1))) |
| k_t = nl.ndarray((dk, 1), dtype=key.dtype, buffer=nl.sbuf) |
| nisa.dma_copy(dst=k_t, src=key[b, h, 0:dk].reshape((dk, 1))) |
| v_t = nl.ndarray((dv, 1), dtype=value.dtype, buffer=nl.sbuf) |
| nisa.dma_copy(dst=v_t, src=value[b, h, 0:dv].reshape((dv, 1))) |
| g_col = nl.ndarray((dk, 1), dtype=g_in.dtype, buffer=nl.sbuf) |
| nisa.dma_copy(dst=g_col, src=g_in[b, h, 0:dk].reshape((dk, 1))) |
| beta_t = nl.ndarray((dk, 1), dtype=beta_in.dtype, buffer=nl.sbuf) |
| nisa.dma_copy(dst=beta_t, src=beta_in[b, h, 0:dk].reshape((dk, 1))) |
|
|
| |
| exp_g_col = nl.ndarray((dk, 1), dtype=nl.float32, buffer=nl.sbuf) |
| nisa.activation(dst=exp_g_col, op=nl.exp, data=g_col, bias=None, scale=1.0) |
|
|
| k_dec = nl.ndarray((dk, 1), dtype=nl.float32, buffer=nl.sbuf) |
| nisa.tensor_tensor(dst=k_dec, data1=k_t, data2=exp_g_col, op=nl.multiply) |
|
|
| kv_mem_psum = nl.ndarray((1, dv), dtype=nl.float32, buffer=nl.psum) |
| nisa.nc_matmul(dst=kv_mem_psum, stationary=k_dec, moving=S_h) |
| kv_mem = nl.ndarray((1, dv), dtype=nl.float32, buffer=nl.sbuf) |
| nisa.tensor_copy(dst=kv_mem, src=kv_mem_psum, engine=nisa.scalar_engine) |
|
|
| v_row_psum = nl.ndarray((1, dv), dtype=nl.float32, buffer=nl.psum) |
| nisa.nc_transpose(dst=v_row_psum, data=v_t) |
| v_row = nl.ndarray((1, dv), dtype=nl.float32, buffer=nl.sbuf) |
| nisa.tensor_copy(dst=v_row, src=v_row_psum, engine=nisa.scalar_engine) |
| diff = nl.ndarray((1, dv), dtype=nl.float32, buffer=nl.sbuf) |
| nisa.tensor_tensor(dst=diff, data1=v_row, data2=kv_mem, op=nl.subtract) |
|
|
| kbeta = nl.ndarray((dk, 1), dtype=nl.float32, buffer=nl.sbuf) |
| nisa.tensor_tensor(dst=kbeta, data1=k_t, data2=beta_t, op=nl.multiply) |
| kbeta_row_psum = nl.ndarray((1, dk), dtype=nl.float32, buffer=nl.psum) |
| nisa.nc_transpose(dst=kbeta_row_psum, data=kbeta) |
| kbeta_row = nl.ndarray((1, dk), dtype=nl.float32, buffer=nl.sbuf) |
| nisa.tensor_copy(dst=kbeta_row, src=kbeta_row_psum, engine=nisa.scalar_engine) |
|
|
| outer_psum = nl.ndarray((dk, dv), dtype=nl.float32, buffer=nl.psum) |
| nisa.nc_matmul(dst=outer_psum, stationary=kbeta_row, moving=diff) |
|
|
| |
| |
| S_updated = nl.ndarray((dk, dv), dtype=nl.float32, buffer=nl.sbuf) |
| nisa.scalar_tensor_tensor( |
| dst=S_updated, |
| data=S_h, |
| op0=nl.multiply, |
| operand0=exp_g_col, |
| op1=nl.add, |
| operand1=outer_psum, |
| ) |
|
|
| |
| o_psum = nl.ndarray((1, dv), dtype=nl.float32, buffer=nl.psum) |
| nisa.nc_matmul(dst=o_psum, stationary=q_t, moving=S_updated) |
| o_row = nl.ndarray((1, dv), dtype=nl.float32, buffer=nl.sbuf) |
| nisa.tensor_copy(dst=o_row, src=o_psum, engine=nisa.scalar_engine) |
|
|
| |
| nisa.dma_copy(dst=out[b, h, 0:dv].reshape((1, dv)), src=o_row) |
| nisa.dma_copy(dst=state_out[b, h, 0:dk, 0:dv], src=S_updated) |
|
|
| return out, state_out |
|
|