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#[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
144pub 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}