1use std::mem::size_of;
2
3use crate::EinsumSubscripts;
4
5pub const EINSUM_EXTENSION_FAMILY_ID: &str = "tenferro.einsum.v1";
7
8pub(crate) const EINSUM_PARSE_CACHE: &str = "parse";
10pub(crate) const EINSUM_RUNTIME_PLANS_CACHE: &str = "runtime_plans";
12#[cfg(feature = "autodiff")]
14pub(crate) const EINSUM_EAGER_EXPANDED_PROGRAMS_CACHE: &str = "eager_expanded_programs";
15
16pub(crate) struct ParsedEinsum {
18 pub(crate) subscripts: EinsumSubscripts,
20}
21
22#[must_use]
24pub(crate) fn einsum_subscripts_retained_bytes(subscripts: &EinsumSubscripts) -> usize {
25 saturating_sum([
26 vec_of_vec_retained_bytes(&subscripts.inputs),
27 vec_retained_bytes(&subscripts.output),
28 ])
29}
30
31pub(crate) fn vec_retained_bytes<T>(values: &Vec<T>) -> usize {
32 values.capacity().saturating_mul(size_of::<T>())
33}
34
35pub(crate) fn vec_of_vec_retained_bytes<T>(values: &[Vec<T>]) -> usize {
36 saturating_sum(values.iter().map(vec_retained_bytes))
37}
38
39pub(crate) fn saturating_sum(values: impl IntoIterator<Item = usize>) -> usize {
40 values.into_iter().fold(0usize, usize::saturating_add)
41}
42
43#[cfg(test)]
44mod tests {
45 use super::saturating_sum;
46
47 #[test]
48 fn retained_byte_sums_saturate() {
49 assert_eq!(saturating_sum([usize::MAX, 1]), usize::MAX);
50 assert_eq!(saturating_sum([usize::MAX - 4, 2, 8]), usize::MAX);
51 }
52}