ferrum_quantization/loader.rs
1//! `WeightLoader` trait — unified interface for loading tensor/linear weights
2//! into a specific backend.
3//!
4//! Implementations (landing in Phase B):
5//! - `SafeTensorsLoader` — reads `.safetensors` files, returns `DenseLinear`
6//! unless `quantize_config.json` indicates GPTQ/AWQ, in which case it
7//! returns `GptqLinear` / `AwqLinear`.
8//! - `GgufLoader` — reads `.gguf` files, returns `GgufLinear`.
9//!
10//! The trait is generic over `B: Backend` so the loader can materialise
11//! tensors directly into backend-native buffers (zero-copy on Apple Silicon
12//! shared memory, dtoh/htod for CUDA, etc.).
13
14use ferrum_kernels::{backend::Backend, MarlinExpertStack};
15use ferrum_types::{FerrumError, Result};
16
17use crate::config::QuantConfig;
18use crate::traits::Linear;
19
20pub trait WeightLoader<B: Backend>: Send + Sync {
21 /// Load a single tensor by fully qualified name
22 /// (e.g. `"model.embed_tokens.weight"`).
23 fn load_tensor(&self, name: &str) -> Result<B::Buffer>;
24
25 /// Load a projection as a `Linear<B>`. The concrete implementation
26 /// (DenseLinear / GptqLinear / AwqLinear / GgufLinear) depends on the
27 /// loader's file format and quant config.
28 ///
29 /// `name` is the module path without the `.weight` suffix, e.g.
30 /// `"model.layers.0.self_attn.qkv_proj"`.
31 fn load_linear(&self, name: &str) -> Result<Box<dyn Linear<B>>>;
32
33 /// Whether a tensor with this name exists in the source.
34 fn has_tensor(&self, name: &str) -> bool;
35
36 /// Quantization metadata (parsed from `quantize_config.json` or a GGUF header).
37 /// `None` means the source is dense.
38 fn quant_config(&self) -> Option<&QuantConfig>;
39
40 /// Load per-expert GPTQ projections into one backend-native stacked expert
41 /// tile. Backends/loaders that do not expose native stacked GPTQ return an
42 /// explicit unsupported error.
43 fn load_stacked_gptq_experts(
44 &self,
45 expert_prefix_fmt: &str,
46 num_experts: usize,
47 proj_names: &[&str],
48 ) -> Result<(std::sync::Arc<dyn MarlinExpertStack<B>>, usize, usize)> {
49 let _ = (expert_prefix_fmt, num_experts, proj_names);
50 Err(FerrumError::unsupported(
51 "load_stacked_gptq_experts not implemented for this weight loader",
52 ))
53 }
54}
55
56/// Adapter that prepends a fixed prefix to every tensor name before
57/// delegating to an underlying loader.
58///
59/// Use case: a single safetensors file contains a sub-model (e.g.
60/// Qwen3-TTS stores the Talker LM under `talker.model.*`) and we want
61/// to reuse a backbone loader like `LlamaFamilyModel::new` that
62/// expects bare `model.*` names. Wrapping with
63/// `PrefixedLoader { inner, prefix: "talker." }` lets the backbone
64/// code stay prefix-agnostic.
65pub struct PrefixedLoader<'a, B: Backend> {
66 inner: &'a dyn WeightLoader<B>,
67 prefix: String,
68}
69
70impl<'a, B: Backend> PrefixedLoader<'a, B> {
71 pub fn new(inner: &'a dyn WeightLoader<B>, prefix: impl Into<String>) -> Self {
72 Self {
73 inner,
74 prefix: prefix.into(),
75 }
76 }
77}
78
79impl<'a, B: Backend> WeightLoader<B> for PrefixedLoader<'a, B> {
80 fn load_tensor(&self, name: &str) -> Result<B::Buffer> {
81 self.inner.load_tensor(&format!("{}{}", self.prefix, name))
82 }
83
84 fn load_linear(&self, name: &str) -> Result<Box<dyn Linear<B>>> {
85 self.inner.load_linear(&format!("{}{}", self.prefix, name))
86 }
87
88 fn has_tensor(&self, name: &str) -> bool {
89 self.inner.has_tensor(&format!("{}{}", self.prefix, name))
90 }
91
92 fn quant_config(&self) -> Option<&QuantConfig> {
93 self.inner.quant_config()
94 }
95
96 fn load_stacked_gptq_experts(
97 &self,
98 expert_prefix_fmt: &str,
99 num_experts: usize,
100 proj_names: &[&str],
101 ) -> Result<(std::sync::Arc<dyn MarlinExpertStack<B>>, usize, usize)> {
102 self.inner.load_stacked_gptq_experts(
103 &format!("{}{}", self.prefix, expert_prefix_fmt),
104 num_experts,
105 proj_names,
106 )
107 }
108}