Skip to main content

datui_lib/formats/
model_files.rs

1//! SafeTensors and GGUF model files as a table of their tensors. Only the header is
2//! read, so a 70 GB checkpoint opens like a 7 MB one. Both parsers are hand-written over
3//! a bounded reader, checking every stated length before allocating or skipping, so a
4//! hostile header errors rather than allocating gigabytes or panicking. Rows are a
5//! small eager frame made lazy; metadata and totals go to a [`ModelSummary`] for Info's
6//! Model tab.
7
8use std::io::Read;
9use std::path::{Path, PathBuf};
10
11use color_eyre::Result;
12use color_eyre::eyre::eyre;
13use polars::prelude::*;
14
15use crate::FileFormat;
16use crate::error_display::FileError;
17
18/// What datui does with a SafeTensors file: see [`crate::formats::readers`].
19pub(crate) const SAFETENSORS: crate::formats::readers::Reader = crate::formats::readers::Reader {
20    scan,
21    signatures: &[crate::formats::readers::Signature {
22        says: |head, _| looks_like_safetensors(head),
23        kind: crate::formats::readers::Kind::Magic,
24        trusted: crate::formats::readers::EVERYWHERE,
25    }],
26    ..crate::formats::readers::BASE
27};
28
29/// What datui does with a GGUF file: see [`crate::formats::readers`].
30pub(crate) const GGUF: crate::formats::readers::Reader = crate::formats::readers::Reader {
31    scan,
32    signatures: &[crate::formats::readers::Signature {
33        says: |head, _| looks_like_gguf(head),
34        kind: crate::formats::readers::Kind::Magic,
35        trusted: crate::formats::readers::EVERYWHERE,
36    }],
37    ..crate::formats::readers::BASE
38};
39
40/// The largest SafeTensors header read: the limit the reference implementation sets.
41pub const MAX_SAFETENSORS_HEADER: u64 = 100_000_000;
42/// The largest `model.safetensors.index.json` read. Real ones are a few hundred KB.
43const MAX_INDEX_JSON: u64 = 64 * 1024 * 1024;
44/// Where a GGUF header must have ended. Real headers, vocabulary and merges included,
45/// are tens of MB; a length that reaches past this is corruption.
46const MAX_GGUF_HEADER: u64 = 1024 * 1024 * 1024;
47/// The longest single GGUF string kept. A chat template is a few KB.
48const MAX_GGUF_STRING: u64 = 16 * 1024 * 1024;
49/// The most dimensions a tensor may have. GGML uses four.
50const MAX_DIMS: usize = 8;
51/// The most tensors or key/value pairs one GGUF file may declare. Real files have a
52/// few thousand tensors and a few dozen pairs; each kept tensor costs about 150 bytes.
53const MAX_GGUF_COUNT: u64 = 1 << 20;
54/// How deep arrays of arrays are followed.
55const MAX_ARRAY_DEPTH: u32 = 4;
56/// An array is listed in full up to this many items, and summarized by length after.
57const LIST_ITEMS_SHOWN: u64 = 16;
58/// A string inside a listed array is cut to this many characters.
59const LIST_ITEM_CHARS: usize = 120;
60/// The most shard files an index or a directory may name.
61const MAX_SHARDS: usize = 100_000;
62
63/// Which of the two formats a model file is.
64#[derive(Debug, Clone, Copy, PartialEq, Eq)]
65pub enum ModelKind {
66    SafeTensors,
67    Gguf { version: u32 },
68}
69
70impl ModelKind {
71    pub fn label(self) -> String {
72        match self {
73            ModelKind::SafeTensors => "SafeTensors".to_string(),
74            ModelKind::Gguf { version } => format!("GGUF v{version}"),
75        }
76    }
77}
78
79/// One metadata value, as the Model tab shows it.
80#[derive(Debug, Clone, PartialEq)]
81pub enum MetaValue {
82    /// A scalar, or a string kept whole (a chat template is shown in full).
83    Text(String),
84    /// An array: listed when short, otherwise only its length and what it holds.
85    List {
86        /// What the items are, plural: `strings`, `integers`.
87        of: &'static str,
88        len: u64,
89        /// Every item, when the array is short enough to list; empty otherwise.
90        items: Vec<String>,
91    },
92}
93
94/// What a model's header says beyond its tensors, and their totals.
95#[derive(Debug, Clone, PartialEq)]
96pub struct ModelSummary {
97    pub kind: ModelKind,
98    /// The files the tensors come from.
99    pub files: usize,
100    pub tensors: usize,
101    /// The sum of every tensor's parameter count.
102    pub params: u64,
103    /// The sum of every tensor's bytes, where they are known.
104    pub bytes: u64,
105    /// Each dtype or quantization type, its tensors and its parameters, most
106    /// parameters first.
107    pub types: Vec<TypeShare>,
108    /// Key and value, in the order the file has them. Across several files the first
109    /// file to name a key gives its value.
110    pub metadata: Vec<(String, MetaValue)>,
111}
112
113/// One dtype's share of a model.
114#[derive(Debug, Clone, PartialEq)]
115pub struct TypeShare {
116    pub name: String,
117    pub tensors: usize,
118    pub params: u64,
119}
120
121/// One tensor, as a header describes it.
122#[derive(Debug, Clone, PartialEq)]
123pub struct Tensor {
124    pub name: String,
125    /// The SafeTensors dtype or the GGML type name.
126    pub dtype: String,
127    pub shape: Vec<u64>,
128    /// `None` when the product of the shape overflows.
129    pub params: Option<u64>,
130    /// `None` for a GGML type datui does not know the size of.
131    pub bytes: Option<u64>,
132    /// SafeTensors: the start of `data_offsets`. GGUF: the offset as written, from the
133    /// start of the tensor data.
134    pub offset: u64,
135    /// SafeTensors only: the end of `data_offsets`.
136    pub offset_end: Option<u64>,
137}
138
139/// Metadata as key and value, in the order the file has them.
140pub type Metadata = Vec<(String, MetaValue)>;
141
142/// One file's header.
143#[derive(Debug, Clone, PartialEq)]
144pub struct Header {
145    pub kind: ModelKind,
146    pub tensors: Vec<Tensor>,
147    pub metadata: Vec<(String, MetaValue)>,
148}
149
150/// A reader that knows where the header must end and refuses any length past it.
151struct Bounded<R> {
152    inner: R,
153    pos: u64,
154    end: u64,
155    big_endian: bool,
156}
157
158impl<R: Read> Bounded<R> {
159    fn left(&self) -> u64 {
160        self.end.saturating_sub(self.pos)
161    }
162
163    /// Fail unless `n` more bytes fit before the end.
164    fn need(&self, n: u64, what: &str) -> Result<()> {
165        if n > self.left() {
166            return Err(eyre!(
167                "{what} runs past the end of the GGUF header ({n} bytes, {} left)",
168                self.left()
169            ));
170        }
171        Ok(())
172    }
173
174    fn fill<const N: usize>(&mut self, what: &str) -> Result<[u8; N]> {
175        self.need(N as u64, what)?;
176        let mut buf = [0u8; N];
177        self.inner
178            .read_exact(&mut buf)
179            .map_err(|e| eyre!("cannot read {what} in the GGUF header: {e}"))?;
180        self.pos += N as u64;
181        Ok(buf)
182    }
183
184    fn u8(&mut self, what: &str) -> Result<u8> {
185        Ok(self.fill::<1>(what)?[0])
186    }
187
188    fn u16(&mut self, what: &str) -> Result<u16> {
189        let b = self.fill::<2>(what)?;
190        Ok(if self.big_endian {
191            u16::from_be_bytes(b)
192        } else {
193            u16::from_le_bytes(b)
194        })
195    }
196
197    fn u32(&mut self, what: &str) -> Result<u32> {
198        let b = self.fill::<4>(what)?;
199        Ok(if self.big_endian {
200            u32::from_be_bytes(b)
201        } else {
202            u32::from_le_bytes(b)
203        })
204    }
205
206    fn u64(&mut self, what: &str) -> Result<u64> {
207        let b = self.fill::<8>(what)?;
208        Ok(if self.big_endian {
209            u64::from_be_bytes(b)
210        } else {
211            u64::from_le_bytes(b)
212        })
213    }
214
215    /// A length-prefixed string, kept. Invalid UTF-8 is replaced, not refused.
216    fn string(&mut self, what: &str) -> Result<String> {
217        let len = self.u64(what)?;
218        if len > MAX_GGUF_STRING {
219            return Err(eyre!(
220                "{what} is {len} bytes, longer than datui reads in a GGUF header"
221            ));
222        }
223        self.need(len, what)?;
224        let mut buf = Vec::new();
225        (&mut self.inner)
226            .take(len)
227            .read_to_end(&mut buf)
228            .map_err(|e| eyre!("cannot read {what} in the GGUF header: {e}"))?;
229        if buf.len() as u64 != len {
230            return Err(eyre!("{what} in the GGUF header is cut short"));
231        }
232        self.pos += len;
233        Ok(String::from_utf8_lossy(&buf).into_owned())
234    }
235
236    /// Read past `n` bytes. Read rather than sought: the skips are a vocabulary's
237    /// strings, a few bytes each, and a seek would throw the read buffer away for each.
238    fn skip(&mut self, n: u64, what: &str) -> Result<()> {
239        self.need(n, what)?;
240        let skipped = std::io::copy(&mut (&mut self.inner).take(n), &mut std::io::sink())
241            .map_err(|e| eyre!("cannot read {what} in the GGUF header: {e}"))?;
242        if skipped != n {
243            return Err(eyre!("{what} in the GGUF header is cut short"));
244        }
245        self.pos += n;
246        Ok(())
247    }
248}
249
250/// Whether the first bytes of a file are a SafeTensors header: a little-endian length
251/// a header could have, then the `{` that opens its JSON.
252pub fn looks_like_safetensors(head: &[u8]) -> bool {
253    if head.len() < 9 {
254        return false;
255    }
256    let len = u64::from_le_bytes(head[..8].try_into().expect("eight bytes"));
257    (2..=MAX_SAFETENSORS_HEADER).contains(&len) && head[8] == b'{'
258}
259
260/// Whether the first bytes of a file are GGUF's magic.
261pub fn looks_like_gguf(head: &[u8]) -> bool {
262    head.starts_with(b"GGUF")
263}
264
265/// The header length a SafeTensors file's first eight bytes state, refused when it is
266/// more than the spec allows or than the `len`-byte file holds.
267fn safetensors_header_len(prefix: [u8; 8], len: u64) -> Result<u64> {
268    let header_len = u64::from_le_bytes(prefix);
269    if header_len > MAX_SAFETENSORS_HEADER {
270        return Err(eyre!(
271            "the SafeTensors header is {header_len} bytes, more than the {MAX_SAFETENSORS_HEADER} allowed"
272        ));
273    }
274    if header_len > len.saturating_sub(8) {
275        return Err(eyre!(
276            "the SafeTensors header claims {header_len} bytes and the file has {}",
277            len.saturating_sub(8)
278        ));
279    }
280    Ok(header_len)
281}
282
283/// Read one SafeTensors header from `reader`, which holds `len` bytes in all.
284pub fn read_safetensors<R: Read>(reader: R, len: u64) -> Result<Header> {
285    let mut reader = reader;
286    let mut prefix = [0u8; 8];
287    reader
288        .read_exact(&mut prefix)
289        .map_err(|_| eyre!("the file is shorter than its SafeTensors header length"))?;
290    let header_len = safetensors_header_len(prefix, len)?;
291    let mut json = Vec::new();
292    reader
293        .take(header_len)
294        .read_to_end(&mut json)
295        .map_err(|e| eyre!("cannot read the SafeTensors header: {e}"))?;
296    if json.len() as u64 != header_len {
297        return Err(eyre!("the SafeTensors header is cut short"));
298    }
299    parse_safetensors_json(&json, len.saturating_sub(8).saturating_sub(header_len))
300}
301
302/// The header's JSON without its length prefix; `data_len` bytes of tensor data follow,
303/// which every `data_offsets` must stay inside. Deserialized straight into what is kept
304/// (a hostile 100 MB header would be gigabytes as a `Value`), unknown fields skipped,
305/// keys in file order (the metadata's display order).
306fn parse_safetensors_json(json: &[u8], data_len: u64) -> Result<Header> {
307    let mut de = serde_json::Deserializer::from_slice(json);
308    let parsed = serde::Deserializer::deserialize_map(&mut de, StHeaderVisitor)
309        .and_then(|header| de.end().map(|()| header))
310        .map_err(|e| eyre!("the SafeTensors header is not valid: {e}"))?;
311    let (mut tensors, metadata) = parsed;
312    for t in &tensors {
313        if t.offset_end.is_some_and(|end| end > data_len) {
314            return Err(eyre!(
315                "tensor \"{}\" runs past the end of the file ({data_len} bytes of data)",
316                t.name
317            ));
318        }
319    }
320    // The order the data is in.
321    tensors.sort_by(|a, b| a.offset.cmp(&b.offset).then_with(|| a.name.cmp(&b.name)));
322    Ok(Header {
323        kind: ModelKind::SafeTensors,
324        tensors,
325        metadata,
326    })
327}
328
329/// One tensor's entry. Any other field is skipped, not kept.
330#[derive(serde::Deserialize)]
331struct StEntry {
332    dtype: String,
333    shape: StShape,
334    data_offsets: (u64, u64),
335}
336
337/// A shape, refused past [`MAX_DIMS`] before a longer list is stored.
338struct StShape(Vec<u64>);
339
340impl<'de> serde::Deserialize<'de> for StShape {
341    fn deserialize<D: serde::Deserializer<'de>>(d: D) -> std::result::Result<Self, D::Error> {
342        struct V;
343        impl<'de> serde::de::Visitor<'de> for V {
344            type Value = StShape;
345            fn expecting(&self, f: &mut std::fmt::Formatter) -> std::fmt::Result {
346                write!(f, "a list of at most {MAX_DIMS} dimensions")
347            }
348            fn visit_seq<A: serde::de::SeqAccess<'de>>(
349                self,
350                mut seq: A,
351            ) -> std::result::Result<StShape, A::Error> {
352                let mut dims = Vec::new();
353                while let Some(d) = seq.next_element::<u64>()? {
354                    if dims.len() == MAX_DIMS {
355                        return Err(serde::de::Error::custom("more dimensions than datui reads"));
356                    }
357                    dims.push(d);
358                }
359                Ok(StShape(dims))
360            }
361        }
362        d.deserialize_seq(V)
363    }
364}
365
366/// A `__metadata__` value. The spec says text; a number or a bool is shown as written,
367/// and anything nested is passed over and named by what it is.
368struct StMetaValue(String);
369
370impl<'de> serde::Deserialize<'de> for StMetaValue {
371    fn deserialize<D: serde::Deserializer<'de>>(d: D) -> std::result::Result<Self, D::Error> {
372        struct V;
373        impl<'de> serde::de::Visitor<'de> for V {
374            type Value = StMetaValue;
375            fn expecting(&self, f: &mut std::fmt::Formatter) -> std::fmt::Result {
376                f.write_str("a metadata value")
377            }
378            fn visit_str<E>(self, v: &str) -> std::result::Result<StMetaValue, E> {
379                Ok(StMetaValue(v.to_string()))
380            }
381            fn visit_string<E>(self, v: String) -> std::result::Result<StMetaValue, E> {
382                Ok(StMetaValue(v))
383            }
384            fn visit_bool<E>(self, v: bool) -> std::result::Result<StMetaValue, E> {
385                Ok(StMetaValue(v.to_string()))
386            }
387            fn visit_i64<E>(self, v: i64) -> std::result::Result<StMetaValue, E> {
388                Ok(StMetaValue(v.to_string()))
389            }
390            fn visit_u64<E>(self, v: u64) -> std::result::Result<StMetaValue, E> {
391                Ok(StMetaValue(v.to_string()))
392            }
393            fn visit_f64<E>(self, v: f64) -> std::result::Result<StMetaValue, E> {
394                Ok(StMetaValue(v.to_string()))
395            }
396            fn visit_unit<E>(self) -> std::result::Result<StMetaValue, E> {
397                Ok(StMetaValue("null".to_string()))
398            }
399            fn visit_seq<A: serde::de::SeqAccess<'de>>(
400                self,
401                mut seq: A,
402            ) -> std::result::Result<StMetaValue, A::Error> {
403                while seq.next_element::<serde::de::IgnoredAny>()?.is_some() {}
404                Ok(StMetaValue("[array]".to_string()))
405            }
406            fn visit_map<A: serde::de::MapAccess<'de>>(
407                self,
408                mut map: A,
409            ) -> std::result::Result<StMetaValue, A::Error> {
410                while map
411                    .next_entry::<serde::de::IgnoredAny, serde::de::IgnoredAny>()?
412                    .is_some()
413                {}
414                Ok(StMetaValue("{object}".to_string()))
415            }
416        }
417        d.deserialize_any(V)
418    }
419}
420
421/// `__metadata__`, in the order the file has it.
422struct StMetadata(Metadata);
423
424impl<'de> serde::Deserialize<'de> for StMetadata {
425    fn deserialize<D: serde::Deserializer<'de>>(d: D) -> std::result::Result<Self, D::Error> {
426        struct V;
427        impl<'de> serde::de::Visitor<'de> for V {
428            type Value = StMetadata;
429            fn expecting(&self, f: &mut std::fmt::Formatter) -> std::fmt::Result {
430                f.write_str("an object of metadata")
431            }
432            fn visit_map<A: serde::de::MapAccess<'de>>(
433                self,
434                mut map: A,
435            ) -> std::result::Result<StMetadata, A::Error> {
436                let mut out: Metadata = Vec::new();
437                while let Some((key, StMetaValue(value))) =
438                    map.next_entry::<String, StMetaValue>()?
439                {
440                    out.push((key, MetaValue::Text(value)));
441                }
442                Ok(StMetadata(out))
443            }
444        }
445        d.deserialize_map(V)
446    }
447}
448
449/// The whole header: its tensors, and `__metadata__`.
450struct StHeaderVisitor;
451
452impl<'de> serde::de::Visitor<'de> for StHeaderVisitor {
453    type Value = (Vec<Tensor>, Metadata);
454    fn expecting(&self, f: &mut std::fmt::Formatter) -> std::fmt::Result {
455        f.write_str("a JSON object of tensors")
456    }
457    fn visit_map<A: serde::de::MapAccess<'de>>(
458        self,
459        mut map: A,
460    ) -> std::result::Result<Self::Value, A::Error> {
461        use serde::de::Error;
462        let mut tensors = Vec::new();
463        let mut seen = std::collections::HashSet::new();
464        let mut metadata = None;
465        while let Some(name) = map.next_key::<String>()? {
466            if name == "__metadata__" {
467                if metadata.is_some() {
468                    return Err(A::Error::custom("__metadata__ appears twice"));
469                }
470                let StMetadata(m) = map
471                    .next_value()
472                    .map_err(|e| A::Error::custom(format!("__metadata__: {e}")))?;
473                metadata = Some(m);
474                continue;
475            }
476            let entry: StEntry = map
477                .next_value()
478                .map_err(|e| A::Error::custom(format!("tensor \"{name}\": {e}")))?;
479            if !seen.insert(name.clone()) {
480                return Err(A::Error::custom(format!("tensor \"{name}\" appears twice")));
481            }
482            let (start, end) = entry.data_offsets;
483            if end < start {
484                return Err(A::Error::custom(format!(
485                    "tensor \"{name}\" ends before it starts"
486                )));
487            }
488            let shape = entry.shape.0;
489            tensors.push(Tensor {
490                name,
491                dtype: entry.dtype,
492                params: product(&shape),
493                shape,
494                bytes: Some(end - start),
495                offset: start,
496                offset_end: Some(end),
497            });
498        }
499        Ok((tensors, metadata.unwrap_or_default()))
500    }
501}
502
503/// The product of a shape; 1 for a scalar, `None` on overflow.
504fn product(shape: &[u64]) -> Option<u64> {
505    shape.iter().try_fold(1u64, |acc, d| acc.checked_mul(*d))
506}
507
508/// A GGML type: its name, and how many elements a block of how many bytes holds.
509fn ggml_type(id: u32) -> Option<(&'static str, u64, u64)> {
510    Some(match id {
511        0 => ("F32", 1, 4),
512        1 => ("F16", 1, 2),
513        2 => ("Q4_0", 32, 18),
514        3 => ("Q4_1", 32, 20),
515        6 => ("Q5_0", 32, 22),
516        7 => ("Q5_1", 32, 24),
517        8 => ("Q8_0", 32, 34),
518        9 => ("Q8_1", 32, 36),
519        10 => ("Q2_K", 256, 84),
520        11 => ("Q3_K", 256, 110),
521        12 => ("Q4_K", 256, 144),
522        13 => ("Q5_K", 256, 176),
523        14 => ("Q6_K", 256, 210),
524        15 => ("Q8_K", 256, 292),
525        16 => ("IQ2_XXS", 256, 66),
526        17 => ("IQ2_XS", 256, 74),
527        18 => ("IQ3_XXS", 256, 98),
528        19 => ("IQ1_S", 256, 50),
529        20 => ("IQ4_NL", 32, 18),
530        21 => ("IQ3_S", 256, 110),
531        22 => ("IQ2_S", 256, 82),
532        23 => ("IQ4_XS", 256, 136),
533        24 => ("I8", 1, 1),
534        25 => ("I16", 1, 2),
535        26 => ("I32", 1, 4),
536        27 => ("I64", 1, 8),
537        28 => ("F64", 1, 8),
538        29 => ("IQ1_M", 256, 56),
539        30 => ("BF16", 1, 2),
540        // Repacked Q4_0 and IQ4_NL, since removed from GGML; files written by
541        // llama.cpp in late 2024 still carry them, with the same block sizes.
542        31 => ("Q4_0_4_4", 32, 18),
543        32 => ("Q4_0_4_8", 32, 18),
544        33 => ("Q4_0_8_8", 32, 18),
545        34 => ("TQ1_0", 256, 54),
546        35 => ("TQ2_0", 256, 66),
547        36 => ("IQ4_NL_4_4", 32, 18),
548        37 => ("IQ4_NL_4_8", 32, 18),
549        38 => ("IQ4_NL_8_8", 32, 18),
550        39 => ("MXFP4", 32, 17),
551        _ => return None,
552    })
553}
554
555/// Where tensor data starts when `general.alignment` does not say.
556const GGUF_DEFAULT_ALIGNMENT: u64 = 32;
557
558/// GGUF metadata value types.
559const GGUF_STRING: u32 = 8;
560const GGUF_ARRAY: u32 = 9;
561
562/// The size of a fixed-width GGUF value type; `None` for a string or an array.
563fn gguf_fixed_size(ty: u32) -> Option<u64> {
564    match ty {
565        0 | 1 | 7 => Some(1),
566        2 | 3 => Some(2),
567        4..=6 => Some(4),
568        10..=12 => Some(8),
569        _ => None,
570    }
571}
572
573/// What an array of `ty` holds, as its summary says it.
574fn gguf_items_noun(ty: u32) -> &'static str {
575    match ty {
576        0..=5 | 10 | 11 => "integers",
577        6 | 12 => "floats",
578        7 => "bools",
579        GGUF_STRING => "strings",
580        _ => "arrays",
581    }
582}
583
584/// One fixed-width value, as text.
585fn gguf_scalar<R: Read>(r: &mut Bounded<R>, ty: u32) -> Result<String> {
586    let what = "a metadata value";
587    Ok(match ty {
588        0 => r.u8(what)?.to_string(),
589        1 => (r.u8(what)? as i8).to_string(),
590        2 => r.u16(what)?.to_string(),
591        3 => (r.u16(what)? as i16).to_string(),
592        4 => r.u32(what)?.to_string(),
593        5 => (r.u32(what)? as i32).to_string(),
594        6 => f32::from_bits(r.u32(what)?).to_string(),
595        7 => (r.u8(what)? != 0).to_string(),
596        10 => r.u64(what)?.to_string(),
597        11 => (r.u64(what)? as i64).to_string(),
598        12 => f64::from_bits(r.u64(what)?).to_string(),
599        other => return Err(eyre!("unknown GGUF metadata type {other}")),
600    })
601}
602
603/// A value of type `ty`, read whole when it is kept and skipped where it is long.
604fn gguf_value<R: Read>(r: &mut Bounded<R>, ty: u32, depth: u32) -> Result<MetaValue> {
605    match ty {
606        GGUF_STRING => Ok(MetaValue::Text(r.string("a metadata string")?)),
607        GGUF_ARRAY => {
608            if depth >= MAX_ARRAY_DEPTH {
609                return Err(eyre!("GGUF arrays nest deeper than datui reads"));
610            }
611            let item_ty = r.u32("an array's type")?;
612            let len = r.u64("an array's length")?;
613            // Every item takes at least this much, so the length is checked against
614            // what is left before a single item is read.
615            let least = match item_ty {
616                GGUF_STRING => 8,
617                GGUF_ARRAY => 12,
618                t => gguf_fixed_size(t).ok_or_else(|| eyre!("unknown GGUF array type {t}"))?,
619            };
620            r.need(
621                len.checked_mul(least)
622                    .ok_or_else(|| eyre!("a GGUF array's length overflows"))?,
623                "an array",
624            )?;
625            let of = gguf_items_noun(item_ty);
626            let listed = len <= LIST_ITEMS_SHOWN && item_ty != GGUF_ARRAY;
627            if !listed {
628                skip_items(r, item_ty, len, depth)?;
629                return Ok(MetaValue::List {
630                    of,
631                    len,
632                    items: Vec::new(),
633                });
634            }
635            let mut items = Vec::with_capacity(len as usize);
636            for _ in 0..len {
637                let item = if item_ty == GGUF_STRING {
638                    let s = r.string("an array's string")?;
639                    let cut: String = s.chars().take(LIST_ITEM_CHARS).collect();
640                    if cut.len() < s.len() {
641                        format!("{cut}...")
642                    } else {
643                        cut
644                    }
645                } else {
646                    gguf_scalar(r, item_ty)?
647                };
648                items.push(item);
649            }
650            Ok(MetaValue::List { of, len, items })
651        }
652        t => Ok(MetaValue::Text(gguf_scalar(r, t)?)),
653    }
654}
655
656/// Step over `len` items of `ty` without keeping them.
657fn skip_items<R: Read>(r: &mut Bounded<R>, ty: u32, len: u64, depth: u32) -> Result<()> {
658    match ty {
659        GGUF_STRING => {
660            for _ in 0..len {
661                let n = r.u64("an array's string")?;
662                r.skip(n, "an array's string")?;
663            }
664        }
665        GGUF_ARRAY => {
666            for _ in 0..len {
667                gguf_value(r, GGUF_ARRAY, depth + 1)?;
668            }
669        }
670        t => {
671            let size = gguf_fixed_size(t).ok_or_else(|| eyre!("unknown GGUF array type {t}"))?;
672            r.skip(len.saturating_mul(size), "an array")?;
673        }
674    }
675    Ok(())
676}
677
678/// Read one GGUF header (versions 2 and 3, either byte order) from `reader`, which
679/// holds `len` bytes in all.
680pub fn read_gguf<R: Read>(reader: R, len: u64) -> Result<Header> {
681    let mut r = Bounded {
682        inner: reader,
683        pos: 0,
684        end: len.min(MAX_GGUF_HEADER),
685        big_endian: false,
686    };
687    let magic = r
688        .fill::<4>("the magic number")
689        .map_err(|_| eyre!("the file is too short to be GGUF"))?;
690    if &magic != b"GGUF" {
691        return Err(eyre!("not a GGUF file: it does not start with GGUF"));
692    }
693    let raw = r.fill::<4>("the version")?;
694    let mut version = u32::from_le_bytes(raw);
695    // The magic reads the same either way; a big-endian file's version does not.
696    if version & 0xFFFF == 0 {
697        r.big_endian = true;
698        version = u32::from_be_bytes(raw);
699    }
700    match version {
701        2 | 3 => {}
702        1 => return Err(eyre!("GGUF version 1 files are not supported")),
703        v => return Err(eyre!("GGUF version {v} is not one datui reads (2 or 3)")),
704    }
705    let tensor_count = r.u64("the tensor count")?;
706    let kv_count = r.u64("the metadata count")?;
707    // A key/value pair takes at least 12 bytes and a tensor's entry at least 24, so
708    // either count is bounded by what is left before anything is allocated for it.
709    if tensor_count > MAX_GGUF_COUNT || tensor_count.saturating_mul(24) > r.left() {
710        return Err(eyre!(
711            "{tensor_count} tensors cannot fit in the GGUF header"
712        ));
713    }
714    if kv_count > MAX_GGUF_COUNT || kv_count.saturating_mul(12) > r.left() {
715        return Err(eyre!(
716            "{kv_count} metadata entries cannot fit in the GGUF header"
717        ));
718    }
719    let mut metadata = Vec::with_capacity(kv_count as usize);
720    for _ in 0..kv_count {
721        let key = r.string("a metadata key")?;
722        let ty = r.u32("a metadata type")?;
723        let value = gguf_value(&mut r, ty, 0).map_err(|e| eyre!("{e} (in \"{key}\")"))?;
724        metadata.push((key, value));
725    }
726    let mut tensors = Vec::with_capacity(tensor_count as usize);
727    for _ in 0..tensor_count {
728        let name = r.string("a tensor name")?;
729        let n_dims = r.u32("a tensor's dimension count")? as usize;
730        if n_dims > MAX_DIMS {
731            return Err(eyre!(
732                "tensor \"{name}\" has {n_dims} dimensions, more than datui reads"
733            ));
734        }
735        let mut shape = Vec::with_capacity(n_dims);
736        for _ in 0..n_dims {
737            shape.push(r.u64("a tensor dimension")?);
738        }
739        let ty = r.u32("a tensor's type")?;
740        let offset = r.u64("a tensor's offset")?;
741        let params = product(&shape);
742        let (dtype, bytes) = match ggml_type(ty) {
743            Some((name, block, size)) => (
744                name.to_string(),
745                params
746                    .filter(|p| p % block == 0)
747                    .and_then(|p| (p / block).checked_mul(size)),
748            ),
749            None => (format!("type {ty}"), None),
750        };
751        tensors.push(Tensor {
752            name,
753            dtype,
754            shape,
755            params,
756            bytes,
757            offset,
758            offset_end: None,
759        });
760    }
761    // Tensor data starts at the next alignment multiple after the header, and every sized
762    // tensor must end inside the file: a truncated download is an error.
763    let alignment = metadata
764        .iter()
765        .find(|(k, _)| k == "general.alignment")
766        .and_then(|(_, v)| match v {
767            MetaValue::Text(t) => t.parse::<u64>().ok(),
768            MetaValue::List { .. } => None,
769        })
770        .filter(|a| a.is_power_of_two())
771        .unwrap_or(GGUF_DEFAULT_ALIGNMENT);
772    let data_start = r.pos.next_multiple_of(alignment);
773    let data_len = len.saturating_sub(data_start);
774    for t in &tensors {
775        let end = t.bytes.and_then(|b| t.offset.checked_add(b));
776        if t.bytes.is_some() && end.is_none_or(|end| end > data_len) {
777            return Err(eyre!(
778                "tensor \"{}\" runs past the end of the file ({data_len} bytes of data)",
779                t.name
780            ));
781        }
782    }
783    Ok(Header {
784        kind: ModelKind::Gguf { version },
785        tensors,
786        metadata,
787    })
788}
789
790/// Parse `bytes` as whichever model header it starts with: both parsers over a slice,
791/// with no file.
792#[cfg(test)]
793pub fn parse_header(bytes: &[u8]) -> Result<Header> {
794    if looks_like_gguf(bytes) {
795        read_gguf(bytes, bytes.len() as u64)
796    } else {
797        read_safetensors(bytes, bytes.len() as u64)
798    }
799}
800
801/// Read one file's header with the parser `format` names.
802fn read_file(path: &Path, format: FileFormat) -> Result<Header> {
803    let file = std::fs::File::open(path)?;
804    let len = file.metadata()?.len();
805    let reader = std::io::BufReader::new(file);
806    match format {
807        FileFormat::Gguf => read_gguf(reader, len),
808        _ => read_safetensors(reader, len),
809    }
810}
811
812/// A `model.safetensors.index.json`: the shard each tensor is in, and metadata. Any
813/// other field is skipped, not kept.
814#[derive(serde::Deserialize)]
815struct StIndex {
816    #[serde(default)]
817    metadata: Option<StMetadata>,
818    weight_map: std::collections::BTreeMap<String, String>,
819}
820
821/// Whether `path` is a SafeTensors index: `model.safetensors.index.json`.
822pub fn is_safetensors_index(path: &Path) -> bool {
823    path.file_name()
824        .and_then(|n| n.to_str())
825        .is_some_and(|n| n.to_ascii_lowercase().ends_with(".safetensors.index.json"))
826}
827
828/// The shards an index names, beside it and in name order, and the index's own
829/// metadata.
830fn read_index(path: &Path) -> Result<(Vec<PathBuf>, Metadata)> {
831    let named = |e: std::io::Error| crate::error_display::in_file(path, e.into());
832    let file = std::fs::File::open(path).map_err(named)?;
833    let len = file.metadata().map_err(named)?.len();
834    if len > MAX_INDEX_JSON {
835        return Err(FileError::new(
836            path,
837            format!("the index is {len} bytes, more than datui reads"),
838        )
839        .into());
840    }
841    let mut text = Vec::new();
842    file.take(MAX_INDEX_JSON)
843        .read_to_end(&mut text)
844        .map_err(named)?;
845    let (names, metadata) = parse_index(&text, &path.display().to_string())?;
846    let dir = path.parent().unwrap_or(Path::new(""));
847    Ok((names.iter().map(|name| dir.join(name)).collect(), metadata))
848}
849
850/// An index's shard names, each once and in name order, and its metadata. `named` is
851/// what errors call the index.
852fn parse_index(text: &[u8], named: &str) -> Result<(Vec<String>, Metadata)> {
853    let refused = |what: String| FileError::new(Path::new(named), what);
854    let index: StIndex = serde_json::from_slice(text)
855        .map_err(|e| refused(format!("not a SafeTensors index: {e}")))?;
856    let names: std::collections::BTreeSet<String> = index.weight_map.into_values().collect();
857    if names.len() > MAX_SHARDS {
858        return Err(refused("the index names too many shards".into()).into());
859    }
860    for name in &names {
861        // A shard is a file beside its index, never a path out of the directory.
862        let path = Path::new(name);
863        if path.components().count() != 1 || path.file_name().is_none() || name.contains('\\') {
864            return Err(refused(format!(
865                "the index names \"{name}\", which is not a file beside it"
866            ))
867            .into());
868        }
869    }
870    let metadata = index.metadata.map(|StMetadata(m)| m).unwrap_or_default();
871    Ok((names.into_iter().collect(), metadata))
872}
873
874/// Bytes of one remote object, fetched a range at a time: an HTTP server, or an object
875/// in a store.
876pub trait RangeSource {
877    /// Bytes `start..end` of the object, fewer only where the object ends first, and
878    /// the object's whole length. `start` is inside the object.
879    fn get(&mut self, start: u64, end: u64) -> std::result::Result<(Vec<u8>, u64), RangeError>;
880}
881
882/// Why a ranged read stopped.
883#[derive(Debug, Clone, PartialEq, Eq)]
884pub enum RangeError {
885    /// The server sent the whole file where a range was asked for: the header cannot be
886    /// read on its own, and the file is downloaded instead.
887    NoRanges,
888    /// Anything else, as the user is told it.
889    Failed(String),
890}
891
892impl From<color_eyre::Report> for RangeError {
893    fn from(e: color_eyre::Report) -> Self {
894        RangeError::Failed(e.to_string())
895    }
896}
897
898/// The first GGUF header read; each next is twice the last up to `MAX_RANGE`, so a few
899/// KB costs one request and a 5-10 MB vocabulary four or five (see
900/// `a_vocabulary_sized_gguf_header_takes_a_few_ranges`).
901pub const FIRST_GGUF_RANGE: u64 = 256 * 1024;
902/// The first SafeTensors read: the header length and usually all its JSON (shard
903/// headers are KB to tens of KB); longer takes one more request.
904pub const FIRST_SAFETENSORS_RANGE: u64 = 64 * 1024;
905/// The most one ranged request asks for.
906const MAX_RANGE: u64 = 16 * 1024 * 1024;
907/// The first read of a remote index; one larger than this takes a second.
908const FIRST_INDEX_RANGE: u64 = 1024 * 1024;
909
910/// `start..end` of `src`, checked: exactly the bytes asked for up to the object's end,
911/// and the same length the first answer gave, if there was one.
912fn fetch(
913    src: &mut dyn RangeSource,
914    start: u64,
915    end: u64,
916    known_len: Option<u64>,
917) -> std::result::Result<(Vec<u8>, u64), RangeError> {
918    let (bytes, len) = src.get(start, end)?;
919    if known_len.is_some_and(|known| known != len) {
920        return Err(RangeError::Failed(format!(
921            "the file changed size while its header was read ({} then {len} bytes)",
922            known_len.unwrap_or_default()
923        )));
924    }
925    let want = end.min(len).saturating_sub(start);
926    if bytes.len() as u64 != want {
927        return Err(RangeError::Failed(format!(
928            "asked for bytes {start}..{} and got {} bytes",
929            end.min(len),
930            bytes.len()
931        )));
932    }
933    Ok((bytes, len))
934}
935
936/// A remote object read forward through ranged requests that grow as the read goes on,
937/// and never past `limit`. One range is held at a time.
938struct Ranged<'a> {
939    src: &'a mut dyn RangeSource,
940    len: u64,
941    limit: u64,
942    buf: Vec<u8>,
943    buf_start: u64,
944    pos: u64,
945    next: u64,
946    stop: &'a dyn Fn() -> bool,
947}
948
949impl Read for Ranged<'_> {
950    fn read(&mut self, out: &mut [u8]) -> std::io::Result<usize> {
951        if self.pos >= self.limit || out.is_empty() {
952            return Ok(0);
953        }
954        let buf_end = self.buf_start + self.buf.len() as u64;
955        if self.pos < self.buf_start || self.pos >= buf_end {
956            if (self.stop)() {
957                return Err(std::io::Error::other("cancelled"));
958            }
959            let end = self.pos.saturating_add(self.next).min(self.limit);
960            let (bytes, _) = fetch(self.src, self.pos, end, Some(self.len)).map_err(|e| {
961                std::io::Error::other(match e {
962                    RangeError::NoRanges => "the server stopped serving byte ranges".to_string(),
963                    RangeError::Failed(message) => message,
964                })
965            })?;
966            self.buf = bytes;
967            self.buf_start = self.pos;
968            self.next = (self.next * 2).min(MAX_RANGE);
969        }
970        let at = (self.pos - self.buf_start) as usize;
971        let n = out.len().min(self.buf.len() - at);
972        out[..n].copy_from_slice(&self.buf[at..at + n]);
973        self.pos += n as u64;
974        Ok(n)
975    }
976}
977
978/// Read one model header from `src` with ranged requests: SafeTensors' first
979/// [`FIRST_SAFETENSORS_RANGE`] plus any remaining JSON; GGUF in growing ranges until the
980/// tensor infos end. File-reader bounds all hold; `stop` is checked before each request.
981pub fn read_header_ranged(
982    src: &mut dyn RangeSource,
983    format: FileFormat,
984    stop: &dyn Fn() -> bool,
985) -> std::result::Result<Header, RangeError> {
986    let first = match format {
987        FileFormat::Gguf => FIRST_GGUF_RANGE,
988        _ => FIRST_SAFETENSORS_RANGE,
989    };
990    read_header_ranged_from(src, format, first, stop)
991}
992
993/// As [`read_header_ranged`] with first range `first` (at least 8 bytes): small for the
994/// fuzz target and tests, to cross many ranges.
995pub fn read_header_ranged_from(
996    src: &mut dyn RangeSource,
997    format: FileFormat,
998    first: u64,
999    stop: &dyn Fn() -> bool,
1000) -> std::result::Result<Header, RangeError> {
1001    if format == FileFormat::Gguf {
1002        let (head, len) = fetch(src, 0, first.max(1), None)?;
1003        let reader = Ranged {
1004            src,
1005            len,
1006            limit: len.min(MAX_GGUF_HEADER),
1007            buf: head,
1008            buf_start: 0,
1009            pos: 0,
1010            next: first.max(1).saturating_mul(2).min(MAX_RANGE),
1011            stop,
1012        };
1013        return Ok(read_gguf(reader, len)?);
1014    }
1015    let (mut head, len) = fetch(src, 0, first.max(8), None)?;
1016    let prefix: [u8; 8] = head
1017        .get(..8)
1018        .and_then(|prefix| prefix.try_into().ok())
1019        .ok_or_else(|| eyre!("the file is shorter than its SafeTensors header length"))?;
1020    let header_len = safetensors_header_len(prefix, len)?;
1021    let end = 8 + header_len;
1022    // The first read holds the whole JSON, or the front of it: the rest is asked for
1023    // once, up to its end and no further.
1024    if (head.len() as u64) < end {
1025        if stop() {
1026            return Err(RangeError::Failed("cancelled".to_string()));
1027        }
1028        let rest = fetch(src, head.len() as u64, end, Some(len))?.0;
1029        head.extend(rest);
1030    }
1031    let json = &head[8..end as usize];
1032    Ok(parse_safetensors_json(json, len - end)?)
1033}
1034
1035/// A remote index: its shard names and metadata, read whole within
1036/// [`MAX_INDEX_JSON`].
1037fn read_index_ranged(
1038    src: &mut dyn RangeSource,
1039    named: &str,
1040) -> std::result::Result<(Vec<String>, Metadata), RangeError> {
1041    let (mut text, len) = fetch(src, 0, FIRST_INDEX_RANGE, None)?;
1042    if len > MAX_INDEX_JSON {
1043        return Err(RangeError::Failed(crate::error_display::file_message(
1044            Path::new(named),
1045            &format!("the index is {len} bytes, more than datui reads"),
1046        )));
1047    }
1048    if len > text.len() as u64 {
1049        text.extend(fetch(src, text.len() as u64, len, Some(len))?.0);
1050    }
1051    Ok(parse_index(&text, named)?)
1052}
1053
1054/// A ranged source for a URL.
1055pub type OpenRanges<'a> =
1056    dyn Fn(&str) -> std::result::Result<Box<dyn RangeSource>, RangeError> + Sync + 'a;
1057
1058/// How a remote model's files are reached: a source per URL and sibling file URLs.
1059/// `open` and `stop` are called from concurrent shard readers ([`SHARD_READS`]).
1060pub struct Remote<'a> {
1061    pub open: &'a OpenRanges<'a>,
1062    pub sibling: &'a dyn Fn(&str, &str) -> String,
1063    pub stop: &'a (dyn Fn() -> bool + Sync),
1064}
1065
1066/// Shard headers read at once: hundreds of shards one at a time add up to minutes; a few
1067/// at once gets most of the gain without a burst.
1068pub const SHARD_READS: usize = 8;
1069
1070/// The last segment of a URL, without a query: what a file's row and its errors call it.
1071pub fn url_file_name(url: &str) -> &str {
1072    let path = url.split(['?', '#']).next().unwrap_or(url);
1073    path.rsplit('/').next().unwrap_or(path)
1074}
1075
1076/// Read `urls` (remote model files, or SafeTensors indexes naming them) as one tensor
1077/// table, fetching only headers, as [`read_model`] does on disk.
1078/// [`RangeError::NoRanges`] only for one standalone file (downloadable instead).
1079pub fn read_remote_model(
1080    urls: &[String],
1081    format: FileFormat,
1082    remote: &Remote,
1083) -> std::result::Result<(LazyFrame, ModelSummary), RangeError> {
1084    let no_ranges = |url: &str| {
1085        RangeError::Failed(crate::error_display::file_message(
1086            Path::new(url),
1087            "the server does not serve byte ranges, which reading a sharded model's headers needs",
1088        ))
1089    };
1090    // Each failure names the file it came from: a shard, or the index naming them.
1091    let named = |url: &str, e: RangeError| match e {
1092        RangeError::Failed(what) => {
1093            RangeError::Failed(crate::error_display::file_message(Path::new(url), &what))
1094        }
1095        e => e,
1096    };
1097    let mut files: Vec<String> = Vec::new();
1098    let mut seen = std::collections::HashSet::new();
1099    let mut metadata: Metadata = Vec::new();
1100    for url in urls {
1101        if format == FileFormat::Safetensors && is_safetensors_index(Path::new(url_file_name(url)))
1102        {
1103            let (names, index_meta) = (remote.open)(url)
1104                .and_then(|mut src| read_index_ranged(src.as_mut(), url))
1105                .map_err(|e| match e {
1106                    RangeError::NoRanges => no_ranges(url),
1107                    e => named(url, e),
1108                })?;
1109            merge_metadata(&mut metadata, index_meta);
1110            for name in names {
1111                let shard = (remote.sibling)(url, &name);
1112                if seen.insert(shard.clone()) {
1113                    files.push(shard);
1114                }
1115            }
1116        } else if seen.insert(url.clone()) {
1117            files.push(url.clone());
1118        }
1119    }
1120    if files.is_empty() {
1121        return Err(RangeError::Failed("no model files to read".to_string()));
1122    }
1123    // Named on its own, a file the server sends whole is downloaded instead.
1124    let alone = files.len() == 1 && urls.len() == 1 && files[0] == urls[0];
1125    let headers = read_headers(&files, format, remote).map_err(|(file, e)| match e {
1126        RangeError::NoRanges if alone => RangeError::NoRanges,
1127        RangeError::NoRanges => no_ranges(file),
1128        e => named(file, e),
1129    })?;
1130    let names: Vec<String> = files.iter().map(|f| url_file_name(f).to_string()).collect();
1131    Ok(build(&headers, &names, metadata)?)
1132}
1133
1134/// Each of `files`' headers in order, [`SHARD_READS`] at a time; the first failure is
1135/// the error, after which (or once stopped) no more requests go out.
1136fn read_headers<'f>(
1137    files: &'f [String],
1138    format: FileFormat,
1139    remote: &Remote,
1140) -> std::result::Result<Vec<Header>, (&'f str, RangeError)> {
1141    use std::sync::Mutex;
1142    use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering};
1143    let next = AtomicUsize::new(0);
1144    let failed = AtomicBool::new(false);
1145    let first_error: Mutex<Option<(usize, RangeError)>> = Mutex::new(None);
1146    let read: Vec<Mutex<Option<Header>>> = files.iter().map(|_| Mutex::new(None)).collect();
1147    let stop = || failed.load(Ordering::Relaxed) || (remote.stop)();
1148    std::thread::scope(|scope| {
1149        for _ in 0..SHARD_READS.min(files.len()) {
1150            scope.spawn(|| {
1151                loop {
1152                    let at = next.fetch_add(1, Ordering::Relaxed);
1153                    if at >= files.len() || stop() {
1154                        return;
1155                    }
1156                    match (remote.open)(&files[at])
1157                        .and_then(|mut src| read_header_ranged(src.as_mut(), format, &stop))
1158                    {
1159                        Ok(header) => {
1160                            *read[at].lock().unwrap_or_else(|e| e.into_inner()) = Some(header);
1161                        }
1162                        Err(e) => {
1163                            // The others stop at their next request: only the first is
1164                            // the reason.
1165                            if !failed.swap(true, Ordering::Relaxed) {
1166                                *first_error.lock().unwrap_or_else(|e| e.into_inner()) =
1167                                    Some((at, e));
1168                            }
1169                            return;
1170                        }
1171                    }
1172                }
1173            });
1174        }
1175    });
1176    if let Some((at, e)) = first_error.into_inner().unwrap_or_else(|e| e.into_inner()) {
1177        return Err((&files[at], e));
1178    }
1179    read.into_iter()
1180        .map(|slot| slot.into_inner().unwrap_or_else(|e| e.into_inner()))
1181        .collect::<Option<Vec<Header>>>()
1182        // Stopped before every file was read.
1183        .ok_or((
1184            files.first().map_or("", String::as_str),
1185            RangeError::Failed("cancelled".to_string()),
1186        ))
1187}
1188
1189/// Read `paths` — model files, or SafeTensors indexes that name them — as one table of
1190/// tensors, with a `file` column when there is more than one file.
1191pub fn read_model(paths: &[PathBuf], format: FileFormat) -> Result<(LazyFrame, ModelSummary)> {
1192    let mut files: Vec<PathBuf> = Vec::new();
1193    // Each file once, however many indexes and names reach it.
1194    let mut seen = std::collections::HashSet::new();
1195    let mut metadata: Vec<(String, MetaValue)> = Vec::new();
1196    for path in paths {
1197        if format == FileFormat::Safetensors && is_safetensors_index(path) {
1198            let (shards, index_meta) = read_index(path)?;
1199            merge_metadata(&mut metadata, index_meta);
1200            for shard in shards {
1201                if seen.insert(shard.clone()) {
1202                    files.push(shard);
1203                }
1204            }
1205        } else if seen.insert(path.clone()) {
1206            files.push(path.clone());
1207        }
1208    }
1209    if files.is_empty() {
1210        return Err(eyre!("No model files to read"));
1211    }
1212    let mut headers = Vec::with_capacity(files.len());
1213    for file in &files {
1214        // The open names the path it was given; of several files, say which one.
1215        let header = read_file(file, format).map_err(|e| match files.len() {
1216            1 => e,
1217            _ => crate::error_display::in_file(file, e),
1218        })?;
1219        headers.push(header);
1220    }
1221    let names: Vec<String> = files
1222        .iter()
1223        .map(|f| {
1224            f.file_name()
1225                .map(|n| n.to_string_lossy().into_owned())
1226                .unwrap_or_else(|| f.display().to_string())
1227        })
1228        .collect();
1229    build(&headers, &names, metadata)
1230}
1231
1232/// Keep each key's first value.
1233fn merge_metadata(into: &mut Vec<(String, MetaValue)>, from: Vec<(String, MetaValue)>) {
1234    let mut seen: std::collections::HashSet<String> = into.iter().map(|(k, _)| k.clone()).collect();
1235    into.extend(from.into_iter().filter(|(key, _)| seen.insert(key.clone())));
1236}
1237
1238/// The table and the summary for headers already read; `names` are their files.
1239pub fn build(
1240    headers: &[Header],
1241    names: &[String],
1242    mut metadata: Vec<(String, MetaValue)>,
1243) -> Result<(LazyFrame, ModelSummary)> {
1244    let kind = headers
1245        .first()
1246        .map(|h| h.kind)
1247        .ok_or_else(|| eyre!("No model files to read"))?;
1248    let safetensors = kind == ModelKind::SafeTensors;
1249    let many = headers.len() > 1;
1250    let rows: usize = headers.iter().map(|h| h.tensors.len()).sum();
1251
1252    let mut file_col = Vec::with_capacity(if many { rows } else { 0 });
1253    let mut name = Vec::with_capacity(rows);
1254    let mut dtype = Vec::with_capacity(rows);
1255    let values: usize = headers
1256        .iter()
1257        .flat_map(|h| &h.tensors)
1258        .map(|t| t.shape.len())
1259        .sum();
1260    let mut shape = ListPrimitiveChunkedBuilder::<UInt64Type>::new(
1261        "shape".into(),
1262        rows,
1263        values,
1264        DataType::UInt64,
1265    );
1266    let mut params = Vec::with_capacity(rows);
1267    let mut bytes = Vec::with_capacity(rows);
1268    let mut start = Vec::with_capacity(rows);
1269    let mut end = Vec::with_capacity(rows);
1270    // By name, then into a list: a hostile header can name a type per tensor.
1271    let mut types: std::collections::HashMap<&str, TypeShare> = Default::default();
1272    let (mut total_params, mut total_bytes) = (0u64, 0u64);
1273    merge_metadata(
1274        &mut metadata,
1275        headers.iter().flat_map(|h| h.metadata.clone()).collect(),
1276    );
1277    for (header, file) in headers.iter().zip(names) {
1278        for t in &header.tensors {
1279            if many {
1280                file_col.push(file.as_str());
1281            }
1282            name.push(t.name.as_str());
1283            dtype.push(t.dtype.as_str());
1284            shape.append_slice(&t.shape);
1285            params.push(t.params);
1286            bytes.push(t.bytes);
1287            start.push(t.offset);
1288            end.push(t.offset_end);
1289            let p = t.params.unwrap_or(0);
1290            total_params = total_params.saturating_add(p);
1291            total_bytes = total_bytes.saturating_add(t.bytes.unwrap_or(0));
1292            let share = types.entry(t.dtype.as_str()).or_insert_with(|| TypeShare {
1293                name: t.dtype.clone(),
1294                tensors: 0,
1295                params: 0,
1296            });
1297            share.tensors += 1;
1298            share.params = share.params.saturating_add(p);
1299        }
1300    }
1301    let mut types: Vec<TypeShare> = types.into_values().collect();
1302    types.sort_by(|a, b| b.params.cmp(&a.params).then_with(|| a.name.cmp(&b.name)));
1303
1304    let shape = shape.finish().into_series();
1305    let mut columns: Vec<Column> = Vec::new();
1306    if many {
1307        columns.push(Series::new("file".into(), file_col).into());
1308    }
1309    columns.push(Series::new("name".into(), name).into());
1310    columns.push(Series::new(if safetensors { "dtype" } else { "type" }.into(), dtype).into());
1311    columns.push(shape.into());
1312    columns.push(Series::new("params".into(), params).into());
1313    columns.push(Series::new("bytes".into(), bytes).into());
1314    if safetensors {
1315        columns.push(Series::new("offset_start".into(), start).into());
1316        columns.push(Series::new("offset_end".into(), end).into());
1317    } else {
1318        columns.push(Series::new("offset".into(), start).into());
1319    }
1320    let df = DataFrame::new(rows, columns)?;
1321    let summary = ModelSummary {
1322        kind,
1323        files: headers.len(),
1324        tensors: rows,
1325        params: total_params,
1326        bytes: total_bytes,
1327        types,
1328        metadata,
1329    };
1330    Ok((df.lazy(), summary))
1331}
1332
1333/// Each type's share of the parameters, most first: `Q4_K 87% · Q6_K 12% · F32 <1%`.
1334/// By tensors when no tensor has a parameter count.
1335fn type_mix(types: &[TypeShare], sep: &str) -> String {
1336    let by_params = types.iter().any(|t| t.params > 0);
1337    let total: u64 = if by_params {
1338        types.iter().map(|t| t.params).fold(0, u64::saturating_add)
1339    } else {
1340        types.iter().map(|t| t.tensors as u64).sum()
1341    };
1342    types
1343        .iter()
1344        .map(|t| {
1345            let part = if by_params {
1346                t.params
1347            } else {
1348                t.tensors as u64
1349            };
1350            let pct = if total == 0 {
1351                0.0
1352            } else {
1353                part as f64 * 100.0 / total as f64
1354            };
1355            if pct > 0.0 && pct < 1.0 {
1356                format!("{} <1%", t.name)
1357            } else {
1358                format!("{} {:.0}%", t.name, pct)
1359            }
1360        })
1361        .collect::<Vec<_>>()
1362        .join(sep)
1363}
1364
1365/// The Model tab: the model's totals, then its metadata as key and value, every value
1366/// whole.
1367pub fn detail(model: &ModelSummary) -> crate::formats::text_formats::Detail {
1368    use crate::widgets::info::{count_of, group_u64, short_count};
1369    let sep = format!(" {} ", crate::glyphs::get().middot);
1370    let mut head = model.kind.label();
1371    head.push_str(&sep);
1372    head.push_str(&count_of(model.tensors as u64, "tensor", "tensors"));
1373    if model.files > 1 {
1374        head.push_str(&sep);
1375        head.push_str(&count_of(model.files as u64, "file", "files"));
1376    }
1377    let mut lines = vec![
1378        head,
1379        format!(
1380            "Parameters: {}{}{sep}Size: {}",
1381            group_u64(model.params),
1382            // The short form only where it is shorter.
1383            if model.params >= 1000 {
1384                format!(" ({})", short_count(model.params))
1385            } else {
1386                String::new()
1387            },
1388            crate::numfmt::bytes(model.bytes)
1389        ),
1390    ];
1391    if !model.types.is_empty() {
1392        lines.push(format!("Types: {}", type_mix(&model.types, &sep)));
1393    }
1394    crate::formats::text_formats::Detail {
1395        tab: crate::formats::text_formats::tab(crate::FileFormat::Safetensors),
1396        lines,
1397        list_title: "Metadata",
1398        list: model.metadata.clone(),
1399        // A model's schema is the same seven columns every time; what is particular
1400        // to it is here.
1401        first: true,
1402        own_columns: true,
1403        ..Default::default()
1404    }
1405}
1406
1407/// What a model's header says besides its tensors, as the dataset takes it.
1408pub(crate) fn opened(summary: &ModelSummary) -> crate::formats::members::Opened {
1409    crate::formats::members::Opened {
1410        detail: Some(std::sync::Arc::new(detail(summary))),
1411        ..Default::default()
1412    }
1413}
1414
1415/// The scan of model files: their tensors, with the header's totals and metadata.
1416fn scan(input: crate::formats::readers::ScanIn<'_>) -> Result<crate::loading::scan::Scan> {
1417    let (lf, summary) = read_model(input.paths, input.format)?;
1418    input.report.opened = Some(std::sync::Arc::new(opened(&summary)));
1419    Ok(lf.into())
1420}
1421
1422#[cfg(test)]
1423pub(crate) mod tests;