Skip to main content

kopitiam_loader/
safetensors.rs

1//! Native SafeTensors parsing.
2//!
3//! SafeTensors (Hugging Face's format) is documented informally by its
4//! reference implementation at `crates/kopitiam-ai/vendor/safetensors/`;
5//! this module is an original Rust implementation studied against that
6//! reference (`safetensors/src/tensor.rs`), not a copy of it.
7//!
8//! # File layout
9//!
10//! ```text
11//! header_len(u64, little-endian)
12//! header_json(header_len bytes) -- a JSON object:
13//!   {
14//!     "__metadata__": { "<key>": "<value>", ... },   // optional, strings only
15//!     "<tensor name>": {
16//!       "dtype": "F32" | "F16" | "BF16" | "I8" | "I32" | ...,
17//!       "shape": [dim0, dim1, ...],
18//!       "data_offsets": [start, end]                 // relative to the byte after the header
19//!     },
20//!     ...
21//!   }
22//! <raw tensor bytes>
23//! ```
24//!
25//! # Dimension order — no trap here
26//!
27//! Unlike GGUF (see [`crate::gguf`]'s module docs), SafeTensors' `shape` is
28//! already outermost-first, row-major — the same convention NumPy, PyTorch
29//! and [`kopitiam_core::Shape`] all use. No reversal is needed; `shape` is
30//! read directly into a [`Shape`].
31//!
32//! # Dtype coverage
33//!
34//! SafeTensors supports more element types than [`kopitiam_core::DType`]
35//! currently models (`U8`, `U16`, `U32`, `I16`, `I64`, `U64`, `F64`,
36//! `BOOL`, the `F8_E*` micro-floats). Only the subset `DType` can represent
37//! today — `F32`, `F16`, `BF16`, `I8`, `I32` — is accepted; everything else
38//! returns [`Error::UnsupportedModelFeature`] naming the unsupported dtype
39//! string, rather than being coerced into a same-size type that would
40//! silently reinterpret the bytes' meaning.
41
42use std::collections::BTreeMap;
43use std::path::Path;
44
45use indexmap::IndexMap;
46use kopitiam_core::{DType, Error, Result, Shape};
47use serde::Deserialize;
48
49use crate::byte_source::ByteSource;
50use crate::metadata::{GgufMetadata, GgufValue, ModelMetadata};
51use crate::model::{LoadedModel, ModelLoader, TensorEntry};
52
53const FORMAT: &str = "safetensors";
54const HEADER_LEN_BYTES: usize = 8;
55const METADATA_KEY: &str = "__metadata__";
56
57fn malformed(reason: impl Into<String>) -> Error {
58    Error::MalformedModel { format: FORMAT, reason: reason.into() }
59}
60
61fn unsupported(feature: impl Into<String>) -> Error {
62    Error::UnsupportedModelFeature { format: FORMAT, feature: feature.into() }
63}
64
65/// One tensor's header entry, deserialized directly from its JSON object.
66#[derive(Debug, Deserialize)]
67struct RawTensorInfo {
68    dtype: String,
69    shape: Vec<u64>,
70    data_offsets: (u64, u64),
71}
72
73/// Maps a SafeTensors dtype string to [`DType`]. See the module doc's
74/// "Dtype coverage" section for why this list is deliberately short.
75fn dtype_from_str(s: &str) -> Result<DType> {
76    match s {
77        "F32" => Ok(DType::F32),
78        "F16" => Ok(DType::F16),
79        "BF16" => Ok(DType::BF16),
80        "I8" => Ok(DType::I8),
81        "I32" => Ok(DType::I32),
82        other => Err(unsupported(format!("safetensors dtype {other:?}"))),
83    }
84}
85
86/// Parses a SafeTensors file, already opened as `source`, into a
87/// [`LoadedModel`].
88fn parse(source: ByteSource) -> Result<LoadedModel> {
89    let bytes = source.as_slice();
90
91    let header_len_bytes = bytes.get(..HEADER_LEN_BYTES).ok_or_else(|| {
92        malformed(format!(
93            "file is {} bytes, shorter than the {HEADER_LEN_BYTES}-byte header length prefix",
94            bytes.len()
95        ))
96    })?;
97    let header_len = u64::from_le_bytes(
98        header_len_bytes.try_into().expect("checked slice is exactly 8 bytes"),
99    );
100    let header_len = usize::try_from(header_len)
101        .map_err(|_| malformed(format!("header length {header_len} does not fit in memory")))?;
102
103    // Bounds-check before parsing: a hostile `header_len` (e.g. near
104    // `u64::MAX`, truncated to `usize::MAX` above) fails this `get` rather
105    // than being handed to `serde_json` as a length to allocate around.
106    let data_start = HEADER_LEN_BYTES
107        .checked_add(header_len)
108        .ok_or_else(|| malformed("header end offset overflows"))?;
109    let header_bytes = bytes.get(HEADER_LEN_BYTES..data_start).ok_or_else(|| {
110        malformed(format!(
111            "declared header length {header_len} extends past end of file ({} bytes)",
112            bytes.len()
113        ))
114    })?;
115
116    // `serde_json::Map` (not `IndexMap`) here: the crate does not enable
117    // `indexmap`'s `serde` feature elsewhere, and tensor iteration order
118    // from the JSON header carries no semantic meaning worth threading a
119    // second dependency feature through to preserve.
120    let header: serde_json::Map<String, serde_json::Value> = serde_json::from_slice(header_bytes)
121        .map_err(|e| malformed(format!("header is not valid JSON: {e}")))?;
122
123    let mut raw_metadata = GgufMetadata::new();
124    let mut tensors = IndexMap::new();
125    let file_len = bytes.len();
126
127    for (name, value) in header {
128        if name == METADATA_KEY {
129            let entries: BTreeMap<String, String> = serde_json::from_value(value).map_err(|e| {
130                malformed(format!("{METADATA_KEY} must map strings to strings: {e}"))
131            })?;
132            for (k, v) in entries {
133                raw_metadata.0.insert(k, GgufValue::String(v));
134            }
135            continue;
136        }
137
138        let info: RawTensorInfo = serde_json::from_value(value).map_err(|e| {
139            malformed(format!("tensor {name:?} header entry is malformed: {e}"))
140        })?;
141
142        let dtype = dtype_from_str(&info.dtype)?;
143
144        let mut dims = Vec::with_capacity(info.shape.len());
145        for d in info.shape {
146            let d = usize::try_from(d).map_err(|_| {
147                malformed(format!("tensor {name:?} has a dimension ({d}) too large to represent"))
148            })?;
149            dims.push(d);
150        }
151        let shape = Shape::new(dims);
152        let elem_count = shape.elem_count();
153
154        let expected_len = dtype.storage_bytes(elem_count).ok_or(Error::PartialQuantizedBlock {
155            dtype,
156            count: elem_count,
157            block_size: dtype.block_size(),
158        })?;
159
160        let (start, end) = info.data_offsets;
161        if start > end {
162            return Err(malformed(format!(
163                "tensor {name:?} has data_offsets start ({start}) after end ({end})"
164            )));
165        }
166        let declared_len = end - start;
167        if declared_len != expected_len as u64 {
168            return Err(malformed(format!(
169                "tensor {name:?} declares {declared_len} data bytes but its dtype ({dtype}) and shape ({shape}) need {expected_len}"
170            )));
171        }
172
173        let abs_offset = (data_start as u64).checked_add(start).ok_or_else(|| {
174            malformed(format!("tensor {name:?} data offset overflows a u64"))
175        })?;
176        let abs_end = (data_start as u64).checked_add(end).ok_or_else(|| {
177            malformed(format!("tensor {name:?} data end offset overflows a u64"))
178        })?;
179        if abs_end > file_len as u64 {
180            return Err(malformed(format!(
181                "tensor {name:?} data range [{abs_offset}, {abs_end}) extends past end of file ({file_len} bytes)"
182            )));
183        }
184        let abs_offset = usize::try_from(abs_offset)
185            .map_err(|_| malformed(format!("tensor {name:?} offset does not fit in memory")))?;
186
187        let entry = TensorEntry {
188            name: name.clone(),
189            dtype,
190            shape,
191            offset: abs_offset,
192            len: expected_len,
193        };
194        if tensors.insert(name.clone(), entry).is_some() {
195            return Err(malformed(format!("duplicate tensor name {name:?}")));
196        }
197    }
198
199    let metadata = ModelMetadata {
200        architecture: None,
201        name: None,
202        n_layers: None,
203        n_heads: None,
204        n_kv_heads: None,
205        embedding_length: None,
206        feed_forward_length: None,
207        context_length: None,
208        vocab_size: None,
209        rope_theta: None,
210        rope_dimension_count: None,
211        norm_epsilon: None,
212        quantization_version: None,
213        file_type: None,
214        raw: raw_metadata,
215    };
216
217    Ok(LoadedModel { metadata, tensors, source, format: FORMAT })
218}
219
220/// Parses SafeTensors (Hugging Face) model files. See the module docs for
221/// the format and its dtype coverage.
222pub struct SafeTensorsLoader;
223
224impl ModelLoader for SafeTensorsLoader {
225    fn format_name(&self) -> &'static str {
226        FORMAT
227    }
228
229    fn probe(&self, bytes: &[u8]) -> bool {
230        // SafeTensors has no magic number, so this is a best-effort
231        // heuristic rather than a proof: the header length prefix must be
232        // in-range for the bytes actually sampled, and (when enough bytes
233        // are available to check) the header's first byte must be `{`,
234        // since the header is always a JSON object.
235        let Some(len_bytes) = bytes.get(..HEADER_LEN_BYTES) else {
236            return false;
237        };
238        let header_len =
239            u64::from_le_bytes(len_bytes.try_into().expect("checked slice is exactly 8 bytes"));
240        if header_len == 0 {
241            return false;
242        }
243        match bytes.get(HEADER_LEN_BYTES) {
244            Some(b'{') => true,
245            Some(_) => false,
246            // Not enough bytes sampled to see the header's first byte;
247            // the length prefix alone is a weak but nonzero signal.
248            None => true,
249        }
250    }
251
252    fn load(&self, path: &Path) -> Result<LoadedModel> {
253        let source = ByteSource::open(path)?;
254        parse(source)
255    }
256}