1use std::collections::HashMap;
4
5use kime_tensor::{Blob, DType};
6
7use crate::error::{Error, Result};
8
9#[derive(Debug, Clone, PartialEq, Eq)]
11pub struct Entry {
12 pub name: String,
14 pub dtype: DType,
16 pub shape: Vec<usize>,
18 pub start: usize,
20 pub end: usize,
22}
23
24impl Entry {
25 #[must_use]
27 pub fn numel(&self) -> usize {
28 self.shape.iter().product()
29 }
30}
31
32#[derive(Debug, Clone, Copy)]
34pub struct View<'a> {
35 pub name: &'a str,
37 pub dtype: DType,
39 pub shape: &'a [usize],
41 pub bytes: &'a [u8],
43}
44
45impl View<'_> {
46 #[must_use]
48 pub fn to_f32(&self) -> Vec<f32> {
49 (0..self.bytes.len() / self.dtype.size())
50 .map(|i| self.dtype.read_f32(self.bytes, i))
51 .collect()
52 }
53}
54
55#[derive(Debug)]
57pub struct Tensors {
58 blob: Blob,
59 entries: Vec<Entry>,
60 by_name: HashMap<String, usize>,
61}
62
63impl Tensors {
64 pub(crate) fn new(blob: Blob, entries: Vec<Entry>) -> Result<Self> {
66 let mut by_name = HashMap::with_capacity(entries.len());
67 for (i, e) in entries.iter().enumerate() {
68 debug_assert!(e.start <= e.end && e.end <= blob.len());
69 if by_name.insert(e.name.clone(), i).is_some() {
70 return Err(Error::format(format!("tensor {:?} appears twice", e.name)));
71 }
72 }
73 Ok(Self { blob, entries, by_name })
74 }
75
76 #[must_use]
78 pub fn entries(&self) -> &[Entry] {
79 &self.entries
80 }
81
82 #[must_use]
84 pub fn index(&self, name: &str) -> Option<usize> {
85 self.by_name.get(name).copied()
86 }
87
88 #[must_use]
90 pub fn view(&self, i: usize) -> View<'_> {
91 let e = &self.entries[i];
92 View { name: &e.name, dtype: e.dtype, shape: &e.shape, bytes: &self.blob[e.start..e.end] }
93 }
94
95 #[must_use]
97 pub fn get(&self, name: &str) -> Option<View<'_>> {
98 self.index(name).map(|i| self.view(i))
99 }
100
101 #[must_use]
103 pub fn blob(&self) -> &Blob {
104 &self.blob
105 }
106
107 #[must_use]
109 pub fn data_bytes(&self) -> usize {
110 self.entries.iter().map(|e| e.end - e.start).sum()
111 }
112}
113
114pub(crate) fn byte_len(dtype: DType, shape: &[usize]) -> Option<usize> {
116 shape.iter().try_fold(dtype.size(), |acc, &d| acc.checked_mul(d))
117}
118
119pub(crate) fn check_disjoint(ranges: &mut [(usize, usize, &str)]) -> Result<()> {
121 ranges.sort_unstable();
122 for w in ranges.windows(2) {
123 if w[1].0 < w[0].1 {
124 return Err(Error::format(format!("{:?} overlaps {:?}", w[1].2, w[0].2)));
125 }
126 }
127 Ok(())
128}