1use crate::op::BackpropOp;
13use crate::storage::Storage;
14use crate::tensor::from_storage;
15use crate::{DType, Device, Error, Result, Tensor, WithDType};
16use safetensors::tensor as st;
17use safetensors::tensor::SafeTensors;
18use std::borrow::Cow;
19use std::collections::HashMap;
20use std::path::Path;
21
22impl From<DType> for st::Dtype {
23 fn from(value: DType) -> Self {
24 match value {
25 DType::U8 => st::Dtype::U8,
26 DType::U32 => st::Dtype::U32,
27 DType::I16 => st::Dtype::I16,
28 DType::I32 => st::Dtype::I32,
29 DType::I64 => st::Dtype::I64,
30 DType::BF16 => st::Dtype::BF16,
31 DType::F16 => st::Dtype::F16,
32 DType::F32 => st::Dtype::F32,
33 DType::F64 => st::Dtype::F64,
34 DType::F8E4M3 => st::Dtype::F8_E4M3,
35 DType::F6E2M3 => st::Dtype::F6_E2M3,
36 DType::F6E3M2 => st::Dtype::F6_E3M2,
37 DType::F4 => st::Dtype::F4,
38 DType::F8E8M0 => st::Dtype::F8_E8M0,
39 }
40 }
41}
42
43impl TryFrom<st::Dtype> for DType {
44 type Error = Error;
45 fn try_from(value: st::Dtype) -> Result<Self> {
46 match value {
47 st::Dtype::U8 => Ok(DType::U8),
48 st::Dtype::U32 => Ok(DType::U32),
49 st::Dtype::I16 => Ok(DType::I16),
50 st::Dtype::I32 => Ok(DType::I32),
51 st::Dtype::I64 => Ok(DType::I64),
52 st::Dtype::BF16 => Ok(DType::BF16),
53 st::Dtype::F16 => Ok(DType::F16),
54 st::Dtype::F32 => Ok(DType::F32),
55 st::Dtype::F64 => Ok(DType::F64),
56 st::Dtype::F8_E4M3 => Ok(DType::F8E4M3),
57 st::Dtype::F6_E2M3 => Ok(DType::F6E2M3),
58 st::Dtype::F6_E3M2 => Ok(DType::F6E3M2),
59 st::Dtype::F4 => Ok(DType::F4),
60 st::Dtype::F8_E8M0 => Ok(DType::F8E8M0),
61 dtype => Err(Error::UnsupportedSafeTensorDtype(dtype)),
62 }
63 }
64}
65
66impl st::View for Tensor {
67 fn dtype(&self) -> st::Dtype {
68 self.dtype().into()
69 }
70 fn shape(&self) -> &[usize] {
71 self.shape().dims()
72 }
73
74 fn data(&self) -> Cow<'_, [u8]> {
75 Cow::Owned(convert_back(self).unwrap())
78 }
79
80 fn data_len(&self) -> usize {
81 let n: usize = self.shape().elem_count();
82 let bytes_per_element = self.dtype().size_in_bytes();
83 n * bytes_per_element
84 }
85}
86
87impl st::View for &Tensor {
88 fn dtype(&self) -> st::Dtype {
89 (*self).dtype().into()
90 }
91 fn shape(&self) -> &[usize] {
92 self.dims()
93 }
94
95 fn data(&self) -> Cow<'_, [u8]> {
96 Cow::Owned(convert_back(self).unwrap())
99 }
100
101 fn data_len(&self) -> usize {
102 let n: usize = self.dims().iter().product();
103 let bytes_per_element = (*self).dtype().size_in_bytes();
104 n * bytes_per_element
105 }
106}
107
108impl Tensor {
109 pub fn save_safetensors<P: AsRef<Path>>(&self, name: &str, filename: P) -> Result<()> {
110 let data = [(name, self.clone())];
111 Ok(st::serialize_to_file(data, None, filename.as_ref())?)
112 }
113}
114
115fn convert_slice<T: WithDType>(data: &[u8], shape: &[usize], device: &Device) -> Result<Tensor> {
116 let size_in_bytes = T::DTYPE.size_in_bytes();
117 let elem_count = data.len() / size_in_bytes;
118 if (data.as_ptr() as usize).is_multiple_of(size_in_bytes) {
119 let data: &[T] =
122 unsafe { std::slice::from_raw_parts(data.as_ptr() as *const T, elem_count) };
123 Tensor::from_slice(data, shape, device)
124 } else {
125 let mut c: Vec<T> = Vec::with_capacity(elem_count);
128 unsafe {
133 std::ptr::copy_nonoverlapping(data.as_ptr(), c.as_mut_ptr() as *mut u8, data.len());
134 c.set_len(elem_count)
135 }
136 Tensor::from_slice(&c, shape, device)
137 }
138}
139
140fn convert_slice_with_cast<T: Sized + Copy, U: WithDType, F: Fn(T) -> Result<U>>(
141 data: &[u8],
142 shape: &[usize],
143 device: &Device,
144 conv: F,
145) -> Result<Tensor> {
146 let size_in_bytes = std::mem::size_of::<T>();
147 let elem_count = data.len() / size_in_bytes;
148 if (data.as_ptr() as usize).is_multiple_of(size_in_bytes) {
149 let data: &[T] =
152 unsafe { std::slice::from_raw_parts(data.as_ptr() as *const T, elem_count) };
153 let data = data.iter().map(|t| conv(*t)).collect::<Result<Vec<_>>>()?;
154 Tensor::from_vec(data, shape, device)
155 } else {
156 let mut c: Vec<T> = Vec::with_capacity(elem_count);
159 unsafe {
164 std::ptr::copy_nonoverlapping(data.as_ptr(), c.as_mut_ptr() as *mut u8, data.len());
165 c.set_len(elem_count)
166 }
167 let c = c.into_iter().map(conv).collect::<Result<Vec<_>>>()?;
168 Tensor::from_vec(c, shape, device)
169 }
170}
171
172fn convert_with_cast_<T: Sized + Copy, U: WithDType, F: Fn(T) -> Result<U>>(
173 view: &st::TensorView<'_>,
174 device: &Device,
175 conv: F,
176) -> Result<Tensor> {
177 convert_slice_with_cast::<T, U, F>(view.data(), view.shape(), device, conv)
178}
179
180fn convert_<T: WithDType>(view: &st::TensorView<'_>, device: &Device) -> Result<Tensor> {
181 convert_slice::<T>(view.data(), view.shape(), device)
182}
183
184fn convert_back_<T: WithDType>(mut vs: Vec<T>) -> Vec<u8> {
185 let size_in_bytes = T::DTYPE.size_in_bytes();
186 let length = vs.len() * size_in_bytes;
187 let capacity = vs.capacity() * size_in_bytes;
188 let ptr = vs.as_mut_ptr() as *mut u8;
189 std::mem::forget(vs);
191 unsafe { Vec::from_raw_parts(ptr, length, capacity) }
196}
197
198pub trait Load {
199 fn load(&self, device: &Device) -> Result<Tensor>;
200}
201
202impl Load for st::TensorView<'_> {
203 fn load(&self, device: &Device) -> Result<Tensor> {
204 convert(self, device)
205 }
206}
207
208impl Tensor {
209 pub fn from_raw_buffer(
210 data: &[u8],
211 dtype: DType,
212 shape: &[usize],
213 device: &Device,
214 ) -> Result<Self> {
215 match dtype {
216 DType::U8 => convert_slice::<u8>(data, shape, device),
217 DType::U32 => convert_slice::<u32>(data, shape, device),
218 DType::I16 => convert_slice::<i16>(data, shape, device),
219 DType::I32 => convert_slice::<i32>(data, shape, device),
220 DType::I64 => convert_slice::<i64>(data, shape, device),
221 DType::BF16 => convert_slice::<half::bf16>(data, shape, device),
222 DType::F16 => convert_slice::<half::f16>(data, shape, device),
223 DType::F32 => convert_slice::<f32>(data, shape, device),
224 DType::F64 => convert_slice::<f64>(data, shape, device),
225 DType::F8E4M3 => convert_slice::<float8::F8E4M3>(data, shape, device),
226 DType::F6E2M3 | DType::F6E3M2 | DType::F4 | DType::F8E8M0 => {
227 let storage = match device {
229 Device::Cpu => {
230 let cpu_storage = match dtype {
231 DType::F6E2M3 => crate::cpu_backend::CpuStorage::F6E2M3(data.to_vec()),
232 DType::F6E3M2 => crate::cpu_backend::CpuStorage::F6E3M2(data.to_vec()),
233 DType::F4 => crate::cpu_backend::CpuStorage::F4(data.to_vec()),
234 DType::F8E8M0 => crate::cpu_backend::CpuStorage::F8E8M0(data.to_vec()),
235 _ => unreachable!(),
236 };
237 Storage::Cpu(cpu_storage)
238 }
239 #[cfg(feature = "cuda")]
240 Device::Cuda(device) => {
241 let mut slice = unsafe { device.alloc::<u8>(data.len())? };
242 device.memcpy_htod(data, &mut slice)?;
243
244 let slice = match dtype {
245 DType::F6E2M3 => crate::cuda_backend::CudaStorageSlice::F6E2M3(slice),
246 DType::F6E3M2 => crate::cuda_backend::CudaStorageSlice::F6E3M2(slice),
247 DType::F4 => crate::cuda_backend::CudaStorageSlice::F4(slice),
248 DType::F8E8M0 => crate::cuda_backend::CudaStorageSlice::F8E8M0(slice),
249 _ => unreachable!(),
250 };
251 let storage = crate::cuda_backend::CudaStorage {
252 slice,
253 device: device.clone(),
254 };
255 Storage::Cuda(storage)
256 }
257 #[cfg(not(feature = "cuda"))]
258 Device::Cuda(_) => {
259 return Err(Error::Msg("CUDA support not compiled".to_string()));
260 }
261 #[cfg(feature = "metal")]
262 Device::Metal(device) => {
263 let buffer = device.new_buffer_with_data(data)?;
264
265 let storage = crate::metal_backend::MetalStorage::new(
266 buffer,
267 device.clone(),
268 data.len(),
269 dtype,
270 );
271 Storage::Metal(storage)
272 }
273 #[cfg(feature = "rocm")]
274 Device::Rocm(_) => crate::bail!("not supported on rocm yet"),
275 #[cfg(feature = "vulkan")]
276 Device::Vulkan(_) => crate::bail!("not supported on vulkan yet"),
277 #[cfg(feature = "wgpu")]
278 Device::Wgpu(_) => crate::bail!("not supported on wgpu yet"),
279 #[cfg(not(feature = "metal"))]
280 Device::Metal(_) => {
281 return Err(Error::Msg("Metal support not compiled".to_string()));
282 }
283 };
284
285 let op = BackpropOp::none();
286 Ok(from_storage(storage, shape, op, false))
287 }
288 }
289 }
290}
291
292fn convert(view: &st::TensorView<'_>, device: &Device) -> Result<Tensor> {
293 match view.dtype() {
294 st::Dtype::U8 => convert_::<u8>(view, device),
295 st::Dtype::U16 => {
296 let conv = |x| Ok(u32::from(x));
297 convert_with_cast_::<u16, u32, _>(view, device, conv)
298 }
299 st::Dtype::U32 => convert_::<u32>(view, device),
300 st::Dtype::I16 => convert_::<i16>(view, device),
301 st::Dtype::I32 => convert_::<i32>(view, device),
302 st::Dtype::I64 => convert_::<i64>(view, device),
303 st::Dtype::BF16 => convert_::<half::bf16>(view, device),
304 st::Dtype::F16 => convert_::<half::f16>(view, device),
305 st::Dtype::F32 => convert_::<f32>(view, device),
306 st::Dtype::F64 => convert_::<f64>(view, device),
307 st::Dtype::F8_E4M3 => convert_::<float8::F8E4M3>(view, device),
308 st::Dtype::F6_E2M3 | st::Dtype::F6_E3M2 | st::Dtype::F4 | st::Dtype::F8_E8M0 => {
309 convert_dummy(view, device)
313 }
314 dtype => Err(Error::UnsupportedSafeTensorDtype(dtype)),
315 }
316}
317
318fn convert_dummy(view: &st::TensorView<'_>, device: &Device) -> Result<Tensor> {
319 let (dtype, _dtype_name) = match view.dtype() {
322 st::Dtype::F6_E2M3 => (DType::F6E2M3, "F6_E2M3 (MX6)"),
323 st::Dtype::F6_E3M2 => (DType::F6E3M2, "F6_E3M2 (MX6)"),
324 st::Dtype::F4 => (DType::F4, "F4 (MX4)"),
325 st::Dtype::F8_E8M0 => (DType::F8E8M0, "F8_E8M0"),
326 _ => unreachable!("convert_dummy called with non-dummy dtype"),
327 };
328
329 let data = view.data();
331 let shape = view.shape();
332
333 let storage = match device {
335 Device::Cpu => {
336 let cpu_storage = match dtype {
337 DType::F6E2M3 => crate::cpu_backend::CpuStorage::F6E2M3(data.to_vec()),
338 DType::F6E3M2 => crate::cpu_backend::CpuStorage::F6E3M2(data.to_vec()),
339 DType::F4 => crate::cpu_backend::CpuStorage::F4(data.to_vec()),
340 DType::F8E8M0 => crate::cpu_backend::CpuStorage::F8E8M0(data.to_vec()),
341 _ => unreachable!(),
342 };
343 Storage::Cpu(cpu_storage)
344 }
345 #[cfg(feature = "cuda")]
346 Device::Cuda(device) => {
347 let mut slice = unsafe { device.alloc::<u8>(data.len())? };
348 device.memcpy_htod(data, &mut slice)?;
349
350 let slice = match dtype {
351 DType::F6E2M3 => crate::cuda_backend::CudaStorageSlice::F6E2M3(slice),
352 DType::F6E3M2 => crate::cuda_backend::CudaStorageSlice::F6E3M2(slice),
353 DType::F4 => crate::cuda_backend::CudaStorageSlice::F4(slice),
354 DType::F8E8M0 => crate::cuda_backend::CudaStorageSlice::F8E8M0(slice),
355 _ => unreachable!(),
356 };
357 let storage = crate::cuda_backend::CudaStorage {
358 slice,
359 device: device.clone(),
360 };
361 Storage::Cuda(storage)
362 }
363 #[cfg(not(feature = "cuda"))]
364 Device::Cuda(_) => {
365 return Err(Error::Msg("CUDA support not compiled".to_string()));
366 }
367 #[cfg(feature = "metal")]
368 Device::Metal(device) => {
369 let buffer = device.new_buffer_with_data(data)?;
370
371 let storage =
372 crate::metal_backend::MetalStorage::new(buffer, device.clone(), data.len(), dtype);
373 Storage::Metal(storage)
374 }
375 #[cfg(feature = "rocm")]
376 Device::Rocm(_) => crate::bail!("not supported on rocm yet"),
377 #[cfg(feature = "vulkan")]
378 Device::Vulkan(_) => crate::bail!("not supported on vulkan yet"),
379 #[cfg(feature = "wgpu")]
380 Device::Wgpu(_) => crate::bail!("not supported on wgpu yet"),
381 #[cfg(not(feature = "metal"))]
382 Device::Metal(_) => {
383 return Err(Error::Msg("Metal support not compiled".to_string()));
384 }
385 };
386
387 let op = BackpropOp::none();
389 Ok(from_storage(storage, shape, op, false))
390}
391
392fn convert_back(tensor: &Tensor) -> Result<Vec<u8>> {
393 let tensor = tensor.flatten_all()?;
395 match tensor.dtype() {
396 DType::U8 => Ok(convert_back_::<u8>(tensor.to_vec1()?)),
397 DType::U32 => Ok(convert_back_::<u32>(tensor.to_vec1()?)),
398 DType::I16 => Ok(convert_back_::<i16>(tensor.to_vec1()?)),
399 DType::I32 => Ok(convert_back_::<i32>(tensor.to_vec1()?)),
400 DType::I64 => Ok(convert_back_::<i64>(tensor.to_vec1()?)),
401 DType::F16 => Ok(convert_back_::<half::f16>(tensor.to_vec1()?)),
402 DType::BF16 => Ok(convert_back_::<half::bf16>(tensor.to_vec1()?)),
403 DType::F32 => Ok(convert_back_::<f32>(tensor.to_vec1()?)),
404 DType::F64 => Ok(convert_back_::<f64>(tensor.to_vec1()?)),
405 DType::F8E4M3 => Ok(convert_back_::<float8::F8E4M3>(tensor.to_vec1()?)),
406 DType::F6E2M3 | DType::F6E3M2 | DType::F4 | DType::F8E8M0 => {
407 Err(Error::Msg("Internal error: dtype mismatch in storage".to_string()).bt())
408 }
409 }
410}
411
412pub fn load<P: AsRef<Path>>(filename: P, device: &Device) -> Result<HashMap<String, Tensor>> {
413 let data = std::fs::read(filename.as_ref())?;
414 load_buffer(&data[..], device)
415}
416
417pub fn load_buffer(data: &[u8], device: &Device) -> Result<HashMap<String, Tensor>> {
418 let st = safetensors::SafeTensors::deserialize(data)?;
419 st.tensors()
420 .into_iter()
421 .map(|(name, view)| Ok((name, view.load(device)?)))
422 .collect()
423}
424
425pub fn save<K: AsRef<str> + Ord + std::fmt::Display, P: AsRef<Path>>(
426 tensors: &HashMap<K, Tensor>,
427 filename: P,
428) -> Result<()> {
429 Ok(st::serialize_to_file(tensors, None, filename.as_ref())?)
430}
431
432#[derive(yoke::Yokeable)]
433struct SafeTensors_<'a>(SafeTensors<'a>);
434
435pub struct MmapedSafetensors {
436 safetensors: Vec<yoke::Yoke<SafeTensors_<'static>, memmap2::Mmap>>,
437 routing: Option<HashMap<String, usize>>,
438}
439
440impl MmapedSafetensors {
441 pub unsafe fn new<P: AsRef<Path>>(p: P) -> Result<Self> {
447 let p = p.as_ref();
448 let file = std::fs::File::open(p).map_err(|e| Error::from(e).with_path(p))?;
449 let file = memmap2::MmapOptions::new()
450 .map(&file)
451 .map_err(|e| Error::from(e).with_path(p))?;
452 let safetensors = yoke::Yoke::<SafeTensors_<'static>, memmap2::Mmap>::try_attach_to_cart(
453 file,
454 |data: &[u8]| {
455 let st = safetensors::SafeTensors::deserialize(data)
456 .map_err(|e| Error::from(e).with_path(p))?;
457 Ok::<_, Error>(SafeTensors_(st))
458 },
459 )?;
460 Ok(Self {
461 safetensors: vec![safetensors],
462 routing: None,
463 })
464 }
465
466 pub unsafe fn multi<P: AsRef<Path>>(paths: &[P]) -> Result<Self> {
474 let mut routing = HashMap::new();
475 let mut safetensors = vec![];
476 for (index, p) in paths.iter().enumerate() {
477 let p = p.as_ref();
478 let file = std::fs::File::open(p).map_err(|e| Error::from(e).with_path(p))?;
479 let file = memmap2::MmapOptions::new()
480 .map(&file)
481 .map_err(|e| Error::from(e).with_path(p))?;
482 let data = yoke::Yoke::<SafeTensors_<'static>, memmap2::Mmap>::try_attach_to_cart(
483 file,
484 |data: &[u8]| {
485 let st = safetensors::SafeTensors::deserialize(data)
486 .map_err(|e| Error::from(e).with_path(p))?;
487 Ok::<_, Error>(SafeTensors_(st))
488 },
489 )?;
490 for k in data.get().0.names() {
491 routing.insert(k.to_string(), index);
492 }
493 safetensors.push(data)
494 }
495 Ok(Self {
496 safetensors,
497 routing: Some(routing),
498 })
499 }
500
501 pub fn load(&self, name: &str, dev: &Device) -> Result<Tensor> {
502 self.get(name)?.load(dev)
503 }
504
505 pub fn tensors(&self) -> Vec<(String, st::TensorView<'_>)> {
506 let mut tensors = vec![];
507 for safetensors in self.safetensors.iter() {
508 tensors.push(safetensors.get().0.tensors())
509 }
510 tensors.into_iter().flatten().collect()
511 }
512
513 pub fn get(&self, name: &str) -> Result<st::TensorView<'_>> {
514 let index = match &self.routing {
515 None => 0,
516 Some(routing) => {
517 let index = routing.get(name).ok_or_else(|| {
518 Error::CannotFindTensor {
519 path: name.to_string(),
520 }
521 .bt()
522 })?;
523 *index
524 }
525 };
526 Ok(self.safetensors[index].get().0.tensor(name)?)
527 }
528}
529
530pub struct SliceSafetensors<'a> {
531 safetensors: SafeTensors<'a>,
532}
533
534impl<'a> SliceSafetensors<'a> {
535 pub fn new(buffer: &'a [u8]) -> Result<Self> {
537 let safetensors = safetensors::SafeTensors::deserialize(buffer)?;
538 Ok(Self { safetensors })
539 }
540
541 pub fn load(&self, name: &str, dev: &Device) -> Result<Tensor> {
542 self.safetensors.tensor(name)?.load(dev)
543 }
544
545 pub fn tensors(&self) -> Vec<(String, st::TensorView<'_>)> {
546 self.safetensors.tensors()
547 }
548
549 pub fn get(&self, name: &str) -> Result<st::TensorView<'_>> {
550 Ok(self.safetensors.tensor(name)?)
551 }
552}
553
554pub struct BufferedSafetensors {
555 safetensors: yoke::Yoke<SafeTensors_<'static>, Vec<u8>>,
556}
557
558impl BufferedSafetensors {
559 pub fn new(buffer: Vec<u8>) -> Result<Self> {
561 let safetensors = yoke::Yoke::<SafeTensors_<'static>, Vec<u8>>::try_attach_to_cart(
562 buffer,
563 |data: &[u8]| {
564 let st = safetensors::SafeTensors::deserialize(data)?;
565 Ok::<_, Error>(SafeTensors_(st))
566 },
567 )?;
568 Ok(Self { safetensors })
569 }
570
571 pub fn load(&self, name: &str, dev: &Device) -> Result<Tensor> {
572 self.get(name)?.load(dev)
573 }
574
575 pub fn tensors(&self) -> Vec<(String, st::TensorView<'_>)> {
576 self.safetensors.get().0.tensors()
577 }
578
579 pub fn get(&self, name: &str) -> Result<st::TensorView<'_>> {
580 Ok(self.safetensors.get().0.tensor(name)?)
581 }
582}
583
584pub struct MmapedFile {
585 path: std::path::PathBuf,
586 inner: memmap2::Mmap,
587}
588
589impl MmapedFile {
590 pub unsafe fn new<P: AsRef<Path>>(p: P) -> Result<Self> {
597 let p = p.as_ref();
598 let file = std::fs::File::open(p).map_err(|e| Error::from(e).with_path(p))?;
599 let inner = memmap2::MmapOptions::new()
600 .map(&file)
601 .map_err(|e| Error::from(e).with_path(p))?;
602 Ok(Self {
603 inner,
604 path: p.to_path_buf(),
605 })
606 }
607
608 pub fn deserialize(&self) -> Result<SafeTensors<'_>> {
609 let st = safetensors::SafeTensors::deserialize(&self.inner)
610 .map_err(|e| Error::from(e).with_path(&self.path))?;
611 Ok(st)
612 }
613}
614
615#[cfg(test)]
616mod tests {
617 use super::*;
618 use std::collections::HashMap;
619
620 #[test]
621 fn save_single_tensor() {
622 let t = Tensor::zeros((2, 2), DType::F32, &Device::Cpu).unwrap();
623 t.save_safetensors("t", "t.safetensors").unwrap();
624 let bytes = std::fs::read("t.safetensors").unwrap();
625 assert_eq!(bytes, b"@\0\0\0\0\0\0\0{\"t\":{\"dtype\":\"F32\",\"shape\":[2,2],\"data_offsets\":[0,16]}} \0\0\0\0\0\0\0\0\0\0\0\0\0\0\0\0");
626 std::fs::remove_file("t.safetensors").unwrap();
627 }
628
629 #[test]
630 fn save_load_multiple_tensors() {
631 let t = Tensor::zeros((2, 2), DType::F32, &Device::Cpu).unwrap();
632 let u = Tensor::zeros((1, 2), DType::F32, &Device::Cpu).unwrap();
633 let map: HashMap<_, _> = [("t", t), ("u", u)].into_iter().collect();
634 save(&map, "multi.safetensors").unwrap();
635
636 let weights = load("multi.safetensors", &Device::Cpu).unwrap();
637 assert_eq!(weights.get("t").unwrap().dims(), &[2, 2]);
638 assert_eq!(weights.get("u").unwrap().dims(), &[1, 2]);
639 let bytes = std::fs::read("multi.safetensors").unwrap();
640 assert_eq!(bytes, b"x\0\0\0\0\0\0\0{\"t\":{\"dtype\":\"F32\",\"shape\":[2,2],\"data_offsets\":[0,16]},\"u\":{\"dtype\":\"F32\",\"shape\":[1,2],\"data_offsets\":[16,24]}} \0\0\0\0\0\0\0\0\0\0\0\0\0\0\0\0\0\0\0\0\0\0\0\0");
641 std::fs::remove_file("multi.safetensors").unwrap();
642 }
643
644 #[test]
645 fn load_u8() {
646 let bytes = b"8\0\0\0\0\0\0\0{\"x\":{\"dtype\":\"U8\",\"shape\":[2],\"data_offsets\":[0,2]}} \x01\x03";
647 std::fs::write("test_u8.safetensors", bytes).unwrap();
648 let weights = load("test_u8.safetensors", &Device::Cpu).unwrap();
649 let tensor = weights.get("x").unwrap();
650 assert_eq!(tensor.dims(), &[2]);
651 assert_eq!(tensor.dtype(), DType::U8);
652 let data: Vec<u8> = tensor.to_vec1().unwrap();
653 assert_eq!(data, vec![1, 3]);
654 std::fs::remove_file("test_u8.safetensors").unwrap();
655 }
656}