use std::collections::HashMap;
use kime_tensor::{Blob, DType};
use crate::error::{Error, Result};
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct Entry {
pub name: String,
pub dtype: DType,
pub shape: Vec<usize>,
pub start: usize,
pub end: usize,
}
impl Entry {
#[must_use]
pub fn numel(&self) -> usize {
self.shape.iter().product()
}
}
#[derive(Debug, Clone, Copy)]
pub struct View<'a> {
pub name: &'a str,
pub dtype: DType,
pub shape: &'a [usize],
pub bytes: &'a [u8],
}
impl View<'_> {
#[must_use]
pub fn to_f32(&self) -> Vec<f32> {
(0..self.bytes.len() / self.dtype.size())
.map(|i| self.dtype.read_f32(self.bytes, i))
.collect()
}
}
#[derive(Debug)]
pub struct Tensors {
blob: Blob,
entries: Vec<Entry>,
by_name: HashMap<String, usize>,
}
impl Tensors {
pub(crate) fn new(blob: Blob, entries: Vec<Entry>) -> Result<Self> {
let mut by_name = HashMap::with_capacity(entries.len());
for (i, e) in entries.iter().enumerate() {
debug_assert!(e.start <= e.end && e.end <= blob.len());
if by_name.insert(e.name.clone(), i).is_some() {
return Err(Error::format(format!("tensor {:?} appears twice", e.name)));
}
}
Ok(Self { blob, entries, by_name })
}
#[must_use]
pub fn entries(&self) -> &[Entry] {
&self.entries
}
#[must_use]
pub fn index(&self, name: &str) -> Option<usize> {
self.by_name.get(name).copied()
}
#[must_use]
pub fn view(&self, i: usize) -> View<'_> {
let e = &self.entries[i];
View { name: &e.name, dtype: e.dtype, shape: &e.shape, bytes: &self.blob[e.start..e.end] }
}
#[must_use]
pub fn get(&self, name: &str) -> Option<View<'_>> {
self.index(name).map(|i| self.view(i))
}
#[must_use]
pub fn blob(&self) -> &Blob {
&self.blob
}
#[must_use]
pub fn data_bytes(&self) -> usize {
self.entries.iter().map(|e| e.end - e.start).sum()
}
}
pub(crate) fn byte_len(dtype: DType, shape: &[usize]) -> Option<usize> {
shape.iter().try_fold(dtype.size(), |acc, &d| acc.checked_mul(d))
}
pub(crate) fn check_disjoint(ranges: &mut [(usize, usize, &str)]) -> Result<()> {
ranges.sort_unstable();
for w in ranges.windows(2) {
if w[1].0 < w[0].1 {
return Err(Error::format(format!("{:?} overlaps {:?}", w[1].2, w[0].2)));
}
}
Ok(())
}