Skip to main content

//clika-runtime/io.clika.runtime/Ops/ssdUpdate

ssdUpdate

[common]
fun ssdUpdate(input: Tensor, dt: Tensor, aRate: Tensor, bMat: Tensor, cMat: Tensor, dSkip: Tensor?, dtBias: Tensor?, gate: Tensor?, state: Tensor, seqLens: Tensor? = null, slotIds: Tensor? = null, dtSoftplus: Boolean = false): Tensor

ssdUpdate(input: Tensor, dt: Tensor, aRate: Tensor, bMat: Tensor, cMat: Tensor, dSkip: Tensor?, dtBias: Tensor?, gate: Tensor?, state: Tensor, seqLens: Tensor? = null, slotIds: Tensor? = null, dtSoftplus: Boolean = false): the ssd_update operator. Mamba2 / SSD selective-state serving step over a per-sequence [H, dim, dstate] Float32 state (updated IN PLACE; the returned tensor is out [B, T, H, dim], dtype following input). Per token: the state decays by exp(dt'*A[h]) (dt' = softplus(dt + dt_bias[h]) when dt_softplus), accumulates dt'*(x (outer) B), and emits S*C + D[h]*x (optionally silu-gated by gate). A/D/dt_bias are per-head Float32 parameters (A carries the NEGATIVE decay rate; fold -exp(A_log) at bind); B/C are [B, T, G, dstate] with H % G == 0. Decode is T == 1; a prefill runs the same sequential law. seq_lens ([B] Int32) bounds ragged rows. slot_ids ([B] Int32, device-resident) addresses state as a SLAB [num_slots, H, dim, dstate]: batch row b reads/updates slab row slot_ids[b] in place (ids in range and DISTINCT per call, the caller's contract); absent keeps state row b.