spark_model/layers/ops/
token_overlay.rs

1// SPDX-License-Identifier: AGPL-3.0-only
2
3//! Token-overlay kernel launchers (Feature 2). Thin `KernelLaunch` wrappers over
4//! `kernels/gb10/common/token_overlay.cu`:
5//! - [`embed_rowdiff`] — build-time: which adapter base rows differ from served.
6//! - [`embed_overlay_routed`] — forward: replace overridden vocab rows post-gather.
7//! - [`lmhead_overlay_routed`] — forward: recompute overridden logit columns.
8//!
9//! Argument order is in LOCKSTEP with the `.cu` signatures (cuLaunchKernel is
10//! type-blind). All device tables are load-time-fixed addresses; the only
11//! per-step arg is `seq_slot` (NULL ⇒ uniform `active`) — graph-capture safe.
12
13use anyhow::Result;
14use spark_runtime::gpu::{DevicePtr, GpuBackend, KernelHandle};
15use spark_runtime::kernel_args::KernelLaunch;
16
17use crate::layers::try_kernel;
18
19/// The four token-overlay kernels, resolved once at model construction via
20/// [`try_kernel`] (null-on-miss ⇒ the feature is silently unused rather than a
21/// hard init failure on a kernel image that predates the overlay).
22#[derive(Clone, Copy)]
23pub struct OverlayKernels {
24    pub rowdiff: KernelHandle,
25    pub embed_overlay: KernelHandle,
26    pub lmhead_overlay_bf16: KernelHandle,
27    pub lmhead_overlay_f32: KernelHandle,
28}
29
30impl Default for OverlayKernels {
31    /// All-null handles: the overlay feature is silently unused (every hook's
32    /// `kernel.0 == 0` guard fires). `KernelHandle` has no `Default` derive, so
33    /// this is spelled out.
34    fn default() -> Self {
35        Self {
36            rowdiff: KernelHandle(0),
37            embed_overlay: KernelHandle(0),
38            lmhead_overlay_bf16: KernelHandle(0),
39            lmhead_overlay_f32: KernelHandle(0),
40        }
41    }
42}
43
44impl OverlayKernels {
45    pub fn new(gpu: &dyn GpuBackend) -> Self {
46        Self {
47            rowdiff: try_kernel(gpu, "token_overlay", "embed_rowdiff_bf16"),
48            embed_overlay: try_kernel(gpu, "token_overlay", "embed_overlay_routed_bf16"),
49            lmhead_overlay_bf16: try_kernel(gpu, "token_overlay", "lmhead_overlay_routed_bf16"),
50            lmhead_overlay_f32: try_kernel(gpu, "token_overlay", "lmhead_overlay_routed_f32"),
51        }
52    }
53}
54
55/// `flags[r] = (max_i |base[r,i] - served[r,i]| > thresh)`. Grid one thread/row.
56#[allow(clippy::too_many_arguments)]
57pub fn embed_rowdiff(
58    gpu: &dyn GpuBackend,
59    kernel: KernelHandle,
60    base: DevicePtr,   // [rows, h] bf16
61    served: DevicePtr, // [rows, h] bf16
62    flags: DevicePtr,  // [rows] u8 out
63    rows: u32,
64    h: u32,
65    thresh: f32,
66    stream: u64,
67) -> Result<()> {
68    KernelLaunch::new(gpu, kernel)
69        .grid([rows.div_ceil(256), 1, 1])
70        .block([256, 1, 1])
71        .arg_ptr(base)
72        .arg_ptr(served)
73        .arg_ptr(flags)
74        .arg_u32(rows)
75        .arg_u32(h)
76        .arg_f32(thresh)
77        .launch(stream)
78}
79
80/// In-place row-replace of overridden vocab rows on the residual stream after
81/// the embed gather (BEFORE `scale_embeddings`). Per row `r`: `s = seq_slot[r]`
82/// (or `active` when NULL); `s<0` skip; `ids[r]>=vocab` skip (no overlay entry
83/// for a token beyond the overlay's served-vocab snapshot — CWE-125 guard);
84/// `slot=slot_map_tab[s][ids[r]]`; `slot<0` or `slot>=n_tab[s]` skip; copy
85/// `rows_tab[s][slot]` over `out[r]`.
86#[allow(clippy::too_many_arguments)]
87pub fn embed_overlay_routed(
88    gpu: &dyn GpuBackend,
89    kernel: KernelHandle,
90    ids: DevicePtr,      // [n] u32 token id per row
91    seq_slot: DevicePtr, // [n] i32 or NULL(0)
92    active: i32,
93    slot_map_tab: DevicePtr, // u64[L]
94    rows_tab: DevicePtr,     // u64[L]
95    n_tab: DevicePtr,        // u32[L] embed n_override per slot
96    out: DevicePtr,          // [n, h] bf16 in place
97    num_tokens: u32,
98    h: u32,
99    vocab: u32, // slot_map length (served vocab at overlay build)
100    stream: u64,
101) -> Result<()> {
102    KernelLaunch::new(gpu, kernel)
103        .grid([num_tokens, 1, 1])
104        .block([256, 1, 1])
105        .arg_ptr(ids)
106        .arg_ptr(seq_slot)
107        .arg_i32(active)
108        .arg_ptr(slot_map_tab)
109        .arg_ptr(rows_tab)
110        .arg_ptr(n_tab)
111        .arg_ptr(out)
112        .arg_u32(h)
113        .arg_u32(vocab)
114        .launch(stream)
115}
116
117/// In-place recompute of overridden logit columns (BEFORE softcap). One warp per
118/// `(row, j)`; `j` indexes the overridden-id slot of that row's adapter. Picks
119/// the bf16 or f32 logits kernel per `is_fp32`.
120#[allow(clippy::too_many_arguments)]
121pub fn lmhead_overlay_routed(
122    gpu: &dyn GpuBackend,
123    kernel: KernelHandle, // bf16 or f32 variant, selected by caller
124    hidden: DevicePtr,    // [m, h] bf16
125    seq_slot: DevicePtr,  // [m] i32 or NULL(0)
126    active: i32,
127    rows_tab: DevicePtr, // u64[L]
128    ids_tab: DevicePtr,  // u64[L]
129    n_tab: DevicePtr,    // u32[L]
130    logits: DevicePtr,   // [m, vocab] bf16 or f32 in place
131    m: u32,
132    max_n_override: u32,
133    h: u32,
134    vocab: u32,
135    stream: u64,
136) -> Result<()> {
137    KernelLaunch::new(gpu, kernel)
138        .grid([m, max_n_override, 1])
139        .block([32, 1, 1])
140        .arg_ptr(hidden)
141        .arg_ptr(seq_slot)
142        .arg_i32(active)
143        .arg_ptr(rows_tab)
144        .arg_ptr(ids_tab)
145        .arg_ptr(n_tab)
146        .arg_ptr(logits)
147        .arg_u32(h)
148        .arg_u32(vocab)
149        .launch(stream)
150}