Skip to main content

tenferro_einsum/
cache.rs

1use std::mem::size_of;
2
3use crate::{EinsumAxis, EinsumNotation, EinsumSubscripts};
4
5/// Stable family identifier for the standard tenferro einsum extension.
6pub const EINSUM_EXTENSION_FAMILY_ID: &str = "tenferro.einsum.v1";
7
8/// Compiler-side subscript parse cache name.
9pub(crate) const EINSUM_PARSE_CACHE: &str = "parse";
10/// Executor-side runtime contraction-plan cache name.
11pub(crate) const EINSUM_RUNTIME_PLANS_CACHE: &str = "runtime_plans";
12/// EagerTensor expanded standard-op program cache name.
13#[cfg(feature = "autodiff")]
14pub(crate) const EINSUM_EAGER_EXPANDED_PROGRAMS_CACHE: &str = "eager_expanded_programs";
15
16/// Parsed einsum notation retained by parse caches.
17pub(crate) struct ParsedEinsum {
18    /// Parsed rank-unresolved notation.
19    pub(crate) notation: EinsumNotation,
20}
21
22/// Return the retained-byte estimate for canonical subscripts.
23#[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 einsum_notation_retained_bytes(notation: &EinsumNotation) -> usize {
32    saturating_sum([
33        vec_of_vec_retained_bytes(&notation.inputs),
34        vec_retained_bytes(&notation.output),
35        (notation.inputs.iter().map(Vec::capacity).sum::<usize>() + notation.output.capacity())
36            .saturating_mul(std::mem::size_of::<EinsumAxis>()),
37    ])
38}
39
40pub(crate) fn vec_retained_bytes<T>(values: &Vec<T>) -> usize {
41    values.capacity().saturating_mul(size_of::<T>())
42}
43
44pub(crate) fn vec_of_vec_retained_bytes<T>(values: &[Vec<T>]) -> usize {
45    saturating_sum(values.iter().map(vec_retained_bytes))
46}
47
48pub(crate) fn saturating_sum(values: impl IntoIterator<Item = usize>) -> usize {
49    values.into_iter().fold(0usize, usize::saturating_add)
50}
51
52#[cfg(test)]
53mod tests {
54    use super::saturating_sum;
55
56    #[test]
57    fn retained_byte_sums_saturate() {
58        assert_eq!(saturating_sum([usize::MAX, 1]), usize::MAX);
59        assert_eq!(saturating_sum([usize::MAX - 4, 2, 8]), usize::MAX);
60    }
61}