Skip to main content

ztensor_compat/
detect.rs

1//! Opening a tensor file of any supported format.
2//!
3//! Detection is by magic bytes wherever the format has them; ONNX (protobuf,
4//! magic-less) falls back to the file extension. Formats whose feature is
5//! disabled are reported as such, never misdetected as something else.
6//!
7//! What comes back is an ordinary [`Source`]. There is no per-format type to
8//! learn: a projected safetensors file and a canonical `.zt` file answer the
9//! same questions, and differ only in what they can honestly say yes to.
10
11use std::fs::File;
12use std::io::Read;
13use std::path::Path;
14
15use ztensor::{Error, Result, Source, Store, Vocabulary};
16
17/// Every label [`detect`] can return.
18///
19/// Enumerable because a consumer usually has a table keyed by these, such as a
20/// display name or an enum of its own. A table that is missing a row says
21/// nothing when the list grows. Checking against this turns that silence into
22/// a failing test.
23pub const FORMATS: &[&str] = &["gguf", "hdf5", "npz", "onnx", "pt", "safetensors", "zt"];
24
25/// Sniffs the format of a tensor file. Returns one of [`FORMATS`].
26pub fn detect(path: impl AsRef<Path>) -> Result<&'static str> {
27    let path = path.as_ref();
28    // A single read() may return fewer bytes than asked for; fill the buffer
29    // so a short read cannot cause a mis-detection.
30    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/// How to open. Mirrors [`ztensor::read::Options`], minus shard resolution:
80/// no foreign format has a shard table.
81#[derive(Clone, Default)]
82pub struct Open {
83    vocab: Option<Vocabulary>,
84    map: Option<bool>,
85}
86
87/// Opening options: a vocabulary to read with, and whether to map.
88pub 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    /// Map the files (the default). With `false`, files are opened but not
99    /// mapped: metadata and addresses are available, borrowed reads are not.
100    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    /// Opens one file of any supported format.
110    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    /// Opens several files as one name space.
153    ///
154    /// What a sharded snapshot is: `model-00001-of-00003.safetensors` and its
155    /// siblings are each a whole file that describes itself, and the index
156    /// beside them is a naming convention outside the format. So this is a
157    /// list of paths and nothing more. The caller decides which files belong
158    /// together, because the files themselves never said.
159    ///
160    /// The set may mix formats. Nothing here requires them to match: what
161    /// makes it a model is that the names do not collide, which the merge
162    /// checks.
163    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
172/// Opens a tensor file of any supported format.
173pub fn open(path: impl AsRef<Path>) -> Result<Source> {
174    options().open(path)
175}
176
177/// Opens without mapping: metadata and addresses only.
178pub fn index(path: impl AsRef<Path>) -> Result<Source> {
179    options().map(false).open(path)
180}
181
182/// Opens several files as one name space. See [`Open::open_all`].
183pub fn open_all(paths: &[impl AsRef<Path>]) -> Result<Source> {
184    options().open_all(paths)
185}
186
187/// Indexes several files as one name space, mapping none of them.
188pub fn index_all(paths: &[impl AsRef<Path>]) -> Result<Source> {
189    options().map(false).open_all(paths)
190}