1use std::fs::File;
12use std::io::Read;
13use std::path::Path;
14
15use ztensor::{Error, Result, Source, Store, Vocabulary};
16
17pub const FORMATS: &[&str] = &["gguf", "hdf5", "npz", "onnx", "pt", "safetensors", "zt"];
24
25pub fn detect(path: impl AsRef<Path>) -> Result<&'static str> {
27 let path = path.as_ref();
28 let mut file = File::open(path)?;
31 let mut head = [0u8; 9];
32 let mut n = 0;
33 while n < head.len() {
34 match file.read(&mut head[n..])? {
35 0 => break,
36 got => n += got,
37 }
38 }
39 let head = &head[..n];
40
41 if head.len() >= 8 && head[..8] == ztensor::format::MAGIC {
42 return Ok("zt");
43 }
44 if head.starts_with(b"GGUF") {
45 return Ok("gguf");
46 }
47 if head.len() >= 8 && &head[..8] == b"\x89HDF\r\n\x1a\n" {
48 return Ok("hdf5");
49 }
50 if head.starts_with(b"PK\x03\x04") {
51 #[cfg(any(feature = "pickle", feature = "npz"))]
52 {
53 let is_pt = zip::ZipArchive::new(File::open(path)?)
54 .ok()
55 .map(|z| z.file_names().any(|n| n.ends_with("data.pkl")))
56 .unwrap_or(false);
57 return Ok(if is_pt { "pt" } else { "npz" });
58 }
59 #[cfg(not(any(feature = "pickle", feature = "npz")))]
60 return Err(Error::Unsupported(
61 "zip-container formats (.pt/.npz) are not compiled in".into(),
62 ));
63 }
64 if head.len() >= 9 && head[8] == b'{' {
65 let header_len = u64::from_le_bytes(head[..8].try_into().unwrap());
66 if header_len > 0 && header_len < (100 << 20) {
67 return Ok("safetensors");
68 }
69 }
70 if path.extension().is_some_and(|e| e == "onnx") {
71 return Ok("onnx");
72 }
73 Err(Error::Unsupported(format!(
74 "cannot detect the format of {}",
75 path.display()
76 )))
77}
78
79#[derive(Clone, Default)]
82pub struct Open {
83 vocab: Option<Vocabulary>,
84 map: Option<bool>,
85}
86
87pub fn options() -> Open {
89 Open::default()
90}
91
92impl Open {
93 pub fn vocabulary(mut self, vocab: &Vocabulary) -> Self {
94 self.vocab = Some(vocab.clone());
95 self
96 }
97
98 pub fn map(mut self, map: bool) -> Self {
101 self.map = Some(map);
102 self
103 }
104
105 fn mapping(&self) -> bool {
106 self.map.unwrap_or(true)
107 }
108
109 pub fn open(self, path: impl AsRef<Path>) -> Result<Source> {
111 let path = path.as_ref();
112 let format = detect(path)?;
113
114 if format == "zt" {
115 let mut opts = ztensor::Source::options().map(self.mapping());
116 if let Some(vocab) = &self.vocab {
117 opts = opts.vocabulary(vocab);
118 }
119 return opts.open(path);
120 }
121
122 let store = if self.mapping() {
123 Store::map(path, format)?
124 } else {
125 Store::index(path, format)?
126 };
127
128 #[allow(unused_variables)]
129 let projection = match format {
130 #[cfg(feature = "safetensors")]
131 "safetensors" => crate::safetensors::project(&store)?,
132 #[cfg(feature = "gguf")]
133 "gguf" => crate::gguf::project(&store)?,
134 #[cfg(feature = "npz")]
135 "npz" => crate::npz::project(&store)?,
136 #[cfg(feature = "pickle")]
137 "pt" => crate::pt::project(&store)?,
138 #[cfg(feature = "hdf5")]
139 "hdf5" => crate::hdf5::project(&store)?,
140 #[cfg(feature = "onnx")]
141 "onnx" => crate::onnx::project(&store)?,
142 other => {
143 return Err(Error::Unsupported(format!(
144 "{other} support is not compiled in (enable the matching \
145 ztensor-compat feature)"
146 )))
147 }
148 };
149 projection.into_source(store, self.vocab.as_ref())
150 }
151
152 pub fn open_all(self, paths: &[impl AsRef<Path>]) -> Result<Source> {
164 let mut sources = Vec::with_capacity(paths.len());
165 for path in paths {
166 sources.push(self.clone().open(path.as_ref())?);
167 }
168 Source::merge(sources)
169 }
170}
171
172pub fn open(path: impl AsRef<Path>) -> Result<Source> {
174 options().open(path)
175}
176
177pub fn index(path: impl AsRef<Path>) -> Result<Source> {
179 options().map(false).open(path)
180}
181
182pub fn open_all(paths: &[impl AsRef<Path>]) -> Result<Source> {
184 options().open_all(paths)
185}
186
187pub fn index_all(paths: &[impl AsRef<Path>]) -> Result<Source> {
189 options().map(false).open_all(paths)
190}