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_notation_retained_bytes, saturating_sum, ParsedEinsum, EINSUM_EXTENSION_FAMILY_ID,
13 EINSUM_PARSE_CACHE,
14};
15use crate::ellipsis::{notation_without_ellipsis, resolve_einsum_notation};
16use crate::extension::EinsumExtensionOp;
17use crate::optimize::{plan_spec_from_optimize, resolve_einsum_strategy_with_spec};
18use crate::{
19 parse_einsum_notation, EinsumAxis, EinsumNotation, EinsumOptimize, EinsumSubscripts, Error,
20 Result, Subscripts, TensorDotAxes,
21};
22
23pub trait TraceContextEinsumExt {
51 fn einsum(&mut self, inputs: &[TraceValue], subscripts: &str) -> Result<TraceValue>;
60
61 fn einsum_notation(
78 &mut self,
79 inputs: &[TraceValue],
80 notation: &EinsumNotation,
81 ) -> Result<TraceValue>;
82
83 fn einsum_subscripts(
91 &mut self,
92 inputs: &[TraceValue],
93 subscripts: &EinsumSubscripts,
94 ) -> Result<TraceValue>;
95
96 fn einsum_with(
105 &mut self,
106 inputs: &[TraceValue],
107 subscripts: &str,
108 optimize: EinsumOptimize,
109 ) -> Result<TraceValue>;
110
111 fn einsum_notation_with(
128 &mut self,
129 inputs: &[TraceValue],
130 notation: &EinsumNotation,
131 optimize: EinsumOptimize,
132 ) -> Result<TraceValue>;
133
134 fn einsum_subscripts_with(
142 &mut self,
143 inputs: &[TraceValue],
144 subscripts: &EinsumSubscripts,
145 optimize: EinsumOptimize,
146 ) -> Result<TraceValue>;
147}
148
149impl TraceContextEinsumExt for TraceContext {
150 fn einsum(&mut self, inputs: &[TraceValue], subscripts: &str) -> Result<TraceValue> {
151 self.einsum_with(inputs, subscripts, EinsumOptimize::default())
152 }
153
154 fn einsum_notation(
155 &mut self,
156 inputs: &[TraceValue],
157 notation: &EinsumNotation,
158 ) -> Result<TraceValue> {
159 self.einsum_notation_with(inputs, notation, EinsumOptimize::default())
160 }
161
162 fn einsum_subscripts(
163 &mut self,
164 inputs: &[TraceValue],
165 subscripts: &EinsumSubscripts,
166 ) -> Result<TraceValue> {
167 self.einsum_subscripts_with(inputs, subscripts, EinsumOptimize::default())
168 }
169
170 fn einsum_with(
171 &mut self,
172 inputs: &[TraceValue],
173 subscripts: &str,
174 optimize: EinsumOptimize,
175 ) -> Result<TraceValue> {
176 let parsed = cached_subscripts(self.extension_caches_mut(), subscripts)?;
177 self.einsum_notation_with(inputs, &parsed.notation, optimize)
178 }
179
180 fn einsum_notation_with(
181 &mut self,
182 inputs: &[TraceValue],
183 notation: &EinsumNotation,
184 optimize: EinsumOptimize,
185 ) -> Result<TraceValue> {
186 trace_context_einsum_notation_with(self, inputs, notation, optimize)
187 }
188
189 fn einsum_subscripts_with(
190 &mut self,
191 inputs: &[TraceValue],
192 subscripts: &EinsumSubscripts,
193 optimize: EinsumOptimize,
194 ) -> Result<TraceValue> {
195 trace_context_einsum_subscripts_with(self, inputs, subscripts, optimize, false)
196 }
197}
198
199fn trace_context_einsum_notation_with(
200 trace: &mut TraceContext,
201 inputs: &[TraceValue],
202 notation: &EinsumNotation,
203 optimize: EinsumOptimize,
204) -> Result<TraceValue> {
205 let allow_broadcast = notation
206 .inputs
207 .iter()
208 .chain(std::iter::once(¬ation.output))
209 .any(|term| term.contains(&EinsumAxis::Ellipsis));
210 let subscripts = if allow_broadcast {
211 let shapes = trace_concrete_shapes(trace, inputs)?.ok_or_else(|| {
212 Error::planning("ellipsis einsum requires concrete input dimensions during tracing")
213 })?;
214 let shape_refs: Vec<_> = shapes.iter().map(Vec::as_slice).collect();
215 resolve_einsum_notation(notation, &shape_refs)?.into()
216 } else {
217 notation_without_ellipsis(notation)?.into()
218 };
219 trace_context_einsum_subscripts_with(trace, inputs, &subscripts, optimize, allow_broadcast)
220}
221
222fn trace_context_einsum_subscripts_with(
223 trace: &mut TraceContext,
224 inputs: &[TraceValue],
225 subscripts: &EinsumSubscripts,
226 optimize: EinsumOptimize,
227 allow_broadcast: bool,
228) -> Result<TraceValue> {
229 if inputs.is_empty() {
230 return Err(Error::invalid_argument(
231 "einsum",
232 "inputs",
233 "einsum requires at least one input tensor",
234 ));
235 }
236 if subscripts.inputs.len() != inputs.len() {
237 return Err(Error::invalid_argument(
238 "einsum",
239 "inputs",
240 format!(
241 "einsum subscripts expect {} inputs, got {}",
242 subscripts.inputs.len(),
243 inputs.len()
244 ),
245 ));
246 }
247
248 let subs = Subscripts::from(subscripts);
249 let op = match optimize {
250 EinsumOptimize::Tree(tree) => {
251 let shapes = trace_concrete_shapes(trace, inputs)?.ok_or_else(|| {
252 Error::planning("precomputed contraction tree requires concrete input shapes")
253 })?;
254 let shape_refs: Vec<_> = shapes.iter().map(Vec::as_slice).collect();
255 let (plan_spec, _tree) =
256 resolve_einsum_strategy_with_spec(EinsumOptimize::Tree(tree), &subs, &shape_refs)?;
257 EinsumExtensionOp::with_plan_spec_and_broadcast(
258 subscripts.clone(),
259 plan_spec,
260 allow_broadcast,
261 )
262 }
263 optimize => EinsumExtensionOp::with_plan_spec_and_broadcast(
264 subscripts.clone(),
265 plan_spec_from_optimize(optimize, &subs)?,
266 allow_broadcast,
267 ),
268 };
269 let outputs = trace
270 .add_extension(Arc::new(op), inputs)
271 .map_err(semantic_trace_error)?;
272 outputs.first().copied().ok_or_else(|| {
273 Error::Runtime(tenferro_runtime::Error::Internal(
274 "einsum semantic extension produced no output".into(),
275 ))
276 })
277}
278
279fn trace_concrete_shapes(
280 trace: &TraceContext,
281 inputs: &[TraceValue],
282) -> Result<Option<Vec<Vec<usize>>>> {
283 let mut shapes = Vec::with_capacity(inputs.len());
284 for &value in inputs {
285 let metadata = trace.value_metadata(value).map_err(semantic_trace_error)?;
286 let Some(shape) = metadata
287 .shape()
288 .iter()
289 .map(|extent| match extent.as_exact() {
290 Some(DimExpr::Const(value)) => Some(*value),
291 _ => None,
292 })
293 .collect::<Option<Vec<_>>>()
294 else {
295 return Ok(None);
296 };
297 shapes.push(shape);
298 }
299 Ok(Some(shapes))
300}
301
302fn semantic_trace_error(source: ProgramBuildError) -> Error {
303 Error::Runtime(tenferro_runtime::Error::extension(
304 "einsum",
305 tenferro_runtime::ErrorPhase::GraphBuild,
306 EINSUM_EXTENSION_FAMILY_ID,
307 ErrorKind::Validation(ValidationKind::InvalidArgument),
308 source,
309 ))
310}
311
312pub trait TracedTensorEinsumExt {
314 fn tensordot(&self, rhs: &TracedTensor, axes: TensorDotAxes<'_>) -> Result<TracedTensor>;
327}
328
329impl TracedTensorEinsumExt for TracedTensor {
330 fn tensordot(&self, rhs: &TracedTensor, axes: TensorDotAxes<'_>) -> Result<TracedTensor> {
331 tensordot(self, rhs, axes)
332 }
333}
334
335fn tensordot(
336 lhs: &TracedTensor,
337 rhs: &TracedTensor,
338 axes: TensorDotAxes<'_>,
339) -> Result<TracedTensor> {
340 let config = crate::tensordot::dot_general_config(axes, lhs.rank, rhs.rank)?;
341 crate::tensordot::validate_traced_contract_dims(lhs, rhs, &config)?;
342 lhs.dot_general(rhs, config).map_err(Error::Runtime)
343}
344
345struct ParsedEinsumCacheEntry {
346 notation: String,
347 parsed: Arc<ParsedEinsum>,
348}
349
350impl ParsedEinsumCacheEntry {
351 fn matches_notation(&self, notation: &str) -> bool {
352 self.notation == notation
353 }
354}
355
356fn cached_subscripts(
357 caches: &mut ExtensionCacheStore,
358 notation: &str,
359) -> Result<Arc<ParsedEinsum>> {
360 let key = ExtensionCacheKey::new(
361 EINSUM_EXTENSION_FAMILY_ID,
362 EINSUM_PARSE_CACHE,
363 hash_value(notation),
364 );
365 if let Some(cached) = caches.get::<ParsedEinsumCacheEntry>(&key) {
366 if cached.matches_notation(notation) {
367 return Ok(Arc::clone(&cached.parsed));
368 }
369 }
370
371 let parsed = Arc::new(ParsedEinsum {
372 notation: parse_einsum_notation(notation)?,
373 });
374 let entry = ParsedEinsumCacheEntry {
375 notation: notation.to_owned(),
376 parsed: Arc::clone(&parsed),
377 };
378 let retained_bytes = saturating_sum([
379 entry.notation.len(),
380 einsum_notation_retained_bytes(&parsed.notation),
381 ]);
382 caches.put(key, entry, retained_bytes);
383 Ok(parsed)
384}
385
386fn hash_value<T: Hash + ?Sized>(value: &T) -> u64 {
387 let mut hasher = DefaultHasher::new();
388 value.hash(&mut hasher);
389 hasher.finish()
390}
391
392#[cfg(test)]
393mod tests;