Skip to main content

rutensor/
einsum.rs

1use std::collections::BTreeMap;
2use crate::{ComputeType, Error, Mode, OperandDescriptor, OperationDescriptor, Plan, Result, TensorDescriptor};
3use crate::plan::merge_extent;
4use ruda_core::tensor::DType;
5use ruda_kernel::{dsl::Runtime, tensor::RudaTensor};
6
7#[derive(Clone, Debug)]
8enum Token { Label(Mode), Ellipsis }
9
10fn tokens(text: &str) -> Result<Vec<Token>> {
11    let mut result = Vec::new();
12    let bytes = text.as_bytes();
13    let mut index = 0;
14    let mut ellipsis = false;
15    while index < bytes.len() {
16        if bytes[index].is_ascii_alphabetic() {
17            result.push(Token::Label(bytes[index] as Mode));
18            index += 1;
19        } else if bytes[index..].starts_with(b"...") && !ellipsis {
20            result.push(Token::Ellipsis);
21            ellipsis = true;
22            index += 3;
23        } else {
24            return Err(Error::InvalidExpression("use ASCII letters and at most one ellipsis per operand".into()));
25        }
26    }
27    Ok(result)
28}
29
30fn expand(tokens: &[Token], ellipsis_modes: &[Mode]) -> Vec<Mode> {
31    let mut modes = Vec::new();
32    for token in tokens {
33        match token {
34            Token::Label(mode) => modes.push(*mode),
35            Token::Ellipsis => modes.extend_from_slice(ellipsis_modes),
36        }
37    }
38    modes
39}
40
41pub(crate) fn output_dtype(inputs: &[OperandDescriptor]) -> DType {
42    let first = inputs[0].tensor().dtype();
43    if inputs.iter().all(|input| input.tensor().dtype() == first) { return first; }
44    if inputs.iter().any(|input| input.tensor().dtype() == DType::F64) { DType::F64 } else { DType::F32 }
45}
46
47pub(crate) fn infer_output(inputs: &[OperandDescriptor], modes: &[Mode], dtype: DType) -> Result<TensorDescriptor> {
48    let mut extents = BTreeMap::new();
49    for input in inputs {
50        for (&mode, &extent) in input.modes().iter().zip(input.tensor().extents()) {
51            merge_extent(&mut extents, mode, extent)?;
52        }
53    }
54    let shape = modes.iter().map(|mode| extents.get(mode).copied().ok_or_else(||
55        Error::InvalidOperation(format!("output mode {mode} is absent from all inputs"))))
56        .collect::<Result<Vec<_>>>()?;
57    TensorDescriptor::contiguous(&shape, dtype)
58}
59
60/// A reusable einsum expression with fixed input shapes, strides and storage types.
61#[derive(Clone, Debug)]
62pub struct EinsumPlan {
63    expression: String,
64    plan: Plan,
65}
66
67impl EinsumPlan {
68    pub fn new(expression: &str, inputs: &[TensorDescriptor]) -> Result<Self> {
69        Self::build(expression, inputs, None)
70    }
71
72    pub fn with_options(
73        expression: &str, inputs: &[TensorDescriptor], output_dtype: DType, compute: ComputeType,
74    ) -> Result<Self> {
75        Self::build(expression, inputs, Some((output_dtype, compute)))
76    }
77
78    fn build(expression: &str, inputs: &[TensorDescriptor], options: Option<(DType, ComputeType)>) -> Result<Self> {
79        if inputs.is_empty() {
80            return Err(Error::InvalidExpression("at least one tensor is required".into()));
81        }
82        let expression: String = expression.chars().filter(|c| !c.is_ascii_whitespace()).collect();
83        let mut sides = expression.split("->");
84        let lhs = sides.next().unwrap_or_default();
85        let rhs = sides.next();
86        if sides.next().is_some() {
87            return Err(Error::InvalidExpression("more than one output arrow".into()));
88        }
89        let input_text: Vec<_> = lhs.split(',').collect();
90        if input_text.len() != inputs.len() {
91            return Err(Error::InputCount { expected: input_text.len(), actual: inputs.len() });
92        }
93        let token_lists: Vec<_> = input_text.iter().map(|text| tokens(text)).collect::<Result<_>>()?;
94        let mut widths = Vec::new();
95        let mut counts = BTreeMap::<Mode, usize>::new();
96        for (tokens, descriptor) in token_lists.iter().zip(inputs) {
97            let explicit = tokens.iter().filter(|token| matches!(token, Token::Label(_))).count();
98            let has_ellipsis = tokens.iter().any(|token| matches!(token, Token::Ellipsis));
99            if explicit > descriptor.rank() || (!has_ellipsis && explicit != descriptor.rank()) {
100                return Err(Error::InvalidExpression("operand labels do not match its tensor rank".into()));
101            }
102            widths.push(descriptor.rank() - explicit);
103            for token in tokens {
104                if let Token::Label(mode) = token { *counts.entry(*mode).or_default() += 1; }
105            }
106        }
107        let width = widths.iter().copied().max().unwrap_or(0);
108        let ellipsis_modes = (0..width).map(|axis| {
109            let axis = i32::try_from(axis).map_err(|_| Error::Overflow)?;
110            i32::MIN.checked_add(axis).ok_or(Error::Overflow)
111        }).collect::<Result<Vec<_>>>()?;
112        let mut operands = Vec::new();
113        for ((tokens, descriptor), local_width) in token_lists.iter().zip(inputs).zip(widths) {
114            let modes = expand(tokens, &ellipsis_modes[width - local_width..]);
115            operands.push(OperandDescriptor::new(descriptor.clone(), &modes)?);
116        }
117        let output_modes = if let Some(rhs) = rhs {
118            expand(&tokens(rhs)?, &ellipsis_modes)
119        } else {
120            let mut modes = ellipsis_modes;
121            modes.extend(counts.iter().filter_map(|(&mode, &count)| (count == 1).then_some(mode)));
122            modes
123        };
124        let (dtype, compute) = options.unwrap_or_else(||
125            (output_dtype(&operands), ComputeType::for_operands(&operands)));
126        let output = infer_output(&operands, &output_modes, dtype)?;
127        let operation = OperationDescriptor::sum_product(operands, None, output, &output_modes, compute)?;
128        Ok(Self { expression, plan: Plan::new(operation)? })
129    }
130
131    pub fn expression(&self) -> &str { &self.expression }
132    pub fn plan(&self) -> &Plan { &self.plan }
133    pub fn output(&self) -> &TensorDescriptor { self.plan.operation().output() }
134
135    pub fn execute<R: Runtime>(&self, inputs: &[&RudaTensor<R>]) -> Result<RudaTensor<R>> {
136        self.plan.execute(inputs, &[1.0])
137    }
138
139    pub fn execute_into<R: Runtime>(&self, inputs: &[&RudaTensor<R>], output: RudaTensor<R>) -> Result<RudaTensor<R>> {
140        self.plan.execute_into(inputs, output, &[1.0])
141    }
142}
143
144/// Evaluate an einsum expression on device tensors without reading values back to the host.
145///
146/// Ellipses are right-aligned for broadcasting. Repeated input labels select diagonals.
147/// Explicit output labels select and order free axes; all other axes are summed.
148pub fn einsum<R: Runtime>(expression: &str, inputs: &[&RudaTensor<R>]) -> Result<RudaTensor<R>> {
149    let descriptors = inputs.iter().map(|input| TensorDescriptor::from_tensor(*input)).collect::<Result<Vec<_>>>()?;
150    EinsumPlan::new(expression, &descriptors)?.execute(inputs)
151}
152
153pub fn einsum_with_options<R: Runtime>(
154    expression: &str, inputs: &[&RudaTensor<R>], output_dtype: DType, compute: ComputeType,
155) -> Result<RudaTensor<R>> {
156    let descriptors = inputs.iter().map(|input| TensorDescriptor::from_tensor(*input)).collect::<Result<Vec<_>>>()?;
157    EinsumPlan::with_options(expression, &descriptors, output_dtype, compute)?.execute(inputs)
158}