1use 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#[derive(Debug, Clone, Copy, Default)]
39pub struct OverlayRaw {
40 pub embed_base: Option<DevicePtr>, pub embed_delta: Option<DevicePtr>, pub embed_r: u32,
43 pub embed_t: u32,
44 pub lmhead_base: Option<DevicePtr>, pub lmhead_delta: Option<DevicePtr>, pub lmhead_r: u32,
47 pub lmhead_t: u32,
48}
49
50pub struct OverlayRawSlot {
52 pub raw: OverlayRaw,
53 pub trainable: Vec<u32>,
56}
57
58#[derive(Debug)]
64pub struct EmbedOverlay {
65 pub rows: DevicePtr, pub ids_dev: DevicePtr, pub slot_map: DevicePtr, pub n_override: u32,
69 pub vocab: u32,
73 pub lmhead: Option<LmHeadOverlay>,
74}
75
76#[derive(Debug)]
79pub struct LmHeadOverlay {
80 pub rows: DevicePtr, pub ids_dev: DevicePtr, pub n_override: u32,
83}
84
85fn 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
93fn 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
114pub 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 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 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
167struct Compact {
171 ids: Vec<u32>,
172 rows_dev: DevicePtr,
173 ids_dev: DevicePtr,
174 n: u32,
175}
176
177#[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 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 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#[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); };
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 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 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 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;