File size: 6,818 Bytes
5532501
 
6709b3f
5532501
 
 
 
 
 
 
 
6709b3f
 
5532501
6709b3f
 
5532501
 
 
 
 
 
 
6709b3f
 
 
 
 
 
5532501
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
6709b3f
 
 
 
 
 
5532501
 
 
 
6709b3f
 
5532501
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
6709b3f
5532501
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
6709b3f
5532501
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
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
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
"""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, # [B, nv, dk] pre-normed + scaled
    key: nl.ndarray, # [B, nv, dk] pre-normed
    value: nl.ndarray, # [B, nv, dv]
    g_in: nl.ndarray, # [B, nv, dk] activated per-channel log-decay
    beta_in: nl.ndarray, # [B, nv, dk] write gate (broadcast across dk)
    state_in: nl.ndarray, # [B, nv, dk, dv] recurrent state IN
):
    """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)

    # Ping-pong state buffers: items are independent, so prefetch item i+1's state
    # while computing item i.
    S_bufs = [
        nl.ndarray((dk, dv), dtype=nl.float32, buffer=nl.sbuf),
        nl.ndarray((dk, dv), dtype=nl.float32, buffer=nl.sbuf),
    ]
    # Seed the first item's state.
    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]
        # Prefetch next item's state (independent of this item's compute).
        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])

        # ---- Load this item's pre-processed vectors as [dk,1] / [dv,1] columns ----
        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)))

        # ---- key-fold body ----
        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)

        # Fresh buffer, not in-place: S_h is a ping-pong buffer that item i+2's
        # prefetch reclaims; in-place would make that prefetch wait on the DMA-out.
        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,
        )

        # out = q @ S_updated -> [1, dv]
        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)

        # Write output and advanced state.
        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