1use std::collections::hash_map::DefaultHasher;
2use std::hash::{Hash, Hasher};
3use std::sync::Arc;
4
5use tenferro_ops::dim_expr::DimExpr;
6use tenferro_runtime::extension::{ExtensionCacheKey, ExtensionCacheStore};
7use tenferro_runtime::program::ProgramBuildError;
8use tenferro_runtime::{TraceContext, TraceValue, TracedTensor};
9use tenferro_tensor::{ErrorKind, ValidationKind};
10
11use crate::cache::{
12 einsum_subscripts_retained_bytes, saturating_sum, ParsedEinsum, EINSUM_EXTENSION_FAMILY_ID,
13 EINSUM_PARSE_CACHE,
14};
15use crate::extension::EinsumExtensionOp;
16use crate::optimize::{plan_spec_from_optimize, resolve_einsum_strategy_with_spec};
17use crate::{
18 parse_einsum_subscripts, EinsumOptimize, EinsumSubscripts, Error, Result, Subscripts,
19 TensorDotAxes,
20};
21
22pub trait TraceContextEinsumExt {
50 fn einsum(&mut self, inputs: &[TraceValue], subscripts: &str) -> Result<TraceValue>;
59
60 fn einsum_subscripts(
68 &mut self,
69 inputs: &[TraceValue],
70 subscripts: &EinsumSubscripts,
71 ) -> Result<TraceValue>;
72
73 fn einsum_with(
82 &mut self,
83 inputs: &[TraceValue],
84 subscripts: &str,
85 optimize: EinsumOptimize,
86 ) -> Result<TraceValue>;
87
88 fn einsum_subscripts_with(
96 &mut self,
97 inputs: &[TraceValue],
98 subscripts: &EinsumSubscripts,
99 optimize: EinsumOptimize,
100 ) -> Result<TraceValue>;
101}
102
103impl TraceContextEinsumExt for TraceContext {
104 fn einsum(&mut self, inputs: &[TraceValue], subscripts: &str) -> Result<TraceValue> {
105 self.einsum_with(inputs, subscripts, EinsumOptimize::default())
106 }
107
108 fn einsum_subscripts(
109 &mut self,
110 inputs: &[TraceValue],
111 subscripts: &EinsumSubscripts,
112 ) -> Result<TraceValue> {
113 self.einsum_subscripts_with(inputs, subscripts, EinsumOptimize::default())
114 }
115
116 fn einsum_with(
117 &mut self,
118 inputs: &[TraceValue],
119 subscripts: &str,
120 optimize: EinsumOptimize,
121 ) -> Result<TraceValue> {
122 let parsed = cached_subscripts(self.extension_caches_mut(), subscripts)?;
123 self.einsum_subscripts_with(inputs, &parsed.subscripts, optimize)
124 }
125
126 fn einsum_subscripts_with(
127 &mut self,
128 inputs: &[TraceValue],
129 subscripts: &EinsumSubscripts,
130 optimize: EinsumOptimize,
131 ) -> Result<TraceValue> {
132 trace_context_einsum_subscripts_with(self, inputs, subscripts, optimize)
133 }
134}
135
136fn trace_context_einsum_subscripts_with(
137 trace: &mut TraceContext,
138 inputs: &[TraceValue],
139 subscripts: &EinsumSubscripts,
140 optimize: EinsumOptimize,
141) -> Result<TraceValue> {
142 if inputs.is_empty() {
143 return Err(Error::invalid_argument(
144 "einsum",
145 "inputs",
146 "einsum requires at least one input tensor",
147 ));
148 }
149 if subscripts.inputs.len() != inputs.len() {
150 return Err(Error::invalid_argument(
151 "einsum",
152 "inputs",
153 format!(
154 "einsum subscripts expect {} inputs, got {}",
155 subscripts.inputs.len(),
156 inputs.len()
157 ),
158 ));
159 }
160
161 let subs = Subscripts::from(subscripts);
162 let op = match optimize {
163 EinsumOptimize::Tree(tree) => {
164 let shapes = trace_concrete_shapes(trace, inputs)?.ok_or_else(|| {
165 Error::planning("precomputed contraction tree requires concrete input shapes")
166 })?;
167 let shape_refs: Vec<_> = shapes.iter().map(Vec::as_slice).collect();
168 let (plan_spec, _tree) =
169 resolve_einsum_strategy_with_spec(EinsumOptimize::Tree(tree), &subs, &shape_refs)?;
170 EinsumExtensionOp::with_plan_spec(subscripts.clone(), plan_spec)
171 }
172 optimize => EinsumExtensionOp::with_plan_spec(
173 subscripts.clone(),
174 plan_spec_from_optimize(optimize, &subs)?,
175 ),
176 };
177 let outputs = trace
178 .add_extension(Arc::new(op), inputs)
179 .map_err(semantic_trace_error)?;
180 outputs.first().copied().ok_or_else(|| {
181 Error::Runtime(tenferro_runtime::Error::Internal(
182 "einsum semantic extension produced no output".into(),
183 ))
184 })
185}
186
187fn trace_concrete_shapes(
188 trace: &TraceContext,
189 inputs: &[TraceValue],
190) -> Result<Option<Vec<Vec<usize>>>> {
191 let mut shapes = Vec::with_capacity(inputs.len());
192 for &value in inputs {
193 let metadata = trace.value_metadata(value).map_err(semantic_trace_error)?;
194 let Some(shape) = metadata
195 .shape()
196 .iter()
197 .map(|extent| match extent.as_exact() {
198 Some(DimExpr::Const(value)) => Some(*value),
199 _ => None,
200 })
201 .collect::<Option<Vec<_>>>()
202 else {
203 return Ok(None);
204 };
205 shapes.push(shape);
206 }
207 Ok(Some(shapes))
208}
209
210fn semantic_trace_error(source: ProgramBuildError) -> Error {
211 Error::Runtime(tenferro_runtime::Error::extension(
212 "einsum",
213 tenferro_runtime::ErrorPhase::GraphBuild,
214 EINSUM_EXTENSION_FAMILY_ID,
215 ErrorKind::Validation(ValidationKind::InvalidArgument),
216 source,
217 ))
218}
219
220pub trait TracedTensorEinsumExt {
222 fn tensordot(&self, rhs: &TracedTensor, axes: TensorDotAxes<'_>) -> Result<TracedTensor>;
235}
236
237impl TracedTensorEinsumExt for TracedTensor {
238 fn tensordot(&self, rhs: &TracedTensor, axes: TensorDotAxes<'_>) -> Result<TracedTensor> {
239 tensordot(self, rhs, axes)
240 }
241}
242
243fn tensordot(
244 lhs: &TracedTensor,
245 rhs: &TracedTensor,
246 axes: TensorDotAxes<'_>,
247) -> Result<TracedTensor> {
248 let config = crate::tensordot::dot_general_config(axes, lhs.rank, rhs.rank)?;
249 crate::tensordot::validate_traced_contract_dims(lhs, rhs, &config)?;
250 lhs.dot_general(rhs, config).map_err(Error::Runtime)
251}
252
253struct ParsedEinsumCacheEntry {
254 notation: String,
255 parsed: Arc<ParsedEinsum>,
256}
257
258impl ParsedEinsumCacheEntry {
259 fn matches_notation(&self, notation: &str) -> bool {
260 self.notation == notation
261 }
262}
263
264fn cached_subscripts(
265 caches: &mut ExtensionCacheStore,
266 notation: &str,
267) -> Result<Arc<ParsedEinsum>> {
268 let key = ExtensionCacheKey::new(
269 EINSUM_EXTENSION_FAMILY_ID,
270 EINSUM_PARSE_CACHE,
271 hash_value(notation),
272 );
273 if let Some(cached) = caches.get::<ParsedEinsumCacheEntry>(&key) {
274 if cached.matches_notation(notation) {
275 return Ok(Arc::clone(&cached.parsed));
276 }
277 }
278
279 let parsed = Arc::new(ParsedEinsum {
280 subscripts: parse_einsum_subscripts(notation)?,
281 });
282 let entry = ParsedEinsumCacheEntry {
283 notation: notation.to_owned(),
284 parsed: Arc::clone(&parsed),
285 };
286 let retained_bytes = saturating_sum([
287 entry.notation.len(),
288 einsum_subscripts_retained_bytes(&parsed.subscripts),
289 ]);
290 caches.put(key, entry, retained_bytes);
291 Ok(parsed)
292}
293
294fn hash_value<T: Hash + ?Sized>(value: &T) -> u64 {
295 let mut hasher = DefaultHasher::new();
296 value.hash(&mut hasher);
297 hasher.finish()
298}
299
300#[cfg(test)]
301mod tests;