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 {
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 pub fn load(dir: impl AsRef<Path>) -> Result<Self> {
83 let dir = dir.as_ref();
84
85 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 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 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 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 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 #[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 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}