1use std::collections::BTreeMap;
17use std::fmt;
18
19use serde_json::Value;
20
21const MAX_HEADER_BYTES: u64 = 64 * 1024 * 1024;
26
27#[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord)]
33pub enum Dtype {
34 Bf16,
36 F32,
38}
39
40impl Dtype {
41 #[must_use]
43 pub const fn size(self) -> usize {
44 match self {
45 Self::Bf16 => 2,
46 Self::F32 => 4,
47 }
48 }
49
50 #[must_use]
52 pub const fn as_str(self) -> &'static str {
53 match self {
54 Self::Bf16 => "BF16",
55 Self::F32 => "F32",
56 }
57 }
58
59 fn parse(raw: &str) -> Option<Self> {
60 match raw {
61 "BF16" => Some(Self::Bf16),
62 "F32" => Some(Self::F32),
63 _ => None,
64 }
65 }
66}
67
68impl fmt::Display for Dtype {
69 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
70 f.write_str(self.as_str())
71 }
72}
73
74#[derive(Clone, Debug, PartialEq, Eq)]
76pub struct TensorEntry {
77 pub name: String,
79 pub dtype: Dtype,
81 pub shape: Vec<usize>,
83 pub begin: usize,
85 pub end: usize,
87}
88
89impl TensorEntry {
90 #[must_use]
92 pub fn element_count(&self) -> usize {
93 self.shape.iter().product()
94 }
95
96 #[must_use]
98 pub const fn byte_len(&self) -> usize {
99 self.end - self.begin
100 }
101
102 #[must_use]
106 pub fn row_len(&self) -> usize {
107 self.shape.iter().skip(1).product()
108 }
109}
110
111#[derive(Clone, Debug, PartialEq, Eq)]
116pub enum WeightsError {
117 TooShortForHeader {
119 len: usize,
121 },
122 HeaderLengthOutOfRange {
124 declared: u64,
126 available: usize,
128 },
129 HeaderNotJson {
131 detail: String,
133 },
134 HeaderNotObject,
136 MalformedEntry {
138 name: String,
140 detail: String,
142 },
143 UnsupportedDtype {
145 name: String,
147 raw: String,
149 },
150 SpanOutOfBounds {
152 name: String,
154 begin: usize,
156 end: usize,
158 payload_len: usize,
160 },
161 ShapeSpanMismatch {
163 name: String,
165 shape: Vec<usize>,
167 expected_bytes: usize,
169 actual_bytes: usize,
171 },
172 ShapeOverflow {
174 name: String,
176 shape: Vec<usize>,
178 },
179}
180
181impl fmt::Display for WeightsError {
182 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
183 match self {
184 Self::TooShortForHeader { len } => {
185 write!(f, "not a safetensors file: {len} bytes, need at least 8")
186 }
187 Self::HeaderLengthOutOfRange {
188 declared,
189 available,
190 } => write!(
191 f,
192 "header length {declared} is out of range (file has {available} bytes, cap is \
193 {MAX_HEADER_BYTES})"
194 ),
195 Self::HeaderNotJson { detail } => write!(f, "header is not valid JSON: {detail}"),
196 Self::HeaderNotObject => f.write_str("header JSON is not an object"),
197 Self::MalformedEntry { name, detail } => {
198 write!(f, "tensor `{name}`: {detail}")
199 }
200 Self::UnsupportedDtype { name, raw } => write!(
201 f,
202 "tensor `{name}`: unsupported dtype `{raw}` (accepted: BF16, F32)"
203 ),
204 Self::SpanOutOfBounds {
205 name,
206 begin,
207 end,
208 payload_len,
209 } => write!(
210 f,
211 "tensor `{name}`: byte span {begin}..{end} escapes the {payload_len}-byte payload"
212 ),
213 Self::ShapeSpanMismatch {
214 name,
215 shape,
216 expected_bytes,
217 actual_bytes,
218 } => write!(
219 f,
220 "tensor `{name}`: shape {shape:?} implies {expected_bytes} bytes but the span \
221 covers {actual_bytes}"
222 ),
223 Self::ShapeOverflow { name, shape } => {
224 write!(f, "tensor `{name}`: shape {shape:?} overflows usize")
225 }
226 }
227 }
228}
229
230impl std::error::Error for WeightsError {}
231
232#[derive(Clone, Debug)]
237pub struct SafetensorsIndex {
238 entries: BTreeMap<String, TensorEntry>,
239 payload_begin: usize,
240}
241
242impl SafetensorsIndex {
243 pub fn parse(bytes: &[u8]) -> Result<Self, WeightsError> {
250 let Some(len_prefix) = bytes.get(..8) else {
251 return Err(WeightsError::TooShortForHeader { len: bytes.len() });
252 };
253 let header_len = u64::from_le_bytes(
255 len_prefix
256 .try_into()
257 .expect("slice of 8 bytes converts to [u8; 8]"),
258 );
259
260 if header_len > MAX_HEADER_BYTES {
261 return Err(WeightsError::HeaderLengthOutOfRange {
262 declared: header_len,
263 available: bytes.len(),
264 });
265 }
266 let header_len_usize =
268 usize::try_from(header_len).map_err(|_| WeightsError::HeaderLengthOutOfRange {
269 declared: header_len,
270 available: bytes.len(),
271 })?;
272 let payload_begin =
273 8usize
274 .checked_add(header_len_usize)
275 .ok_or(WeightsError::HeaderLengthOutOfRange {
276 declared: header_len,
277 available: bytes.len(),
278 })?;
279 let Some(header_bytes) = bytes.get(8..payload_begin) else {
280 return Err(WeightsError::HeaderLengthOutOfRange {
281 declared: header_len,
282 available: bytes.len(),
283 });
284 };
285
286 let parsed: Value =
287 serde_json::from_slice(header_bytes).map_err(|error| WeightsError::HeaderNotJson {
288 detail: error.to_string(),
289 })?;
290 let Value::Object(directory) = parsed else {
291 return Err(WeightsError::HeaderNotObject);
292 };
293
294 let payload_len = bytes.len() - payload_begin;
295 let mut entries = BTreeMap::new();
296 for (name, value) in directory {
297 if name == "__metadata__" {
300 continue;
301 }
302 let entry = parse_entry(&name, &value, payload_begin, payload_len)?;
303 entries.insert(name, entry);
304 }
305
306 Ok(Self {
307 entries,
308 payload_begin,
309 })
310 }
311
312 #[must_use]
314 pub const fn payload_begin(&self) -> usize {
315 self.payload_begin
316 }
317
318 #[must_use]
320 pub fn len(&self) -> usize {
321 self.entries.len()
322 }
323
324 #[must_use]
326 pub fn is_empty(&self) -> bool {
327 self.entries.is_empty()
328 }
329
330 #[must_use]
332 pub fn entry(&self, name: &str) -> Option<&TensorEntry> {
333 self.entries.get(name)
334 }
335
336 pub fn entries(&self) -> impl Iterator<Item = &TensorEntry> {
338 self.entries.values()
339 }
340
341 pub fn names(&self) -> impl Iterator<Item = &str> {
343 self.entries.keys().map(String::as_str)
344 }
345
346 #[must_use]
350 pub fn total_tensor_bytes(&self) -> usize {
351 self.entries.values().map(TensorEntry::byte_len).sum()
352 }
353
354 #[must_use]
359 pub fn view<'a>(&self, name: &str, bytes: &'a [u8]) -> Option<TensorView<'a>> {
360 let entry = self.entries.get(name)?;
361 let raw = bytes.get(entry.begin..entry.end)?;
362 Some(TensorView {
363 dtype: entry.dtype,
364 shape: entry.shape.clone(),
365 raw,
366 })
367 }
368}
369
370fn parse_entry(
371 name: &str,
372 value: &Value,
373 payload_begin: usize,
374 payload_len: usize,
375) -> Result<TensorEntry, WeightsError> {
376 let object = value
377 .as_object()
378 .ok_or_else(|| WeightsError::MalformedEntry {
379 name: name.to_owned(),
380 detail: "entry is not a JSON object".to_owned(),
381 })?;
382
383 let raw_dtype = object.get("dtype").and_then(Value::as_str).ok_or_else(|| {
384 WeightsError::MalformedEntry {
385 name: name.to_owned(),
386 detail: "missing string field `dtype`".to_owned(),
387 }
388 })?;
389 let dtype = Dtype::parse(raw_dtype).ok_or_else(|| WeightsError::UnsupportedDtype {
390 name: name.to_owned(),
391 raw: raw_dtype.to_owned(),
392 })?;
393
394 let raw_shape = object
395 .get("shape")
396 .and_then(Value::as_array)
397 .ok_or_else(|| WeightsError::MalformedEntry {
398 name: name.to_owned(),
399 detail: "missing array field `shape`".to_owned(),
400 })?;
401 let mut shape = Vec::with_capacity(raw_shape.len());
402 for dim in raw_shape {
403 let dim = dim
404 .as_u64()
405 .and_then(|d| usize::try_from(d).ok())
406 .ok_or_else(|| WeightsError::MalformedEntry {
407 name: name.to_owned(),
408 detail: "shape contains a non-usize dimension".to_owned(),
409 })?;
410 shape.push(dim);
411 }
412
413 let offsets = object
414 .get("data_offsets")
415 .and_then(Value::as_array)
416 .ok_or_else(|| WeightsError::MalformedEntry {
417 name: name.to_owned(),
418 detail: "missing array field `data_offsets`".to_owned(),
419 })?;
420 if offsets.len() != 2 {
421 return Err(WeightsError::MalformedEntry {
422 name: name.to_owned(),
423 detail: format!("`data_offsets` has {} entries, expected 2", offsets.len()),
424 });
425 }
426 let mut bound = [0usize; 2];
427 for (slot, raw) in bound.iter_mut().zip(offsets) {
428 *slot = raw
429 .as_u64()
430 .and_then(|v| usize::try_from(v).ok())
431 .ok_or_else(|| WeightsError::MalformedEntry {
432 name: name.to_owned(),
433 detail: "`data_offsets` contains a non-usize value".to_owned(),
434 })?;
435 }
436 let [begin, end] = bound;
437
438 if begin > end || end > payload_len {
440 return Err(WeightsError::SpanOutOfBounds {
441 name: name.to_owned(),
442 begin,
443 end,
444 payload_len,
445 });
446 }
447
448 let mut elements = 1usize;
450 for dim in &shape {
451 elements = elements
452 .checked_mul(*dim)
453 .ok_or_else(|| WeightsError::ShapeOverflow {
454 name: name.to_owned(),
455 shape: shape.clone(),
456 })?;
457 }
458 let expected_bytes =
459 elements
460 .checked_mul(dtype.size())
461 .ok_or_else(|| WeightsError::ShapeOverflow {
462 name: name.to_owned(),
463 shape: shape.clone(),
464 })?;
465 let actual_bytes = end - begin;
466 if expected_bytes != actual_bytes {
467 return Err(WeightsError::ShapeSpanMismatch {
468 name: name.to_owned(),
469 shape,
470 expected_bytes,
471 actual_bytes,
472 });
473 }
474
475 Ok(TensorEntry {
476 name: name.to_owned(),
477 dtype,
478 shape,
479 begin: payload_begin + begin,
480 end: payload_begin + end,
481 })
482}
483
484#[derive(Clone, Copy, Debug)]
488pub struct TensorViewRef<'a> {
489 dtype: Dtype,
490 raw: &'a [u8],
491}
492
493#[derive(Clone, Debug)]
495pub struct TensorView<'a> {
496 dtype: Dtype,
497 shape: Vec<usize>,
498 raw: &'a [u8],
499}
500
501#[must_use]
507pub const fn bf16_bits_to_f32(bits: u16) -> f32 {
508 f32::from_bits((bits as u32) << 16)
509}
510
511impl<'a> TensorView<'a> {
512 #[must_use]
514 pub const fn dtype(&self) -> Dtype {
515 self.dtype
516 }
517
518 #[must_use]
520 pub fn shape(&self) -> &[usize] {
521 &self.shape
522 }
523
524 #[must_use]
526 pub fn len(&self) -> usize {
527 self.raw.len() / self.dtype.size()
528 }
529
530 #[must_use]
532 pub fn is_empty(&self) -> bool {
533 self.len() == 0
534 }
535
536 #[must_use]
538 pub fn row_len(&self) -> usize {
539 self.shape.iter().skip(1).product()
540 }
541
542 #[must_use]
547 pub fn get_f32(&self, index: usize) -> Option<f32> {
548 let size = self.dtype.size();
549 let start = index.checked_mul(size)?;
550 let chunk = self.raw.get(start..start.checked_add(size)?)?;
551 Some(match self.dtype {
552 Dtype::Bf16 => bf16_bits_to_f32(u16::from_le_bytes([chunk[0], chunk[1]])),
553 Dtype::F32 => f32::from_le_bytes([chunk[0], chunk[1], chunk[2], chunk[3]]),
554 })
555 }
556
557 #[must_use]
566 pub fn copy_row_f32(&self, row: usize, out: &mut [f32]) -> bool {
567 let row_len = self.row_len();
568 if row_len == 0 || out.len() != row_len {
569 return false;
570 }
571 let Some(base) = row.checked_mul(row_len) else {
572 return false;
573 };
574 if base.checked_add(row_len).is_none_or(|end| end > self.len()) {
575 return false;
576 }
577 for (offset, slot) in out.iter_mut().enumerate() {
578 match self.get_f32(base + offset) {
580 Some(value) => *slot = value,
581 None => return false,
582 }
583 }
584 true
585 }
586
587 #[must_use]
589 pub const fn as_bytes(&self) -> &'a [u8] {
590 self.raw
591 }
592
593 #[must_use]
595 pub const fn as_ref(&self) -> TensorViewRef<'a> {
596 TensorViewRef {
597 dtype: self.dtype,
598 raw: self.raw,
599 }
600 }
601}
602
603impl TensorViewRef<'_> {
604 #[must_use]
606 pub const fn dtype(&self) -> Dtype {
607 self.dtype
608 }
609
610 #[must_use]
612 pub fn len(&self) -> usize {
613 self.raw.len() / self.dtype.size()
614 }
615
616 #[must_use]
618 pub fn is_empty(&self) -> bool {
619 self.len() == 0
620 }
621}
622
623#[cfg(test)]
624mod tests {
625 use super::*;
626
627 fn build(parts: &[(&str, Dtype, &[usize], &[u8])]) -> Vec<u8> {
629 let mut directory = serde_json::Map::new();
630 let mut payload = Vec::new();
631 for (name, dtype, shape, bytes) in parts {
632 let begin = payload.len();
633 payload.extend_from_slice(bytes);
634 directory.insert(
635 (*name).to_owned(),
636 serde_json::json!({
637 "dtype": dtype.as_str(),
638 "shape": shape,
639 "data_offsets": [begin, payload.len()],
640 }),
641 );
642 }
643 assemble(&Value::Object(directory), &payload)
644 }
645
646 fn assemble(header: &Value, payload: &[u8]) -> Vec<u8> {
647 let header_bytes = serde_json::to_vec(header).expect("header serializes");
648 let mut out = Vec::new();
649 out.extend_from_slice(&(header_bytes.len() as u64).to_le_bytes());
650 out.extend_from_slice(&header_bytes);
651 out.extend_from_slice(payload);
652 out
653 }
654
655 fn bf16_payload(values: &[u16]) -> Vec<u8> {
656 values.iter().flat_map(|v| v.to_le_bytes()).collect()
657 }
658
659 #[test]
660 fn parses_a_two_tensor_directory() {
661 let buffer = build(&[
662 ("a", Dtype::Bf16, &[2, 2], &bf16_payload(&[0, 1, 2, 3])),
663 ("b", Dtype::F32, &[2], &1.0f32.to_le_bytes().repeat(2)),
664 ]);
665 let index = SafetensorsIndex::parse(&buffer).expect("parses");
666
667 assert_eq!(index.len(), 2);
668 assert_eq!(index.names().collect::<Vec<_>>(), vec!["a", "b"]);
669 let a = index.entry("a").expect("entry a");
670 assert_eq!(a.dtype, Dtype::Bf16);
671 assert_eq!(a.shape, vec![2, 2]);
672 assert_eq!(a.element_count(), 4);
673 assert_eq!(a.byte_len(), 8);
674 assert_eq!(a.row_len(), 2);
675 assert_eq!(index.total_tensor_bytes(), 16);
676 }
677
678 #[test]
679 fn metadata_key_is_not_a_tensor() {
680 let mut directory = serde_json::Map::new();
681 directory.insert(
682 "__metadata__".to_owned(),
683 serde_json::json!({"format": "pt"}),
684 );
685 directory.insert(
686 "w".to_owned(),
687 serde_json::json!({"dtype": "F32", "shape": [1], "data_offsets": [0, 4]}),
688 );
689 let buffer = assemble(&Value::Object(directory), &1.0f32.to_le_bytes());
690 let index = SafetensorsIndex::parse(&buffer).expect("parses");
691 assert_eq!(index.len(), 1);
692 assert!(index.entry("__metadata__").is_none());
693 }
694
695 #[test]
696 fn widening_bf16_is_exact_for_representable_values() {
697 for bits in [0x0000u16, 0x3f80, 0xbf80, 0x7f80, 0xff80, 0x0001, 0x8000] {
700 let widened = bf16_bits_to_f32(bits);
701 assert_eq!(widened.to_bits() >> 16, u32::from(bits));
702 assert_eq!(widened.to_bits() & 0x0000_ffff, 0);
703 }
704 assert_eq!(bf16_bits_to_f32(0x3f80), 1.0);
705 assert_eq!(bf16_bits_to_f32(0xbf80), -1.0);
706 assert_eq!(bf16_bits_to_f32(0x0000), 0.0);
707 assert!(bf16_bits_to_f32(0x7f80).is_infinite());
708 assert!(bf16_bits_to_f32(0x7fc0).is_nan());
709 }
710
711 #[test]
712 fn bf16_widening_round_trips_every_bit_pattern() {
713 for bits in 0..=u16::MAX {
716 let widened = bf16_bits_to_f32(bits);
717 assert_eq!(
718 (widened.to_bits() >> 16) as u16,
719 bits,
720 "bit pattern {bits:#06x} did not survive widening"
721 );
722 let exponent = bits & 0x7f80;
723 let mantissa = bits & 0x007f;
724 if exponent == 0x7f80 && mantissa != 0 {
725 assert!(widened.is_nan(), "{bits:#06x} should widen to NaN");
726 } else {
727 assert!(!widened.is_nan(), "{bits:#06x} should not widen to NaN");
728 }
729 }
730 }
731
732 #[test]
733 fn reads_elements_and_rows_without_materializing() {
734 let payload = bf16_payload(&[0x3f80, 0xbf80, 0x4000, 0xc000]);
735 let buffer = build(&[("w", Dtype::Bf16, &[2, 2], &payload)]);
736 let index = SafetensorsIndex::parse(&buffer).expect("parses");
737 let view = index.view("w", &buffer).expect("view");
738
739 assert_eq!(view.len(), 4);
740 assert_eq!(view.row_len(), 2);
741 assert_eq!(view.get_f32(0), Some(1.0));
742 assert_eq!(view.get_f32(1), Some(-1.0));
743 assert_eq!(view.get_f32(3), Some(-2.0));
744 assert_eq!(view.get_f32(4), None);
745
746 let mut row = [0.0f32; 2];
747 assert!(view.copy_row_f32(1, &mut row));
748 assert_eq!(row, [2.0, -2.0]);
749
750 assert!(!view.copy_row_f32(2, &mut row));
752 let mut wrong = [0.0f32; 3];
753 assert!(!view.copy_row_f32(0, &mut wrong));
754 }
755
756 #[test]
757 fn refuses_a_truncated_file() {
758 assert_eq!(
759 SafetensorsIndex::parse(&[0u8; 4]).expect_err("must refuse"),
760 WeightsError::TooShortForHeader { len: 4 }
761 );
762 }
763
764 #[test]
765 fn refuses_an_absurd_header_length() {
766 let mut buffer = u64::MAX.to_le_bytes().to_vec();
767 buffer.extend_from_slice(b"{}");
768 let error = SafetensorsIndex::parse(&buffer).expect_err("must refuse");
769 assert!(matches!(error, WeightsError::HeaderLengthOutOfRange { .. }));
770 }
771
772 #[test]
773 fn refuses_a_header_longer_than_the_file() {
774 let mut buffer = 4096u64.to_le_bytes().to_vec();
775 buffer.extend_from_slice(b"{}");
776 let error = SafetensorsIndex::parse(&buffer).expect_err("must refuse");
777 assert!(matches!(error, WeightsError::HeaderLengthOutOfRange { .. }));
778 }
779
780 #[test]
781 fn refuses_malformed_json_and_non_objects() {
782 let mut buffer = 5u64.to_le_bytes().to_vec();
783 buffer.extend_from_slice(b"{ not");
784 assert!(matches!(
785 SafetensorsIndex::parse(&buffer).expect_err("must refuse"),
786 WeightsError::HeaderNotJson { .. }
787 ));
788
789 let buffer = assemble(&serde_json::json!([1, 2]), &[]);
790 assert_eq!(
791 SafetensorsIndex::parse(&buffer).expect_err("must refuse"),
792 WeightsError::HeaderNotObject
793 );
794 }
795
796 #[test]
797 fn refuses_an_unsupported_dtype() {
798 let buffer = assemble(
799 &serde_json::json!({"w": {"dtype": "I64", "shape": [1], "data_offsets": [0, 8]}}),
800 &[0u8; 8],
801 );
802 assert_eq!(
803 SafetensorsIndex::parse(&buffer).expect_err("must refuse"),
804 WeightsError::UnsupportedDtype {
805 name: "w".to_owned(),
806 raw: "I64".to_owned(),
807 }
808 );
809 }
810
811 #[test]
812 fn refuses_a_span_past_the_payload() {
813 let buffer = assemble(
815 &serde_json::json!({"w": {"dtype": "F32", "shape": [16], "data_offsets": [0, 64]}}),
816 &[0u8; 8],
817 );
818 let error = SafetensorsIndex::parse(&buffer).expect_err("must refuse");
819 assert!(matches!(error, WeightsError::SpanOutOfBounds { .. }));
820 }
821
822 #[test]
823 fn refuses_reversed_offsets() {
824 let buffer = assemble(
825 &serde_json::json!({"w": {"dtype": "F32", "shape": [1], "data_offsets": [8, 4]}}),
826 &[0u8; 8],
827 );
828 let error = SafetensorsIndex::parse(&buffer).expect_err("must refuse");
829 assert!(matches!(error, WeightsError::SpanOutOfBounds { .. }));
830 }
831
832 #[test]
833 fn refuses_a_shape_that_disagrees_with_its_span() {
834 let buffer = assemble(
836 &serde_json::json!({"w": {"dtype": "F32", "shape": [4], "data_offsets": [0, 8]}}),
837 &[0u8; 8],
838 );
839 let error = SafetensorsIndex::parse(&buffer).expect_err("must refuse");
840 assert!(
841 matches!(
842 error,
843 WeightsError::ShapeSpanMismatch {
844 expected_bytes: 16,
845 actual_bytes: 8,
846 ..
847 }
848 ),
849 "wrong error: {error}"
850 );
851 }
852
853 #[test]
854 fn refuses_a_shape_that_overflows() {
855 let huge = usize::MAX;
856 let buffer = assemble(
857 &serde_json::json!({
858 "w": {"dtype": "F32", "shape": [huge, huge], "data_offsets": [0, 8]}
859 }),
860 &[0u8; 8],
861 );
862 let error = SafetensorsIndex::parse(&buffer).expect_err("must refuse");
863 assert!(matches!(error, WeightsError::ShapeOverflow { .. }));
864 }
865
866 #[test]
867 fn view_refuses_a_buffer_that_is_not_the_parsed_one() {
868 let buffer = build(&[("w", Dtype::F32, &[1], &1.0f32.to_le_bytes())]);
869 let index = SafetensorsIndex::parse(&buffer).expect("parses");
870 assert!(index.view("w", &buffer[..4]).is_none());
871 assert!(index.view("missing", &buffer).is_none());
872 }
873}
874
875#[derive(Debug)]
884pub struct SafetensorsFile {
885 mapping: ftts_kernels::mmap::MappedFile,
886 index: SafetensorsIndex,
887}
888
889impl SafetensorsFile {
890 pub fn open(path: impl AsRef<std::path::Path>) -> Result<Self, OpenError> {
897 let mapping = ftts_kernels::mmap::MappedFile::open(path).map_err(OpenError::Io)?;
898 let index = SafetensorsIndex::parse(mapping.as_slice()).map_err(OpenError::Weights)?;
899 Ok(Self { mapping, index })
900 }
901
902 #[cfg(not(unix))] pub fn from_bytes(bytes: Vec<u8>) -> Result<Self, OpenError> {
909 let mapping = ftts_kernels::mmap::MappedFile::from_bytes(bytes);
910 let index = SafetensorsIndex::parse(mapping.as_slice()).map_err(OpenError::Weights)?;
911 Ok(Self { mapping, index })
912 }
913
914 pub fn advise_random(&self) {
919 self.mapping.advise_random();
920 }
921
922 #[must_use]
924 pub const fn index(&self) -> &SafetensorsIndex {
925 &self.index
926 }
927
928 #[must_use]
930 pub const fn mapped_len(&self) -> usize {
931 self.mapping.len()
932 }
933
934 #[must_use]
936 pub fn view(&self, name: &str) -> Option<TensorView<'_>> {
937 self.index.view(name, self.mapping.as_slice())
938 }
939}
940
941#[derive(Debug)]
943pub enum OpenError {
944 Io(std::io::Error),
946 Weights(WeightsError),
948}
949
950impl fmt::Display for OpenError {
951 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
952 match self {
953 Self::Io(error) => write!(f, "cannot open checkpoint: {error}"),
954 Self::Weights(error) => write!(f, "{error}"),
955 }
956 }
957}
958
959impl std::error::Error for OpenError {
960 fn source(&self) -> Option<&(dyn std::error::Error + 'static)> {
961 match self {
962 Self::Io(error) => Some(error),
963 Self::Weights(error) => Some(error),
964 }
965 }
966}