kopitiam_loader/
safetensors.rs1use 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#[derive(Debug, Deserialize)]
67struct RawTensorInfo {
68 dtype: String,
69 shape: Vec<u64>,
70 data_offsets: (u64, u64),
71}
72
73fn 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
86fn 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 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 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
220pub 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 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 None => true,
249 }
250 }
251
252 fn load(&self, path: &Path) -> Result<LoadedModel> {
253 let source = ByteSource::open(path)?;
254 parse(source)
255 }
256}