Skip to main content

combs_formats/
safetensors.rs

1//! HuggingFace safetensors adapter: `config.json` + `model.safetensors`
2//! (single-file or sharded with `model.safetensors.index.json`), mmap-backed.
3
4use std::collections::HashMap;
5use std::path::{Path, PathBuf};
6
7use memmap2::Mmap;
8use safetensors::SafeTensors;
9
10use crate::metadata::ModelMetadata;
11use crate::source::{ModelSource, SamplerConfig, TensorDtype, TensorReader};
12use crate::tokenizer::TokenizerSpec;
13use crate::{FormatError, Result};
14
15/// One memory-mapped safetensors shard.
16struct Shard {
17    mmap: Mmap,
18}
19
20/// Lightweight per-tensor index entry (populated once at load).
21struct TensorEntry {
22    shard: usize,
23    dtype: TensorDtype,
24    shape: Vec<usize>,
25}
26
27/// [`ModelSource`] over a HuggingFace-format directory:
28///
29/// ```text
30/// <dir>/config.json                   (required)
31/// <dir>/generation_config.json        (optional)
32/// <dir>/tokenizer.json                (required by the runtime)
33/// <dir>/tokenizer_config.json         (optional; chat special tokens)
34/// <dir>/model.safetensors             (single-file), or
35/// <dir>/model.safetensors.index.json  (sharded)
36/// ```
37///
38/// Files are memory-mapped; [`ModelSource::open_tensor`] returns zero-copy
39/// views into the mapping.
40pub struct SafetensorsSource {
41    metadata: ModelMetadata,
42    tokenizer: TokenizerSpec,
43    sampler: Option<SamplerConfig>,
44    shards: Vec<Shard>,
45    index: HashMap<String, TensorEntry>,
46}
47
48fn read_json(path: &Path, required: bool) -> Result<Option<serde_json::Value>> {
49    match std::fs::read_to_string(path) {
50        Ok(text) => serde_json::from_str(&text)
51            .map(Some)
52            .map_err(|source| FormatError::Json {
53                context: path.display().to_string(),
54                source,
55            }),
56        Err(e) if e.kind() == std::io::ErrorKind::NotFound && !required => Ok(None),
57        Err(e) if e.kind() == std::io::ErrorKind::NotFound => {
58            Err(FormatError::MissingFile(path.display().to_string()))
59        }
60        Err(e) => Err(FormatError::Io(e)),
61    }
62}
63
64fn convert_dtype(
65    name: &str,
66    dtype: safetensors::Dtype,
67) -> Result<TensorDtype> {
68    match dtype {
69        safetensors::Dtype::F32 => Ok(TensorDtype::F32),
70        safetensors::Dtype::F16 => Ok(TensorDtype::F16),
71        safetensors::Dtype::BF16 => Ok(TensorDtype::BF16),
72        safetensors::Dtype::U8 => Ok(TensorDtype::U8),
73        other => Err(FormatError::UnsupportedDtype {
74            tensor: name.to_string(),
75            dtype: format!("{other:?}"),
76        }),
77    }
78}
79
80impl SafetensorsSource {
81    /// Opens a model directory (see type docs for the expected layout).
82    pub fn load(dir: impl AsRef<Path>) -> Result<Self> {
83        let dir = dir.as_ref();
84
85        // --- config + metadata -------------------------------------------------
86        let config = read_json(&dir.join("config.json"), true)?.expect("required");
87        let generation_config = read_json(&dir.join("generation_config.json"), false)?;
88        let metadata =
89            ModelMetadata::from_hf_config(&config, generation_config.as_ref())?;
90
91        // --- tokenizer ----------------------------------------------------------
92        let tokenizer_json = dir.join("tokenizer.json");
93        if !tokenizer_json.exists() {
94            return Err(FormatError::MissingFile(
95                tokenizer_json.display().to_string(),
96            ));
97        }
98        let mut added_tokens = HashMap::new();
99        let mut chat_template = None;
100        if let Some(tc) = read_json(&dir.join("tokenizer_config.json"), false)? {
101            if let Some(map) = tc.get("added_tokens_decoder").and_then(|v| v.as_object()) {
102                for (id, entry) in map {
103                    if let (Ok(id), Some(content)) = (
104                        id.parse::<u32>(),
105                        entry.get("content").and_then(|c| c.as_str()),
106                    ) {
107                        added_tokens.insert(id, content.to_string());
108                    }
109                }
110            }
111            chat_template = tc
112                .get("chat_template")
113                .and_then(|v| v.as_str())
114                .map(|s| s.to_string());
115        }
116        let tokenizer = TokenizerSpec {
117            tokenizer_json,
118            added_tokens,
119            chat_template,
120        };
121
122        // --- sampler defaults ----------------------------------------------------
123        let sampler = generation_config.as_ref().map(|gc| SamplerConfig {
124            temperature: gc
125                .get("temperature")
126                .and_then(|v| v.as_f64())
127                .map(|v| v as f32),
128            top_p: gc.get("top_p").and_then(|v| v.as_f64()).map(|v| v as f32),
129            top_k: gc
130                .get("top_k")
131                .and_then(|v| v.as_u64())
132                .map(|v| v as usize),
133            repetition_penalty: gc
134                .get("repetition_penalty")
135                .and_then(|v| v.as_f64())
136                .map(|v| v as f32),
137            max_new_tokens: gc
138                .get("max_new_tokens")
139                .or_else(|| gc.get("max_length"))
140                .and_then(|v| v.as_u64())
141                .map(|v| v as usize),
142        });
143
144        // --- weight shards ---------------------------------------------------------
145        let index_json = dir.join("model.safetensors.index.json");
146        let single = dir.join("model.safetensors");
147        let shard_files: Vec<PathBuf> = if index_json.exists() {
148            let idx = read_json(&index_json, true)?.expect("required");
149            let weight_map = idx
150                .get("weight_map")
151                .and_then(|v| v.as_object())
152                .ok_or_else(|| FormatError::MissingField("weight_map".to_string()))?;
153            let mut files: Vec<PathBuf> = weight_map
154                .values()
155                .filter_map(|v| v.as_str())
156                .map(|f| dir.join(f))
157                .collect();
158            files.sort();
159            files.dedup();
160            files
161        } else if single.exists() {
162            vec![single]
163        } else {
164            return Err(FormatError::MissingFile(format!(
165                "{single:?} (or model.safetensors.index.json)"
166            )));
167        };
168
169        let mut shards = Vec::with_capacity(shard_files.len());
170        let mut index = HashMap::new();
171        for (shard_idx, file) in shard_files.iter().enumerate() {
172            let f = std::fs::File::open(file)?;
173            // SAFETY: the file is opened read-only and never mutated by us;
174            // external mutation of an mmap'd model file is out of scope.
175            let mmap = unsafe { Mmap::map(&f)? };
176            let st = SafeTensors::deserialize(&mmap).map_err(|e| {
177                FormatError::Safetensors(format!("{}: {e}", file.display()))
178            })?;
179            for (name, view) in st.tensors() {
180                index.insert(
181                    name.clone(),
182                    TensorEntry {
183                        shard: shard_idx,
184                        dtype: convert_dtype(&name, view.dtype())?,
185                        shape: view.shape().to_vec(),
186                    },
187                );
188            }
189            shards.push(Shard { mmap });
190        }
191
192        Ok(SafetensorsSource {
193            metadata,
194            tokenizer,
195            sampler,
196            shards,
197            index,
198        })
199    }
200}
201
202impl ModelSource for SafetensorsSource {
203    fn metadata(&self) -> &ModelMetadata {
204        &self.metadata
205    }
206
207    fn tensor_names(&self) -> Vec<String> {
208        let mut names: Vec<String> = self.index.keys().cloned().collect();
209        names.sort();
210        names
211    }
212
213    fn open_tensor(&self, name: &str) -> Result<TensorReader<'_>> {
214        let entry = self
215            .index
216            .get(name)
217            .ok_or_else(|| FormatError::TensorNotFound(name.to_string()))?;
218        let shard = &self.shards[entry.shard];
219        let st = SafeTensors::deserialize(&shard.mmap)
220            .map_err(|e| FormatError::Safetensors(e.to_string()))?;
221        let view = st
222            .tensor(name)
223            .map_err(|e| FormatError::Safetensors(e.to_string()))?;
224        Ok(TensorReader::new(
225            name.to_string(),
226            entry.shape.clone(),
227            entry.dtype,
228            view.data(),
229        ))
230    }
231
232    fn tokenizer(&self) -> Result<TokenizerSpec> {
233        Ok(self.tokenizer.clone())
234    }
235
236    fn sampler_defaults(&self) -> Option<SamplerConfig> {
237        self.sampler.clone()
238    }
239}
240
241#[cfg(test)]
242mod tests {
243    use super::*;
244
245    /// Builds a tiny synthetic HF model dir and checks the adapter surface.
246    #[test]
247    fn lists_and_reads_tensors() {
248        let dir = tempfile::tempdir().unwrap();
249        let root = dir.path();
250
251        std::fs::write(
252            root.join("config.json"),
253            serde_json::to_string(&serde_json::json!({
254                "model_type": "llama",
255                "hidden_size": 8,
256                "intermediate_size": 16,
257                "num_hidden_layers": 1,
258                "num_attention_heads": 2,
259                "num_key_value_heads": 1,
260                "vocab_size": 32,
261                "max_position_embeddings": 128,
262                "rope_theta": 10000,
263                "rms_norm_eps": 1e-5,
264                "tie_word_embeddings": true,
265                "eos_token_id": 0,
266                "bos_token_id": 0
267            }))
268            .unwrap(),
269        )
270        .unwrap();
271        std::fs::write(root.join("tokenizer.json"), "{}").unwrap();
272        std::fs::write(
273            root.join("tokenizer_config.json"),
274            serde_json::to_string(&serde_json::json!({
275                "added_tokens_decoder": {
276                    "0": {"content": "<|endoftext|>", "special": true},
277                    "2": {"content": "<|im_end|>", "special": true}
278                }
279            }))
280            .unwrap(),
281        )
282        .unwrap();
283
284        // Two tensors: one F32, one BF16.
285        let w1_bytes: Vec<u8> = (0..16)
286            .flat_map(|i| (i as f32).to_le_bytes())
287            .collect();
288        let w2_bytes: Vec<u8> = (0..8)
289            .flat_map(|i| half::bf16::from_f32(i as f32 * 0.5).to_le_bytes())
290            .collect();
291        let v1 = safetensors::tensor::TensorView::new(
292            safetensors::Dtype::F32,
293            vec![4, 4],
294            &w1_bytes,
295        )
296        .unwrap();
297        let v2 = safetensors::tensor::TensorView::new(
298            safetensors::Dtype::BF16,
299            vec![8],
300            &w2_bytes,
301        )
302        .unwrap();
303        safetensors::serialize_to_file(
304            vec![("a.weight", v1), ("b.weight", v2)],
305            None,
306            root.join("model.safetensors").as_path(),
307        )
308        .unwrap();
309
310        let src = SafetensorsSource::load(root).unwrap();
311        assert_eq!(src.metadata().architecture, "llama");
312        assert_eq!(src.metadata().head_dim, 4);
313        assert_eq!(src.tensor_names(), vec!["a.weight", "b.weight"]);
314
315        let a = src.open_tensor("a.weight").unwrap();
316        assert_eq!(a.shape(), &[4, 4]);
317        assert_eq!(a.dtype(), TensorDtype::F32);
318        let data = a.load_data().unwrap();
319        let vals: Vec<f32> = data.to_vec().unwrap();
320        assert_eq!(vals[15], 15.0);
321
322        let b = src.open_tensor("b.weight").unwrap();
323        assert_eq!(b.dtype(), TensorDtype::BF16);
324        let vals: Vec<f32> = b.load_data().unwrap().to_vec().unwrap();
325        assert!((vals[3] - 1.5).abs() < 1e-3);
326
327        let tok = src.tokenizer().unwrap();
328        assert_eq!(tok.special_token_id("<|im_end|>"), Some(2));
329
330        assert!(src.open_tensor("missing").is_err());
331    }
332}