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/// Standard layout:
30///
31/// ```text
32/// <dir>/config.json                   (required)
33/// <dir>/generation_config.json        (optional)
34/// <dir>/tokenizer.json                (required by the runtime)
35/// <dir>/tokenizer_config.json         (optional; chat special tokens)
36/// <dir>/model.safetensors             (single-file), or
37/// <dir>/model.safetensors.index.json  (sharded)
38/// ```
39///
40/// For diffusion sub-dirs (`unet/`, `vae/`, `text_encoder/`) use
41/// [`SafetensorsSource::load_weights_only`], which tolerates
42/// `diffusion_pytorch_model.safetensors` (or `model.safetensors`) and no
43/// tokenizer/config.
44///
45/// Files are memory-mapped; [`ModelSource::open_tensor`] returns zero-copy
46/// views into the mapping.
47pub struct SafetensorsSource {
48    metadata: ModelMetadata,
49    tokenizer: TokenizerSpec,
50    sampler: Option<SamplerConfig>,
51    shards: Vec<Shard>,
52    index: HashMap<String, TensorEntry>,
53}
54
55fn read_json(path: &Path, required: bool) -> Result<Option<serde_json::Value>> {
56    match std::fs::read_to_string(path) {
57        Ok(text) => serde_json::from_str(&text)
58            .map(Some)
59            .map_err(|source| FormatError::Json {
60                context: path.display().to_string(),
61                source,
62            }),
63        Err(e) if e.kind() == std::io::ErrorKind::NotFound && !required => Ok(None),
64        Err(e) if e.kind() == std::io::ErrorKind::NotFound => {
65            Err(FormatError::MissingFile(path.display().to_string()))
66        }
67        Err(e) => Err(FormatError::Io(e)),
68    }
69}
70
71fn convert_dtype(
72    _name: &str,
73    dtype: safetensors::Dtype,
74) -> Result<TensorDtype> {
75    match dtype {
76        safetensors::Dtype::F64 => Ok(TensorDtype::F64),
77        safetensors::Dtype::F32 => Ok(TensorDtype::F32),
78        safetensors::Dtype::F16 => Ok(TensorDtype::F16),
79        safetensors::Dtype::BF16 => Ok(TensorDtype::BF16),
80        safetensors::Dtype::I64 => Ok(TensorDtype::I64),
81        safetensors::Dtype::I32 => Ok(TensorDtype::I32),
82        safetensors::Dtype::I16 => Ok(TensorDtype::I16),
83        safetensors::Dtype::I8 => Ok(TensorDtype::I8),
84        safetensors::Dtype::U64 => Ok(TensorDtype::U64),
85        safetensors::Dtype::U32 => Ok(TensorDtype::U32),
86        safetensors::Dtype::U16 => Ok(TensorDtype::U16),
87        safetensors::Dtype::U8 => Ok(TensorDtype::U8),
88        safetensors::Dtype::BOOL => Ok(TensorDtype::Bool),
89        other => Err(FormatError::UnsupportedDtype {
90            tensor: _name.to_string(),
91            dtype: format!("{other:?}"),
92        }),
93    }
94}
95
96impl SafetensorsSource {
97    /// Opens a model directory (see type docs for the expected layout).
98    pub fn load(dir: impl AsRef<Path>) -> Result<Self> {
99        let dir = dir.as_ref();
100
101        // --- config + metadata -------------------------------------------------
102        let config = read_json(&dir.join("config.json"), true)?.expect("required");
103        let generation_config = read_json(&dir.join("generation_config.json"), false)?;
104        let metadata =
105            ModelMetadata::from_hf_config(&config, generation_config.as_ref())?;
106
107        // --- tokenizer ----------------------------------------------------------
108        let tokenizer_json = dir.join("tokenizer.json");
109        if !tokenizer_json.exists() {
110            return Err(FormatError::MissingFile(
111                tokenizer_json.display().to_string(),
112            ));
113        }
114        let mut added_tokens = HashMap::new();
115        let mut chat_template = None;
116        let mut add_bos = None;
117        if let Some(tc) = read_json(&dir.join("tokenizer_config.json"), false)? {
118            if let Some(map) = tc.get("added_tokens_decoder").and_then(|v| v.as_object()) {
119                for (id, entry) in map {
120                    if let (Ok(id), Some(content)) = (
121                        id.parse::<u32>(),
122                        entry.get("content").and_then(|c| c.as_str()),
123                    ) {
124                        added_tokens.insert(id, content.to_string());
125                    }
126                }
127            }
128            chat_template = tc
129                .get("chat_template")
130                .and_then(|v| v.as_str())
131                .map(|s| s.to_string());
132            add_bos = tc.get("add_bos_token").and_then(|v| v.as_bool());
133        }
134        let tokenizer = TokenizerSpec {
135            tokenizer_json,
136            added_tokens,
137            chat_template,
138            add_bos,
139        };
140
141        // --- sampler defaults ----------------------------------------------------
142        let sampler = generation_config.as_ref().map(|gc| SamplerConfig {
143            temperature: gc
144                .get("temperature")
145                .and_then(|v| v.as_f64())
146                .map(|v| v as f32),
147            top_p: gc.get("top_p").and_then(|v| v.as_f64()).map(|v| v as f32),
148            top_k: gc
149                .get("top_k")
150                .and_then(|v| v.as_u64())
151                .map(|v| v as usize),
152            repetition_penalty: gc
153                .get("repetition_penalty")
154                .and_then(|v| v.as_f64())
155                .map(|v| v as f32),
156            max_new_tokens: gc
157                .get("max_new_tokens")
158                .or_else(|| gc.get("max_length"))
159                .and_then(|v| v.as_u64())
160                .map(|v| v as usize),
161        });
162
163        // --- weight shards ---------------------------------------------------------
164        let shard_files = Self::collect_shard_files(dir, &["model.safetensors"])?;
165        let (shards, index) = Self::mmap_shards(&shard_files)?;
166
167        Ok(SafetensorsSource {
168            metadata,
169            tokenizer,
170            sampler,
171            shards,
172            index,
173        })
174    }
175
176    /// Opens a directory purely for weight loading. No `config.json` or
177    /// `tokenizer.json` is required; `metadata()` returns a placeholder and
178    /// `tokenizer()` errors. Accepts diffusion-style
179    /// `diffusion_pytorch_model.safetensors` as well as `model.safetensors`.
180    pub fn load_weights_only(dir: impl AsRef<Path>, architecture: &str) -> Result<Self> {
181        let dir = dir.as_ref();
182        let shard_files = Self::collect_shard_files(
183            dir,
184            &[
185                "diffusion_pytorch_model.safetensors",
186                "model.safetensors",
187            ],
188        )?;
189        let (shards, index) = Self::mmap_shards(&shard_files)?;
190        Ok(SafetensorsSource {
191            metadata: ModelMetadata::diffusion_placeholder(architecture),
192            tokenizer: TokenizerSpec::placeholder(),
193            sampler: None,
194            shards,
195            index,
196        })
197    }
198
199    fn collect_shard_files(dir: &Path, base_names: &[&str]) -> Result<Vec<PathBuf>> {
200        for base in base_names {
201            let index_json = dir.join(format!("{base}.index.json"));
202            if index_json.exists() {
203                let idx = read_json(&index_json, true)?.expect("required");
204                let weight_map = idx
205                    .get("weight_map")
206                    .and_then(|v| v.as_object())
207                    .ok_or_else(|| FormatError::MissingField("weight_map".to_string()))?;
208                let mut files: Vec<PathBuf> = weight_map
209                    .values()
210                    .filter_map(|v| v.as_str())
211                    .map(|f| dir.join(f))
212                    .collect();
213                files.sort();
214                files.dedup();
215                return Ok(files);
216            }
217            let single = dir.join(base);
218            if single.exists() {
219                return Ok(vec![single]);
220            }
221        }
222        Err(FormatError::MissingFile(format!(
223            "{:?} (or .index.json)",
224            base_names
225                .iter()
226                .map(|b| dir.join(b))
227                .collect::<Vec<_>>()
228        )))
229    }
230
231    fn mmap_shards(shard_files: &[PathBuf]) -> Result<(Vec<Shard>, HashMap<String, TensorEntry>)> {
232        let mut shards = Vec::with_capacity(shard_files.len());
233        let mut index = HashMap::new();
234        for (shard_idx, file) in shard_files.iter().enumerate() {
235            let f = std::fs::File::open(file)?;
236            // SAFETY: the file is opened read-only and never mutated by us;
237            // external mutation of an mmap'd model file is out of scope.
238            let mmap = unsafe { Mmap::map(&f)? };
239            let st = SafeTensors::deserialize(&mmap).map_err(|e| {
240                FormatError::Safetensors(format!("{}: {e}", file.display()))
241            })?;
242            for (name, view) in st.tensors() {
243                index.insert(
244                    name.clone(),
245                    TensorEntry {
246                        shard: shard_idx,
247                        dtype: convert_dtype(&name, view.dtype())?,
248                        shape: view.shape().to_vec(),
249                    },
250                );
251            }
252            shards.push(Shard { mmap });
253        }
254        Ok((shards, index))
255    }
256}
257
258impl ModelSource for SafetensorsSource {
259    fn metadata(&self) -> &ModelMetadata {
260        &self.metadata
261    }
262
263    fn tensor_names(&self) -> Vec<String> {
264        let mut names: Vec<String> = self.index.keys().cloned().collect();
265        names.sort();
266        names
267    }
268
269    fn open_tensor(&self, name: &str) -> Result<TensorReader<'_>> {
270        let entry = self
271            .index
272            .get(name)
273            .ok_or_else(|| FormatError::TensorNotFound(name.to_string()))?;
274        let shard = &self.shards[entry.shard];
275        let st = SafeTensors::deserialize(&shard.mmap)
276            .map_err(|e| FormatError::Safetensors(e.to_string()))?;
277        let view = st
278            .tensor(name)
279            .map_err(|e| FormatError::Safetensors(e.to_string()))?;
280        Ok(TensorReader::new(
281            name.to_string(),
282            entry.shape.clone(),
283            entry.dtype,
284            view.data(),
285        ))
286    }
287
288    fn tokenizer(&self) -> Result<TokenizerSpec> {
289        if self.tokenizer.tokenizer_json.as_os_str().is_empty() {
290            return Err(FormatError::MissingFile(
291                "tokenizer.json (weights-only source has no tokenizer)".to_string(),
292            ));
293        }
294        Ok(self.tokenizer.clone())
295    }
296
297    fn sampler_defaults(&self) -> Option<SamplerConfig> {
298        self.sampler.clone()
299    }
300}
301
302#[cfg(test)]
303mod tests {
304    use super::*;
305
306    /// Builds a tiny synthetic HF model dir and checks the adapter surface.
307    #[test]
308    fn lists_and_reads_tensors() {
309        let dir = tempfile::tempdir().unwrap();
310        let root = dir.path();
311
312        std::fs::write(
313            root.join("config.json"),
314            serde_json::to_string(&serde_json::json!({
315                "model_type": "llama",
316                "hidden_size": 8,
317                "intermediate_size": 16,
318                "num_hidden_layers": 1,
319                "num_attention_heads": 2,
320                "num_key_value_heads": 1,
321                "vocab_size": 32,
322                "max_position_embeddings": 128,
323                "rope_theta": 10000,
324                "rms_norm_eps": 1e-5,
325                "tie_word_embeddings": true,
326                "eos_token_id": 0,
327                "bos_token_id": 0
328            }))
329            .unwrap(),
330        )
331        .unwrap();
332        std::fs::write(root.join("tokenizer.json"), "{}").unwrap();
333        std::fs::write(
334            root.join("tokenizer_config.json"),
335            serde_json::to_string(&serde_json::json!({
336                "added_tokens_decoder": {
337                    "0": {"content": "<|endoftext|>", "special": true},
338                    "2": {"content": "<|im_end|>", "special": true}
339                }
340            }))
341            .unwrap(),
342        )
343        .unwrap();
344
345        // Two tensors: one F32, one BF16.
346        let w1_bytes: Vec<u8> = (0..16)
347            .flat_map(|i| (i as f32).to_le_bytes())
348            .collect();
349        let w2_bytes: Vec<u8> = (0..8)
350            .flat_map(|i| half::bf16::from_f32(i as f32 * 0.5).to_le_bytes())
351            .collect();
352        let v1 = safetensors::tensor::TensorView::new(
353            safetensors::Dtype::F32,
354            vec![4, 4],
355            &w1_bytes,
356        )
357        .unwrap();
358        let v2 = safetensors::tensor::TensorView::new(
359            safetensors::Dtype::BF16,
360            vec![8],
361            &w2_bytes,
362        )
363        .unwrap();
364
365        // Add an I64 position_ids buffer like some CLIP text encoders ship.
366        let pos_ids_bytes: Vec<u8> = (0..8i64)
367            .flat_map(|i| i.to_le_bytes())
368            .collect();
369        let v_pos = safetensors::tensor::TensorView::new(
370            safetensors::Dtype::I64,
371            vec![8],
372            &pos_ids_bytes,
373        )
374        .unwrap();
375
376        safetensors::serialize_to_file(
377            vec![("a.weight", v1), ("b.weight", v2), ("text_model.embeddings.position_ids", v_pos)],
378            None,
379            root.join("model.safetensors").as_path(),
380        )
381        .unwrap();
382
383        let src = SafetensorsSource::load(root).unwrap();
384        assert_eq!(src.metadata().architecture, "llama");
385        assert_eq!(src.metadata().head_dim, 4);
386        assert_eq!(
387            src.tensor_names(),
388            vec!["a.weight", "b.weight", "text_model.embeddings.position_ids"]
389        );
390
391        let pos = src.open_tensor("text_model.embeddings.position_ids").unwrap();
392        assert_eq!(pos.dtype(), TensorDtype::I64);
393        let pos_vals: Vec<f32> = pos.load_data().unwrap().to_vec().unwrap();
394        assert_eq!(pos_vals[7], 7.0);
395
396        let a = src.open_tensor("a.weight").unwrap();
397        assert_eq!(a.shape(), &[4, 4]);
398        assert_eq!(a.dtype(), TensorDtype::F32);
399        let data = a.load_data().unwrap();
400        let vals: Vec<f32> = data.to_vec().unwrap();
401        assert_eq!(vals[15], 15.0);
402
403        let b = src.open_tensor("b.weight").unwrap();
404        assert_eq!(b.dtype(), TensorDtype::BF16);
405        let vals: Vec<f32> = b.load_data().unwrap().to_vec().unwrap();
406        assert!((vals[3] - 1.5).abs() < 1e-3);
407
408        let tok = src.tokenizer().unwrap();
409        assert_eq!(tok.special_token_id("<|im_end|>"), Some(2));
410
411        assert!(src.open_tensor("missing").is_err());
412    }
413
414    /// `load_weights_only` accepts diffusion-style filenames and no config.
415    #[test]
416    fn loads_weights_only_without_config_or_tokenizer() {
417        let dir = tempfile::tempdir().unwrap();
418        let root = dir.path();
419
420        let bytes: Vec<u8> = (0..16)
421            .flat_map(|i| (i as f32).to_le_bytes())
422            .collect();
423        let t = safetensors::tensor::TensorView::new(
424            safetensors::Dtype::F32,
425            vec![4, 4],
426            &bytes,
427        )
428        .unwrap();
429        safetensors::serialize_to_file(
430            vec![("conv_in.weight", t)],
431            None,
432            root.join("diffusion_pytorch_model.safetensors")
433                .as_path(),
434        )
435        .unwrap();
436
437        let src =
438            SafetensorsSource::load_weights_only(root, "stable-diffusion-unet").unwrap();
439        assert_eq!(src.metadata().architecture, "stable-diffusion-unet");
440        assert_eq!(src.tensor_names(), vec!["conv_in.weight"]);
441        assert!(src.tokenizer().is_err());
442    }
443}