1use crate::gpu::GpuBackend;
20use crate::weights::{
21 WeightLoader, WeightStore, WeightTensor, check_oom_guard, estimate_has_fp8,
22 estimate_load_bytes, evict_page_cache, f16_to_bf16_bytes,
23};
24use anyhow::{Context, Result, bail};
25use std::collections::HashMap;
26use std::fs::File;
27use std::path::{Path, PathBuf};
28use std::sync::mpsc::sync_channel;
29
30mod direct_io;
31mod header;
32
33use header::{parse_header, resolve_shards};
34
35pub struct FastSafetensorsLoader {
38 pub ep_rank: usize,
39 pub ep_world_size: usize,
40 pub num_experts: usize,
41 pub peak_memory_multiplier: Option<f64>,
42 pub skip_activation_scales: bool,
53 pub skip_mtp: bool,
63 pub try_direct_io: bool,
66 pub direct_io_tensor_cap: usize,
76 pub prefetch_shards: bool,
80 pub skip_vision: bool,
89}
90
91pub fn is_vision_tensor(name: &str) -> bool {
96 name.starts_with("model.visual.")
97 || name.starts_with("model.vision")
98 || name.starts_with("visual.")
99}
100
101pub const DEFAULT_DIRECT_IO_TENSOR_CAP: usize = 5000;
105
106impl Default for FastSafetensorsLoader {
107 fn default() -> Self {
108 Self::new()
109 }
110}
111
112#[path = "skip.rs"]
113mod skip;
114
115impl FastSafetensorsLoader {
116 pub fn new() -> Self {
117 Self {
118 ep_rank: 0,
119 ep_world_size: 1,
120 num_experts: 0,
121 peak_memory_multiplier: None,
122 skip_activation_scales: false,
123 skip_mtp: false,
124 try_direct_io: true,
125 direct_io_tensor_cap: DEFAULT_DIRECT_IO_TENSOR_CAP,
126 prefetch_shards: false,
127 skip_vision: false,
128 }
129 }
130
131 pub fn with_ep(ep_rank: usize, ep_world_size: usize, num_experts: usize) -> Self {
132 Self {
133 ep_rank,
134 ep_world_size,
135 num_experts,
136 peak_memory_multiplier: None,
137 skip_activation_scales: false,
138 skip_mtp: false,
139 try_direct_io: true,
140 direct_io_tensor_cap: DEFAULT_DIRECT_IO_TENSOR_CAP,
141 prefetch_shards: false,
142 skip_vision: false,
143 }
144 }
145}
146
147impl WeightLoader for FastSafetensorsLoader {
148 fn load(
149 &self,
150 model_dir: &Path,
151 gpu: &dyn GpuBackend,
152 oom_reserve_bytes: usize,
153 ) -> Result<WeightStore> {
154 let skip_fn = |name: &str| self.should_skip_tensor(name);
155
156 let (shard_files, tensor_to_shard): (Vec<PathBuf>, Option<HashMap<String, String>>) =
158 resolve_shards(model_dir)?;
159
160 let preflight_skip = |name: &str| skip_fn(name) || crate::weights::is_ngram_table(name);
167 {
168 let estimated = estimate_load_bytes(&shard_files, &preflight_skip)?;
169 let has_fp8 = estimate_has_fp8(&shard_files, &preflight_skip)?;
170 let mult = self
171 .peak_memory_multiplier
172 .unwrap_or(if has_fp8 { 1.5 } else { 1.3 });
173 let peak = (estimated as f64 * mult) as usize;
174 let free = gpu.free_memory()?;
175 let gib = |b: usize| b as f64 / (1024.0 * 1024.0 * 1024.0);
176 tracing::info!(
177 "Fast-load pre-flight: {:.2} GB on-disk, {:.1}x overhead = {:.2} GB peak, \
178 {:.2} GB free, {:.1} GB reserve (FP8: {})",
179 gib(estimated),
180 mult,
181 gib(peak),
182 gib(free),
183 gib(oom_reserve_bytes),
184 has_fp8,
185 );
186 crate::progress::preflight(gib(estimated), gib(free));
187 if peak + oom_reserve_bytes > free {
188 bail!(
189 "OOM pre-flight: peak {:.2} GB + {:.2} GB reserve exceeds {:.2} GB free. \
190 Use a smaller quantization or add more GPUs for EP.",
191 gib(peak),
192 gib(oom_reserve_bytes),
193 gib(free),
194 );
195 }
196 }
197
198 let mut weights: HashMap<String, WeightTensor> = HashMap::new();
200 let mut deferred: HashMap<String, crate::weights::DeferredTensor> = HashMap::new();
202 let total_shards = shard_files.len();
203 let initial_free = gpu.free_memory()?;
204 let mut offload_logged = false;
205
206 for (i, shard_path) in shard_files.iter().enumerate() {
207 let shard_name = shard_path
210 .file_name()
211 .and_then(|n| n.to_str())
212 .unwrap_or_default();
213 let tensor_filter: Option<Vec<String>> = tensor_to_shard.as_ref().map(|map| {
214 map.iter()
215 .filter(|(_, s)| *s == shard_name)
216 .map(|(t, _)| t.clone())
217 .collect()
218 });
219
220 tracing::info!(
221 "Fast-loading shard {}/{}: {}{}",
222 i + 1,
223 total_shards,
224 shard_name,
225 tensor_filter
226 .as_ref()
227 .map(|v| format!(" ({} tensors)", v.len()))
228 .unwrap_or_default(),
229 );
230 crate::progress::shard_start(i + 1, total_shards, shard_name);
231
232 load_shard_fast(
233 shard_path,
234 tensor_filter.as_deref(),
235 gpu,
236 &skip_fn,
237 self.try_direct_io,
238 self.direct_io_tensor_cap,
239 self.prefetch_shards,
240 &mut weights,
241 &mut deferred,
242 &mut offload_logged,
243 )?;
244
245 let free_now = gpu.free_memory().unwrap_or(0);
246 let used = initial_free.saturating_sub(free_now);
247 tracing::info!(
248 " Shard {}/{} done — GPU memory: {:.2} GB used, {:.2} GB free",
249 i + 1,
250 total_shards,
251 used as f64 / (1024.0 * 1024.0 * 1024.0),
252 free_now as f64 / (1024.0 * 1024.0 * 1024.0),
253 );
254 crate::progress::shard_done(
255 i + 1,
256 total_shards,
257 used as f64 / (1024.0 * 1024.0 * 1024.0),
258 free_now as f64 / (1024.0 * 1024.0 * 1024.0),
259 );
260 if !offload_logged {
261 check_oom_guard(
262 gpu,
263 oom_reserve_bytes,
264 &format!("fast weight loading (shard {}/{})", i + 1, total_shards),
265 )?;
266 }
267 }
268
269 let no_skip = |_: &str| false;
271 let extra = model_dir.join("extra_weights.safetensors");
272 if extra.exists() {
273 tracing::info!("Fast-loading extra_weights.safetensors");
274 let mut extra_offload = false;
275 load_shard_fast(
276 &extra,
277 None,
278 gpu,
279 &no_skip,
280 self.try_direct_io,
281 self.direct_io_tensor_cap,
282 self.prefetch_shards,
283 &mut weights,
284 &mut deferred,
285 &mut extra_offload,
286 )?;
287 }
288
289 tracing::info!("Fast-loaded {} weight tensors", weights.len());
290 let mut store = WeightStore::from_map(weights);
291 for (name, d) in deferred {
292 store.defer(name, d);
293 }
294 Ok(store)
295 }
296}
297
298#[allow(clippy::too_many_arguments)]
308fn load_shard_fast(
309 shard_path: &Path,
310 tensor_filter: Option<&[String]>,
311 gpu: &dyn GpuBackend,
312 skip_fn: &dyn Fn(&str) -> bool,
313 try_direct_io: bool,
314 direct_io_tensor_cap: usize,
315 prefetch_shards: bool,
316 out: &mut HashMap<String, WeightTensor>,
317 deferred_out: &mut HashMap<String, crate::weights::DeferredTensor>,
318 offload_logged: &mut bool,
319) -> Result<()> {
320 let mut meta_file = File::open(shard_path)
323 .with_context(|| format!("Failed to open {}", shard_path.display()))?;
324 let mut tensors = parse_header(&mut meta_file)?;
325 let file_size = meta_file.metadata()?.len();
326
327 if let Some(allow) = tensor_filter {
329 let allow_set: std::collections::HashSet<&str> = allow.iter().map(|s| s.as_str()).collect();
330 tensors.retain(|t| allow_set.contains(t.name.as_str()));
331 }
332 let mut deferred_here: Vec<(String, crate::weights::DeferredTensor)> = Vec::new();
339 #[allow(clippy::items_after_statements)]
340 tensors.retain(|t| {
341 if crate::weights::is_ngram_table(&t.name) {
342 deferred_here.push((
343 t.name.clone(),
344 crate::weights::DeferredTensor {
345 path: shard_path.to_path_buf(),
346 offset: t.abs_offset,
347 shape: t.shape.clone(),
348 dtype: t.dtype,
349 },
350 ));
351 return false;
352 }
353 !skip_fn(&t.name)
354 });
355 if !deferred_here.is_empty() {
356 tracing::info!(
357 "Deferred {} n-gram table(s) in {} — served from disk, not uploaded",
358 deferred_here.len(),
359 shard_path.display()
360 );
361 deferred_out.extend(deferred_here);
362 }
363
364 let wants_direct = try_direct_io && tensors.len() <= direct_io_tensor_cap;
369 if try_direct_io && !wants_direct {
370 tracing::info!(
371 " Shard has {} tensors (> {} cap) — using buffered+pipelined path",
372 tensors.len(),
373 direct_io_tensor_cap
374 );
375 }
376
377 let (direct_file, using_direct) = match wants_direct
379 .then(|| direct_io::open_direct(shard_path))
380 .transpose()
381 {
382 Ok(Some(f)) => (Some(f), true),
383 Ok(None) => (None, false),
384 Err(e) => {
385 tracing::warn!(
386 "O_DIRECT open failed for {} ({e}); falling back to buffered reads",
387 shard_path.display()
388 );
389 (None, false)
390 }
391 };
392 let buffered_file = File::open(shard_path)?;
393 let data_fd = direct_file.as_ref().unwrap_or(&buffered_file);
394 if prefetch_shards && !using_direct {
395 advise_prefetch_shard(&buffered_file, shard_path, file_size);
396 }
397
398 type ReadMsg = (usize, direct_io::AlignedBuffer, usize);
400 let (tx, rx) = sync_channel::<Result<ReadMsg>>(1);
401 let tensors_for_reader: Vec<(u64, usize)> =
402 tensors.iter().map(|t| (t.abs_offset, t.len)).collect();
403 let raw_fd = {
404 use std::os::unix::io::AsRawFd;
405 data_fd.as_raw_fd()
406 };
407
408 let _ = file_size; let reader_handle = std::thread::spawn(move || {
410 for (idx, (abs_offset, len)) in tensors_for_reader.iter().enumerate() {
411 let msg = direct_io::read_tensor_aligned(raw_fd, *abs_offset, *len, using_direct)
412 .map(|(buf, slice_start)| (idx, buf, slice_start));
413 if tx.send(msg).is_err() {
414 break; }
416 }
417 });
418
419 for result in rx {
421 let (idx, buf, slice_start) = result?;
422 let meta = &tensors[idx];
423 let raw = &buf.as_slice()[slice_start..slice_start + meta.len];
424 let converted: Vec<u8>;
427 let src: &[u8] = if meta.from_f16 {
428 converted = f16_to_bf16_bytes(raw);
429 &converted
430 } else {
431 raw
432 };
433
434 let ptr = match gpu.alloc(meta.len) {
435 Ok(p) => {
436 gpu.copy_h2d(src, p)?;
437 p
438 }
439 Err(_) => {
440 if !*offload_logged {
441 tracing::warn!(
442 "GPU alloc failed for {} ({} bytes) — switching to managed (UVM) memory",
443 meta.name,
444 meta.len
445 );
446 *offload_logged = true;
447 }
448 let p = gpu.alloc_managed(meta.len)?;
449 unsafe {
450 std::ptr::copy_nonoverlapping(src.as_ptr(), p.0 as *mut u8, meta.len);
451 }
452 p
453 }
454 };
455
456 out.insert(
457 meta.name.clone(),
458 WeightTensor {
459 ptr,
460 shape: meta.shape.clone(),
461 dtype: meta.dtype,
462 },
463 );
464 }
465
466 reader_handle
467 .join()
468 .map_err(|_| anyhow::anyhow!("reader thread panicked"))?;
469
470 drop(direct_file);
474 evict_page_cache(&buffered_file);
475 drop(buffered_file);
476 Ok(())
477}
478
479#[cfg(target_os = "linux")]
480fn advise_prefetch_shard(file: &File, shard_path: &Path, file_size: u64) {
481 use std::os::unix::io::AsRawFd;
482
483 let fd = file.as_raw_fd();
484 let seq_rc = unsafe { libc::posix_fadvise(fd, 0, 0, libc::POSIX_FADV_SEQUENTIAL) };
485 let willneed_rc = unsafe { libc::posix_fadvise(fd, 0, 0, libc::POSIX_FADV_WILLNEED) };
486 if seq_rc == 0 && willneed_rc == 0 {
487 tracing::info!(
488 " NFS/shard prefetch requested for {} ({:.2} GB)",
489 shard_path.display(),
490 file_size as f64 / (1024.0 * 1024.0 * 1024.0)
491 );
492 } else {
493 tracing::warn!(
494 " NFS/shard prefetch hint failed for {}: sequential_rc={}, willneed_rc={}",
495 shard_path.display(),
496 seq_rc,
497 willneed_rc
498 );
499 }
500}
501
502#[cfg(not(target_os = "linux"))]
503fn advise_prefetch_shard(_file: &File, _shard_path: &Path, _file_size: u64) {}
504
505#[cfg(test)]
506mod skip_vision_tests {
507 use super::{FastSafetensorsLoader, is_vision_tensor};
508
509 fn loader(skip_vision: bool, ep: usize) -> FastSafetensorsLoader {
510 let mut l = FastSafetensorsLoader::with_ep(0, ep, 288);
511 l.skip_vision = skip_vision;
512 l
513 }
514
515 #[test]
516 fn vision_names_are_recognised() {
517 assert!(is_vision_tensor("model.visual.blocks.0.attn.proj.weight"));
518 assert!(is_vision_tensor(
519 "model.vision_tower.encoder.layer.0.weight"
520 ));
521 assert!(is_vision_tensor("visual.merger.proj.weight"));
522 assert!(!is_vision_tensor(
523 "model.language_model.layers.45.eh_proj.weight"
524 ));
525 assert!(!is_vision_tensor(
527 "model.language_model.layers.3.mlp.revision.weight"
528 ));
529 }
530
531 #[test]
532 fn skip_vision_drops_only_the_tower() {
533 let l = loader(true, 2);
534 assert!(l.should_skip_tensor("model.visual.blocks.0.attn.proj.weight"));
535 assert!(!l.should_skip_tensor("model.language_model.layers.45.eh_proj.weight"));
536 assert!(!l.should_skip_tensor("lm_head.weight"));
537 }
538
539 #[test]
540 fn skip_vision_applies_without_ep() {
541 let l = loader(true, 1);
543 assert!(l.should_skip_tensor("model.visual.blocks.0.attn.proj.weight"));
544 assert!(!l.should_skip_tensor("model.layers.0.self_attn.q_proj.weight"));
545 }
546
547 #[test]
548 fn default_loader_keeps_the_tower() {
549 let l = loader(false, 2);
550 assert!(!l.should_skip_tensor("model.visual.blocks.0.attn.proj.weight"));
551 assert!(!FastSafetensorsLoader::new().skip_vision);
552 }
553
554 #[test]
555 fn ep_expert_filtering_is_unchanged_by_the_vision_rule() {
556 let l = loader(true, 2); assert!(
558 !l.should_skip_tensor("model.language_model.layers.4.mlp.experts.7.up_proj.weight")
559 );
560 assert!(
561 l.should_skip_tensor("model.language_model.layers.4.mlp.experts.200.up_proj.weight")
562 );
563 }
564}