spark_model/lora/
overlay_build.rs

1// SPDX-License-Identifier: AGPL-3.0-only
2
3//! Token-overlay GPU build (Feature 2, Stage 1 + Stage 2). Splits the overlay
4//! materialization across the model-build ordering gap: the RAW adapter tensors
5//! live only while the adapter [`WeightStore`] is alive (loader time), but the
6//! served embed/lm_head tables the row-diff needs only exist after weight load.
7//! So:
8//!
9//! - **Stage 1** ([`stage_overlay_raw`], loader): copy the raw overlay tensors
10//!   into owned device scratch while the store is alive; stash an [`OverlayRaw`].
11//! - **Stage 2** ([`build_overlay`], `set_lora_weights`): served tables now
12//!   exist → row-diff kernel → [`build_override_set`] → compact override rows +
13//!   `slot_map` + `ids` → [`EmbedOverlay`]; free the raw scratch.
14//!
15//! The pure row-selection math ([`clamp_trainable_to_vocab`],
16//! [`build_override_set`], [`override_source`], [`ROWDIFF_THRESH`]) lives in
17//! `super::overlay`; the device pointer tables in `super::overlay_tables`; the
18//! forward hooks in `crate::model::token_overlay`.
19
20use anyhow::{Result, bail};
21use atlas_core::config::PeftAdapterConfig;
22use spark_runtime::gpu::{DevicePtr, GpuBackend};
23use spark_runtime::weights::WeightStore;
24
25use super::overlay::{
26    OverlayTensors, ROWDIFF_THRESH, build_override_set, clamp_trainable_to_vocab, override_source,
27};
28use crate::layers::ops::token_overlay::{self, OverlayKernels};
29
30const BF16_BYTES: usize = 2;
31const F32_BYTES: usize = 4;
32
33/// Stage-1 raw upload: the adapter's own overlay tensors copied into owned
34/// device scratch (the source `WeightStore` is freed after the loader). Dims
35/// are recorded so Stage 2 can shape the row-diff and compaction without the
36/// store. `embed_base`/`lmhead_base` also carry a `modules_to_save` full-weight
37/// replacement (row-diff finds the changed rows either way).
38#[derive(Debug, Clone, Copy, Default)]
39pub struct OverlayRaw {
40    pub embed_base: Option<DevicePtr>,  // [embed_r, h] bf16
41    pub embed_delta: Option<DevicePtr>, // [embed_t, h] f32
42    pub embed_r: u32,
43    pub embed_t: u32,
44    pub lmhead_base: Option<DevicePtr>,  // [lmhead_r, h] bf16
45    pub lmhead_delta: Option<DevicePtr>, // [lmhead_t, h] f32
46    pub lmhead_r: u32,
47    pub lmhead_t: u32,
48}
49
50/// Stage-1 raw upload + the (host) clamped trainable ids the delta rows align to.
51pub struct OverlayRawSlot {
52    pub raw: OverlayRaw,
53    /// PEFT `trainable_token_indices` VERBATIM (clamped to served vocab in
54    /// Stage 2 so the loader stays independent of the served geometry).
55    pub trainable: Vec<u32>,
56}
57
58/// Stage-2 compact result for one adapter slot's embed overlay. `rows` holds the
59/// `n_override` replacement embedding rows in ascending-`ids` order; `slot_map`
60/// is the `[vocab]` i32 lookup (`-1` default, `[id] = compact-row-index`) the
61/// embed kernel indexes by token id. `lmhead` is `Some` only for an UNTIED head
62/// that ships its own overlay tensors.
63#[derive(Debug)]
64pub struct EmbedOverlay {
65    pub rows: DevicePtr,     // [n_override, h] bf16
66    pub ids_dev: DevicePtr,  // u32[n_override] ascending vocab ids
67    pub slot_map: DevicePtr, // i32[vocab], -1 default
68    pub n_override: u32,
69    /// `slot_map` length — the served vocab this overlay was built against.
70    /// Recorded HERE (where `slot_map` is allocated) as the SSOT bound the
71    /// embed kernel's `ids[r] < vocab` guard checks against (CWE-125).
72    pub vocab: u32,
73    pub lmhead: Option<LmHeadOverlay>,
74}
75
76/// Distinct output-projection overlay for an untied lm_head. Rows are recomputed
77/// into logits by a `dot(hidden, row)` per overridden id (no `slot_map`).
78#[derive(Debug)]
79pub struct LmHeadOverlay {
80    pub rows: DevicePtr,    // [n_override, h] bf16
81    pub ids_dev: DevicePtr, // u32[n_override]
82    pub n_override: u32,
83}
84
85/// Round-to-nearest-even f32 → bf16 bit pattern (build-time host cast for the
86/// f32 `trainable_tokens_delta` replacement rows).
87fn f32_to_bf16(x: f32) -> u16 {
88    let bits = x.to_bits();
89    let round = ((bits >> 16) & 1) + 0x7fff;
90    (bits.wrapping_add(round) >> 16) as u16
91}
92
93/// Copy one on-device tensor into a fresh owned device buffer, returning the
94/// buffer and its leading dim (`shape[0]`). Validates the trailing dim == `h`.
95fn stage_tensor(
96    store: &WeightStore,
97    name: &str,
98    h: usize,
99    gpu: &dyn GpuBackend,
100) -> Result<(DevicePtr, u32)> {
101    let t = store.get(name)?;
102    if t.shape.len() != 2 || t.shape[1] != h {
103        bail!(
104            "REJECT[overlay-shape]: '{name}' is {:?}, expected [rows, {h}] (hidden)",
105            t.shape
106        );
107    }
108    let bytes: usize = t.shape.iter().product::<usize>() * t.dtype.byte_size();
109    let dst = gpu.alloc(bytes)?;
110    gpu.copy_d2d(t.ptr, dst, bytes)?;
111    Ok((dst, t.shape[0] as u32))
112}
113
114/// Stage 1 (loader): upload the classified overlay tensors of one adapter into
115/// owned device scratch. `Ok(None)` when the adapter ships no overlay tensors.
116/// `embed_full`/`lmhead_full` (`modules_to_save`) map onto the `*_base` slot.
117pub fn stage_overlay_raw(
118    store: &WeightStore,
119    overlay: &OverlayTensors,
120    peft: &PeftAdapterConfig,
121    h: usize,
122    gpu: &dyn GpuBackend,
123) -> Result<Option<OverlayRawSlot>> {
124    if overlay.is_empty() || overlay.lora_embedding_seen {
125        // lora_embedding is a Tier-2 reject handled at audit; never staged.
126        return Ok(None);
127    }
128    let mut raw = OverlayRaw::default();
129    let embed_base_name = overlay.embed_base.as_ref().or(overlay.embed_full.as_ref());
130    if let Some(name) = embed_base_name {
131        let (p, r) = stage_tensor(store, name, h, gpu)?;
132        raw.embed_base = Some(p);
133        raw.embed_r = r;
134    }
135    if let Some(name) = &overlay.embed_delta {
136        let (p, t) = stage_tensor(store, name, h, gpu)?;
137        raw.embed_delta = Some(p);
138        raw.embed_t = t;
139    }
140    let lmhead_base_name = overlay
141        .lmhead_base
142        .as_ref()
143        .or(overlay.lmhead_full.as_ref());
144    if let Some(name) = lmhead_base_name {
145        let (p, r) = stage_tensor(store, name, h, gpu)?;
146        raw.lmhead_base = Some(p);
147        raw.lmhead_r = r;
148    }
149    if let Some(name) = &overlay.lmhead_delta {
150        let (p, t) = stage_tensor(store, name, h, gpu)?;
151        raw.lmhead_delta = Some(p);
152        raw.lmhead_t = t;
153    }
154    if raw.embed_base.is_none() && raw.lmhead_base.is_none() {
155        // A bare delta with no base to diff against is meaningless; a real
156        // trainable_tokens adapter always ships base_layer.weight.
157        bail!(
158            "REJECT[overlay-no-base]: overlay tensors present but no embed/lm_head base row table"
159        );
160    }
161    Ok(Some(OverlayRawSlot {
162        raw,
163        trainable: peft.trainable_token_indices.clone(),
164    }))
165}
166
167/// The compact-override intermediate shared by the embed and lm_head builds:
168/// the ascending overridden ids and their `[n, h]` bf16 replacement rows on
169/// device. `None` when nothing is overridden.
170struct Compact {
171    ids: Vec<u32>,
172    rows_dev: DevicePtr,
173    ids_dev: DevicePtr,
174    n: u32,
175}
176
177/// Row-diff `base` vs `served`, union with `kept` trainable ids, and materialize
178/// the compact bf16 replacement rows (delta wins over base per `override_source`).
179#[allow(clippy::too_many_arguments)]
180fn compact_override(
181    gpu: &dyn GpuBackend,
182    kernels: &OverlayKernels,
183    base: DevicePtr,
184    delta: Option<DevicePtr>,
185    r: u32,
186    served: DevicePtr,
187    vocab: usize,
188    h: usize,
189    kept: &[u32],
190    stream: u64,
191) -> Result<Option<Compact>> {
192    let r_eff = (r as usize).min(vocab);
193    // Row-diff kernel: flags[row] = max_i |base-served| > thresh.
194    let flags_dev = gpu.alloc(r_eff.max(1))?;
195    if r_eff > 0 {
196        token_overlay::embed_rowdiff(
197            gpu,
198            kernels.rowdiff,
199            base,
200            served,
201            flags_dev,
202            r_eff as u32,
203            h as u32,
204            ROWDIFF_THRESH,
205            stream,
206        )?;
207        gpu.synchronize(stream)?;
208    }
209    let mut flags = vec![0u8; r_eff];
210    if r_eff > 0 {
211        gpu.copy_d2h(flags_dev, &mut flags)?;
212    }
213    let _ = gpu.free(flags_dev);
214    let row_diff: Vec<bool> = flags.iter().map(|&b| b != 0).collect();
215    let ids = build_override_set(&row_diff, kept);
216    if ids.is_empty() {
217        return Ok(None);
218    }
219    let n = ids.len();
220    // Materialize compact rows host-side (build-time, n is small): delta rows
221    // f32→bf16, base rows copied verbatim bf16.
222    let mut compact = vec![0u8; n * h * BF16_BYTES];
223    let mut frow = vec![0u8; h * F32_BYTES];
224    for (ci, &id) in ids.iter().enumerate() {
225        let dst = &mut compact[ci * h * BF16_BYTES..(ci + 1) * h * BF16_BYTES];
226        match override_source(id, kept) {
227            Some(k) => {
228                let d = delta.ok_or_else(|| {
229                    anyhow::anyhow!(
230                        "REJECT[overlay-delta-missing]: trainable id {id} but no delta tensor"
231                    )
232                })?;
233                gpu.copy_d2h(d.offset(k * h * F32_BYTES), &mut frow)?;
234                for i in 0..h {
235                    let x = f32::from_le_bytes([
236                        frow[i * 4],
237                        frow[i * 4 + 1],
238                        frow[i * 4 + 2],
239                        frow[i * 4 + 3],
240                    ]);
241                    dst[i * 2..i * 2 + 2].copy_from_slice(&f32_to_bf16(x).to_le_bytes());
242                }
243            }
244            None => {
245                gpu.copy_d2h(base.offset(id as usize * h * BF16_BYTES), dst)?;
246            }
247        }
248    }
249    let rows_dev = gpu.alloc(compact.len())?;
250    gpu.copy_h2d(&compact, rows_dev)?;
251    let ids_bytes: Vec<u8> = ids.iter().flat_map(|i| i.to_le_bytes()).collect();
252    let ids_dev = gpu.alloc(ids_bytes.len())?;
253    gpu.copy_h2d(&ids_bytes, ids_dev)?;
254    Ok(Some(Compact {
255        ids,
256        rows_dev,
257        ids_dev,
258        n: n as u32,
259    }))
260}
261
262/// Stage 2 (`set_lora_weights`): row-diff against the served tables, compact the
263/// override rows, build the `slot_map`, and free the Stage-1 raw scratch.
264/// `Ok(None)` when the slot overrides nothing (silently inert = correct).
265///
266/// `tied` means the lm_head aliases the embed table (buffer aliasing OR a
267/// quantized head derived from embed) — the caller then reuses the embed rows
268/// for the logit recompute (`overlay_tables`), so no distinct lm_head build.
269#[allow(clippy::too_many_arguments)]
270pub fn build_overlay(
271    gpu: &dyn GpuBackend,
272    kernels: &OverlayKernels,
273    slot: &OverlayRawSlot,
274    served_embed: DevicePtr,
275    served_lmhead: DevicePtr,
276    vocab: usize,
277    h: usize,
278    tied: bool,
279    stream: u64,
280) -> Result<Option<EmbedOverlay>> {
281    if kernels.rowdiff.0 == 0 || kernels.embed_overlay.0 == 0 {
282        bail!(
283            "REJECT[overlay-kernels-missing]: adapter ships token-overlay tensors but the \
284             token_overlay CUDA kernels are not loaded (rebuild with the kernel image)"
285        );
286    }
287    let raw = &slot.raw;
288    let Some(embed_base) = raw.embed_base else {
289        return Ok(None); // lm_head-only overlay without an embed base: not supported here.
290    };
291    let (kept, skipped) = clamp_trainable_to_vocab(&slot.trainable, raw.embed_r as usize, vocab)?;
292    if skipped > 0 {
293        tracing::warn!(
294            "LoRA overlay: dropped {skipped} vocab-extension trainable id(s) beyond served vocab {vocab}"
295        );
296    }
297    let Some(embed) = compact_override(
298        gpu,
299        kernels,
300        embed_base,
301        raw.embed_delta,
302        raw.embed_r,
303        served_embed,
304        vocab,
305        h,
306        &kept,
307        stream,
308    )?
309    else {
310        return Ok(None);
311    };
312    // Host slot_map: [vocab] i32, -1 default, [id] = compact index.
313    let mut slot_map = vec![-1i32; vocab];
314    for (ci, &id) in embed.ids.iter().enumerate() {
315        slot_map[id as usize] = ci as i32;
316    }
317    let sm_bytes: Vec<u8> = slot_map.iter().flat_map(|i| i.to_le_bytes()).collect();
318    let slot_map_dev = gpu.alloc(sm_bytes.len())?;
319    gpu.copy_h2d(&sm_bytes, slot_map_dev)?;
320
321    // lm_head branch: distinct build only for an untied head that ships its own
322    // overlay tensors; tied heads reuse the embed rows (handled in the tables).
323    let lmhead = if let Some(base) = raw.lmhead_base.filter(|_| !tied) {
324        compact_override(
325            gpu,
326            kernels,
327            base,
328            raw.lmhead_delta,
329            raw.lmhead_r,
330            served_lmhead,
331            vocab,
332            h,
333            &kept,
334            stream,
335        )?
336        .map(|c| LmHeadOverlay {
337            rows: c.rows_dev,
338            ids_dev: c.ids_dev,
339            n_override: c.n,
340        })
341    } else {
342        None
343    };
344
345    // Free Stage-1 raw scratch (build-time; on error the buffers leak, which
346    // matches the existing adapter-load leak tolerance).
347    for p in [
348        raw.embed_base,
349        raw.embed_delta,
350        raw.lmhead_base,
351        raw.lmhead_delta,
352    ]
353    .into_iter()
354    .flatten()
355    {
356        let _ = gpu.free(p);
357    }
358
359    Ok(Some(EmbedOverlay {
360        rows: embed.rows_dev,
361        ids_dev: embed.ids_dev,
362        slot_map: slot_map_dev,
363        n_override: embed.n,
364        vocab: vocab as u32,
365        lmhead,
366    }))
367}
368
369#[cfg(test)]
370#[path = "overlay_build_tests.rs"]
371mod tests;