1use 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
15struct Shard {
17 mmap: Mmap,
18}
19
20struct TensorEntry {
22 shard: usize,
23 dtype: TensorDtype,
24 shape: Vec<usize>,
25}
26
27pub 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 pub fn load(dir: impl AsRef<Path>) -> Result<Self> {
99 let dir = dir.as_ref();
100
101 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 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 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 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 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 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 #[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 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 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 #[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}