1use std::collections::BTreeMap;
9use std::sync::atomic::{AtomicU64, AtomicUsize};
10
11use anyhow::{Result, bail};
12use atlas_core::config::{ModelConfig, PeftAdapterConfig};
13use spark_runtime::gpu::{DevicePtr, GpuBackend};
14use spark_runtime::weights::WeightStore;
15
16use super::*;
17use crate::layers::ops::lora_delta::LoraPair;
18use crate::weight_map::DenseWeight;
19
20fn check_family(cfg: &ModelConfig) -> Result<()> {
30 if !(cfg.is_qwen35_dense()
31 || cfg.model_type == "holo3_1_moe"
32 || cfg.model_type == "qwen3_6_moe")
33 {
34 bail!(
35 "REJECT[unvalidated-family]: LoRA v0 is validated on qwen3_5 dense \
36 (holo-3.1-0.8b), holo3_1_moe (holo-3.1-35b-a3b), and qwen3_6_moe \
37 (Qwen3.6-35B-A3B) only; model_type='{}', num_experts={}",
38 cfg.model_type,
39 cfg.num_experts
40 );
41 }
42 Ok(())
43}
44
45#[allow(clippy::type_complexity)]
53fn pack_slot(
54 slot: usize,
55 name: &str,
56 adapter_store: &WeightStore,
57 peft: &PeftAdapterConfig,
58 found: &BTreeMap<(usize, LoraModule), [Option<String>; 2]>,
59 cfg: &ModelConfig,
60 gpu: &dyn GpuBackend,
61 pool: DevicePtr,
62 max_lora_rank: usize,
63) -> Result<(
64 Vec<Option<LoraLayerWeights>>,
65 BTreeMap<(usize, LoraModule), (u64, u64)>,
66)> {
67 let scale = peft.scaling();
68 let slot_bytes = pool_slot_bytes(cfg, max_lora_rank);
69 let mut layers: Vec<Option<LoraLayerWeights>> =
70 (0..cfg.num_hidden_layers).map(|_| None).collect();
71 let mut slot_ptrs: BTreeMap<(usize, LoraModule), (u64, u64)> = BTreeMap::new();
72 let mut off = slot * slot_bytes; for layer_idx in 0..cfg.num_hidden_layers {
77 let mut lw = LoraLayerWeights::empty(layer_idx);
78 let mut any = false;
79 for module in LoraModule::ALL {
80 if !module.applies_to_layer(cfg, layer_idx) {
81 continue;
82 }
83 let (out_dim, in_dim) = module.dims(cfg);
84 let a_off = off;
85 let b_off = off + max_lora_rank * in_dim * BF16_BYTES;
86 off = b_off + out_dim * max_lora_rank * BF16_BYTES;
87 let a_ptr = DevicePtr(pool.0 + a_off as u64);
88 let b_ptr = DevicePtr(pool.0 + b_off as u64);
89
90 let mut this = (0u64, 0u64); if let Some([Some(a_key), Some(b_key)]) = found.get(&(layer_idx, module)) {
92 let a_t = adapter_store.get(a_key)?;
94 let mut a_host = vec![0u8; peft.r * in_dim * BF16_BYTES];
95 gpu.copy_d2h(a_t.ptr, &mut a_host)?;
96 gpu.copy_h2d(&a_host, a_ptr)?;
97 let b_t = adapter_store.get(b_key)?;
99 let mut b_src = vec![0u8; out_dim * peft.r * BF16_BYTES];
100 gpu.copy_d2h(b_t.ptr, &mut b_src)?;
101 let mut b_host = vec![0u8; out_dim * max_lora_rank * BF16_BYTES];
102 for row in 0..out_dim {
103 let d = row * max_lora_rank * BF16_BYTES;
104 let s = row * peft.r * BF16_BYTES;
105 b_host[d..d + peft.r * BF16_BYTES]
106 .copy_from_slice(&b_src[s..s + peft.r * BF16_BYTES]);
107 }
108 gpu.copy_h2d(&b_host, b_ptr)?;
109
110 let pair = LoraPair {
111 a: DenseWeight { weight: a_ptr },
112 b: DenseWeight { weight: b_ptr },
113 rank: peft.r as u32,
114 k_in: in_dim as u32,
115 n_out: out_dim as u32,
116 scale,
117 max_rank: max_lora_rank as u32,
120 };
121 tracing::info!(
122 "LoRA: slot {slot} '{name}' layer {layer_idx} {module:?} r={} \
123 scale={:.6} A=[{},{}] B=[{},{}] (padded to max_rank={})",
124 peft.r,
125 scale,
126 peft.r,
127 in_dim,
128 out_dim,
129 peft.r,
130 max_lora_rank
131 );
132 match module {
133 LoraModule::QProj => lw.q_proj = Some(pair),
134 LoraModule::KProj => lw.k_proj = Some(pair),
135 LoraModule::VProj => lw.v_proj = Some(pair),
136 LoraModule::OProj => lw.o_proj = Some(pair),
137 LoraModule::GateProj => lw.gate_proj = Some(pair),
138 LoraModule::UpProj => lw.up_proj = Some(pair),
139 LoraModule::DownProj => lw.down_proj = Some(pair),
140 LoraModule::OutProj => lw.out_proj = Some(pair),
141 }
142 this = (a_ptr.0, b_ptr.0);
143 any = true;
144 }
145 slot_ptrs.insert((layer_idx, module), this);
146 }
147 if any {
148 layers[layer_idx] = Some(lw);
149 }
150 }
151 debug_assert_eq!(off, (slot + 1) * slot_bytes); Ok((layers, slot_ptrs))
153}
154
155pub fn load_lora_adapters_multi(
166 adapters: &[LoraAdapterInput<'_>],
167 cfg: &ModelConfig,
168 gpu: &dyn GpuBackend,
169 max_loras: usize,
170 max_lora_rank: usize,
171) -> Result<LoraWeights> {
172 check_family(cfg)?;
173 if adapters.is_empty() {
174 bail!("REJECT[no-adapters]: load_lora_adapters_multi called with an empty set");
175 }
176 if adapters.len() > max_loras {
177 bail!(
178 "REJECT[too-many-adapters]: {} --lora-adapter given but --max-loras={} \
179 (pool has {} slots); raise --max-loras or stage the extras on an \
180 $ATLAS_LORA_PEER for on-demand RDMA swap",
181 adapters.len(),
182 max_loras,
183 max_loras
184 );
185 }
186
187 let mut audited: Vec<AuditedAdapter> = Vec::with_capacity(adapters.len());
190 for a in adapters {
191 audited.push(audit_adapter(a.store, &a.peft, cfg, max_lora_rank)?);
192 }
193
194 let expert_rank = max_lora_expert_rank();
197 let expert_total: usize = adapters
198 .iter()
199 .zip(&audited)
200 .map(|(_, au)| {
201 let (ek, rl) = expert_pack::key_lists(&au.router, &au.experts);
202 expert_router_bytes(cfg, &ek, &rl, expert_rank)
203 })
204 .sum();
205
206 let pool_bytes = pool_slot_bytes(cfg, max_lora_rank) * max_loras;
209 let free = gpu.free_memory()?;
210 if (pool_bytes + expert_total) * 2 > free {
211 bail!(
212 "OOM pre-flight (LoRA pool): {:.1} MiB attn pool ({} slots) + {:.1} MiB \
213 expert/router pool would leave < 1× headroom of {:.1} MiB free; every \
214 pool byte comes directly out of the KV-cache budget on GB10 unified memory",
215 pool_bytes as f64 / (1024.0 * 1024.0),
216 max_loras,
217 expert_total as f64 / (1024.0 * 1024.0),
218 free as f64 / (1024.0 * 1024.0),
219 );
220 }
221 let pool = gpu.alloc(pool_bytes)?;
222 gpu.memset(pool, 0, pool_bytes)?;
223 let expert_pool = if expert_total > 0 {
226 let ep = gpu.alloc(expert_total)?;
227 gpu.memset(ep, 0, expert_total)?;
228 Some(ep)
229 } else {
230 None
231 };
232 let mut expert_off = 0usize;
233
234 let mut slots: Vec<AdapterSlot> = Vec::with_capacity(adapters.len());
237 let mut a_tabs: BTreeMap<(usize, LoraModule), Vec<u64>> = BTreeMap::new();
238 let mut b_tabs: BTreeMap<(usize, LoraModule), Vec<u64>> = BTreeMap::new();
239 let mut overlay_raw: Vec<Option<OverlayRawSlot>> = Vec::with_capacity(adapters.len());
243 for (k, a) in adapters.iter().enumerate() {
244 overlay_raw.push(stage_overlay_raw(
245 a.store,
246 &audited[k].overlay,
247 &a.peft,
248 cfg.hidden_size,
249 gpu,
250 )?);
251 let (mut layers, slot_ptrs) = pack_slot(
252 k,
253 &a.name,
254 a.store,
255 &a.peft,
256 &audited[k].attn,
257 cfg,
258 gpu,
259 pool,
260 max_lora_rank,
261 )?;
262 if let Some(ep) = expert_pool {
265 let packed = expert_pack::pack_into(
266 &mut layers,
267 a.store,
268 &a.peft,
269 &audited[k].router,
270 &audited[k].experts,
271 cfg,
272 gpu,
273 ep,
274 expert_rank,
275 &mut expert_off,
276 )?;
277 if packed > 0 {
278 tracing::info!(
279 "LoRA: slot {k} '{}' packed {packed} router/expert pair(s) \
280 (expert_rank={expert_rank})",
281 a.name
282 );
283 }
284 }
285 for ((layer, module), (a_ptr, b_ptr)) in slot_ptrs {
286 a_tabs
287 .entry((layer, module))
288 .or_insert_with(|| vec![0u64; max_loras])[k] = a_ptr;
289 b_tabs
290 .entry((layer, module))
291 .or_insert_with(|| vec![0u64; max_loras])[k] = b_ptr;
292 }
293 slots.push(AdapterSlot {
294 name: a.name.clone(),
295 adapter_config: a.peft.clone(),
296 layers,
297 generation: 0, });
299 }
300
301 let pinned = slots.len();
311 let num_layers = cfg.num_hidden_layers;
312 while slots.len() < max_loras {
313 slots.push(AdapterSlot {
314 name: String::new(),
315 adapter_config: PeftAdapterConfig {
316 r: 1,
317 lora_alpha: 0.0,
318 target_modules: Vec::new(),
319 target_modules_pattern: None,
320 use_rslora: false,
321 layers_to_transform: None,
322 trainable_token_indices: Vec::new(),
323 modules_to_save: Vec::new(),
324 lora_embedding: false,
325 },
326 layers: vec![None; num_layers],
327 generation: 0,
328 });
329 }
330
331 let mk = |tab: &[u64]| -> Result<DevicePtr> {
335 let bytes: Vec<u8> = tab.iter().flat_map(|p| p.to_le_bytes()).collect();
336 let d = gpu.alloc(bytes.len())?;
337 gpu.copy_h2d(&bytes, d)?;
338 Ok(d)
339 };
340 let mut tables = BTreeMap::new();
341 for (key, a_tab) in &a_tabs {
342 let b_tab = &b_tabs[key];
343 tables.insert(*key, (mk(a_tab)?, mk(b_tab)?));
344 }
345
346 debug_assert_eq!(expert_off, expert_total, "expert pool filled exactly");
350 let scale_vals = scale_table_values(adapters, max_loras);
351 let scale_bytes: Vec<u8> = scale_vals.iter().flat_map(|s| s.to_le_bytes()).collect();
352 let scale_table = gpu.alloc(scale_bytes.len())?;
353 gpu.copy_h2d(&scale_bytes, scale_table)?;
354
355 Ok(LoraWeights {
356 name: slots[0].name.clone(),
357 adapter_config: slots[0].adapter_config.clone(),
358 max_rank: max_lora_rank,
359 max_loras,
360 pool,
361 pool_bytes,
362 expert_pool,
363 expert_pool_bytes: expert_total,
364 slots,
365 active: 0,
366 tables,
367 scale_table,
368 ref_counts: (0..max_loras).map(|_| AtomicUsize::new(0)).collect(),
371 pinned,
372 last_used: (0..max_loras).map(|_| AtomicU64::new(0)).collect(),
373 lru_tick: AtomicU64::new(0),
374 overlay_raw,
375 })
376}
377
378pub fn pack_store_into_slot(
389 lw: &mut LoraWeights,
390 slot: usize,
391 name: &str,
392 store: &WeightStore,
393 peft: &PeftAdapterConfig,
394 cfg: &ModelConfig,
395 gpu: &dyn GpuBackend,
396) -> Result<Vec<Option<LoraLayerWeights>>> {
397 if slot >= lw.max_loras {
398 bail!(
399 "LoRA disk swap: slot {slot} >= max_loras {} (pool has {} slots)",
400 lw.max_loras,
401 lw.max_loras
402 );
403 }
404 let busy = lw.slot_ref_count(slot);
409 if busy > 0 {
410 bail!(
411 "LoRA disk swap REFUSED: slot {slot} has {busy} in-flight sequence(s) \
412 (ref_count>0); cannot replace an adapter mid-decode"
413 );
414 }
415 validate_peft_config(peft, lw.max_rank)?;
416 let audited = audit_adapter(store, peft, cfg, lw.max_rank)?;
417 if expert_pack::present(&audited.router, &audited.experts) {
418 bail!(
419 "LoRA disk swap REFUSED: adapter '{name}' carries router/expert deltas \
420 (Feature-1); runtime slot-swap of the expert pool is a phase-2 followup"
421 );
422 }
423 if !audited.overlay.is_empty() {
424 bail!(
425 "LoRA disk swap REFUSED: adapter '{name}' ships token-overlay tensors \
426 (Feature-2); runtime slot-swap of the overlay tables is a phase-2 \
427 followup (would silently drop the overlay otherwise)"
428 );
429 }
430 let found = audited.attn;
431 let slot_bytes = pool_slot_bytes(cfg, lw.max_rank);
432 gpu.memset(
433 DevicePtr(lw.pool.0 + (slot * slot_bytes) as u64),
434 0,
435 slot_bytes,
436 )?;
437 let (layers, _slot_ptrs) = pack_slot(
438 slot,
439 name,
440 store,
441 peft,
442 &found,
443 cfg,
444 gpu,
445 lw.pool,
446 lw.max_rank,
447 )?;
448 lw.slots[slot].name = name.to_string();
449 lw.slots[slot].adapter_config = peft.clone();
450 lw.slots[slot].layers = layers.clone();
451 lw.refresh_slot_tables(slot, &layers, peft.scaling(), gpu)?;
455 lw.slots[slot].generation = lw.slots[slot].generation.wrapping_add(1);
459 Ok(layers)
460}
461
462pub fn load_lora_adapters_generic(
466 adapter_store: &WeightStore,
467 peft: &PeftAdapterConfig,
468 cfg: &ModelConfig,
469 gpu: &dyn GpuBackend,
470 max_loras: usize,
471 max_lora_rank: usize,
472) -> Result<LoraWeights> {
473 let inputs = [LoraAdapterInput {
474 name: String::new(),
475 store: adapter_store,
476 peft: peft.clone(),
477 }];
478 load_lora_adapters_multi(&inputs, cfg, gpu, max_loras, max_lora_rank)
479}