1use std::any::Any;
15use std::hash::Hasher;
16use std::marker::PhantomData;
17use std::sync::Arc;
18
19use tenferro_ops::ext_op::{ExtensionAliasDeclaration, ExtensionEffectDeclaration};
20use tenferro_runtime::extension::{ExtensionOp, ExtensionShapeContext, SymDim};
21use tenferro_runtime::{
22 CoreCapabilityKind, EngineId, ErasedExecutionContext, ErrorPhase, ExecutionContextIdentity,
23 ExtensionCacheStore, ExtensionEngine, ExtensionModule, ExtensionModuleId,
24 ExtensionModuleRegistrar, ExtensionPlanningConfig, ExtensionPrepareRequest, PrepareCapability,
25 PrepareError, PreparedOperation, PreparedOperationBinding, PreparedOperationExecutor,
26 PreparedOperationPlan, ProviderContractError, RuntimeConfigError, SpecializationProjection,
27};
28use tenferro_tensor::{DType, DynRank, Host, Tensor, TensorBackend, TensorRead, TypedTensor};
29
30use crate::Bf16;
31
32pub const BF16_EINSUM_FAMILY: &str = "tenferro-bf16-proof.einsum.v1";
34
35pub const BF16_SCALAR_IDENTITY: &str = "tenferro-bf16-proof.bf16.v1";
37
38#[derive(Clone, Debug, PartialEq, Eq)]
53pub struct Bf16Einsum {
54 lhs: Vec<u32>,
55 rhs: Vec<u32>,
56 out: Vec<u32>,
57}
58
59impl Bf16Einsum {
60 pub fn new(lhs: &[u32], rhs: &[u32], out: &[u32]) -> tenferro_runtime::Result<Self> {
77 let invalid = |message: &str| {
78 tenferro_runtime::Error::from(tenferro_tensor::Error::invalid_argument(
79 "bf16_einsum",
80 "pattern",
81 message,
82 ))
83 };
84 if lhs.is_empty() || rhs.is_empty() {
85 return Err(invalid("an operand must carry at least one label"));
86 }
87 for labels in [lhs, rhs] {
88 let mut seen = labels.to_vec();
89 seen.sort_unstable();
90 seen.dedup();
91 if seen.len() != labels.len() {
92 return Err(invalid(
93 "a label repeats within one operand, which is a trace and not supported here",
94 ));
95 }
96 }
97 for label in out {
98 if !lhs.contains(label) && !rhs.contains(label) {
99 return Err(invalid(
100 "an output label must appear in at least one operand",
101 ));
102 }
103 }
104 Ok(Self {
105 lhs: lhs.to_vec(),
106 rhs: rhs.to_vec(),
107 out: out.to_vec(),
108 })
109 }
110
111 #[must_use]
122 pub fn labels(&self) -> (&[u32], &[u32], &[u32]) {
123 (&self.lhs, &self.rhs, &self.out)
124 }
125}
126
127impl ExtensionOp for Bf16Einsum {
128 fn family_id(&self) -> &'static str {
129 BF16_EINSUM_FAMILY
130 }
131
132 fn payload_hash(&self, hasher: &mut dyn Hasher) {
133 for labels in [&self.lhs, &self.rhs, &self.out] {
134 hasher.write_usize(labels.len());
135 for label in labels {
136 hasher.write_u32(*label);
137 }
138 }
139 }
140
141 fn payload_eq(&self, other: &dyn ExtensionOp) -> bool {
142 other
143 .as_any()
144 .downcast_ref::<Self>()
145 .is_some_and(|other| other == self)
146 }
147
148 fn clone_arc(&self) -> Arc<dyn ExtensionOp> {
149 Arc::new(self.clone())
150 }
151
152 fn as_any(&self) -> &dyn Any {
153 self
154 }
155
156 fn input_count(&self) -> usize {
157 2
158 }
159
160 fn output_count(&self) -> usize {
161 1
162 }
163
164 fn semantic_effects(&self) -> ExtensionEffectDeclaration<'_> {
165 ExtensionEffectDeclaration::Declared(&[])
166 }
167
168 fn semantic_aliases(&self) -> ExtensionAliasDeclaration<'_> {
169 ExtensionAliasDeclaration::AllFresh
170 }
171
172 fn scalar_identity(&self) -> Option<&'static str> {
173 Some(BF16_SCALAR_IDENTITY)
174 }
175
176 fn infer_output_meta(
177 &self,
178 ctx: &mut ExtensionShapeContext<'_>,
179 ) -> tenferro_tensor::Result<Vec<(DType, Vec<SymDim>)>> {
180 let dtype = ctx.input_dtype(0)?;
181 if !matches!(dtype, DType::External(_)) {
182 return Err(tenferro_tensor::Error::unsupported_dtype(
183 "bf16_einsum",
184 dtype,
185 "bf16_einsum takes an externally defined scalar",
186 ));
187 }
188 if ctx.input_dtype(1)? != dtype {
189 return Err(tenferro_tensor::Error::invalid_argument(
190 "bf16_einsum",
191 "inputs",
192 "both operands must carry the same scalar",
193 ));
194 }
195 let lhs = ctx.input_shape(0)?;
196 let rhs = ctx.input_shape(1)?;
197 if lhs.len() != self.lhs.len() || rhs.len() != self.rhs.len() {
198 return Err(tenferro_tensor::Error::rank_mismatch(
199 "bf16_einsum",
200 self.lhs.len().max(self.rhs.len()),
201 lhs.len().min(rhs.len()),
202 ));
203 }
204 let mut out_shape = Vec::with_capacity(self.out.len());
205 for label in &self.out {
206 let extent = self
207 .lhs
208 .iter()
209 .position(|candidate| candidate == label)
210 .map(|axis| lhs[axis].clone())
211 .or_else(|| {
212 self.rhs
213 .iter()
214 .position(|candidate| candidate == label)
215 .map(|axis| rhs[axis].clone())
216 })
217 .ok_or_else(|| {
218 tenferro_tensor::Error::invalid_argument(
219 "bf16_einsum",
220 "pattern",
221 "an output label must appear in at least one operand",
222 )
223 })?;
224 out_shape.push(extent);
225 }
226 Ok(vec![(dtype, out_shape)])
227 }
228}
229
230fn element_count(shape: &[usize]) -> usize {
232 shape.iter().product()
233}
234
235fn advance(index: &mut [usize], shape: &[usize]) {
237 for axis in 0..shape.len() {
238 index[axis] += 1;
239 if index[axis] < shape[axis] {
240 return;
241 }
242 index[axis] = 0;
243 }
244}
245
246fn offset_for(
248 input_labels: &[u32],
249 input_shape: &[usize],
250 out_labels: &[u32],
251 out_index: &[usize],
252 summed_labels: &[u32],
253 summed_index: &[usize],
254) -> usize {
255 let mut offset = 0usize;
256 let mut stride = 1usize;
257 for (axis, label) in input_labels.iter().enumerate() {
258 let position = out_labels
259 .iter()
260 .position(|candidate| candidate == label)
261 .map(|index| out_index[index])
262 .or_else(|| {
263 summed_labels
264 .iter()
265 .position(|candidate| candidate == label)
266 .map(|index| summed_index[index])
267 })
268 .unwrap_or(0);
269 offset += position * stride;
270 stride *= input_shape[axis];
271 }
272 offset
273}
274
275fn values_of(op: &'static str, tensor: &Tensor) -> tenferro_runtime::Result<Vec<f32>> {
277 match tensor.external_payload() {
278 Some(payload) => payload
279 .downcast_ref::<Bf16>()
280 .map(|stored| {
281 stored
282 .as_slice()
283 .iter()
284 .map(|value| value.to_f32())
285 .collect()
286 })
287 .ok_or_else(|| {
288 tenferro_runtime::Error::from(tenferro_tensor::Error::invalid_argument(
289 op,
290 "input",
291 "the operand does not carry a bfloat16 payload",
292 ))
293 }),
294 None => Err(tenferro_runtime::Error::from(
295 tenferro_tensor::Error::unsupported_dtype(
296 op,
297 tensor.dtype(),
298 "bf16_einsum takes an externally defined bfloat16 scalar",
299 ),
300 )),
301 }
302}
303
304fn contract(op: &dyn ExtensionOp, inputs: &[&Tensor]) -> tenferro_runtime::Result<Vec<Tensor>> {
306 let name = "bf16_einsum";
307 let Some(contraction) = op.as_any().downcast_ref::<Bf16Einsum>() else {
308 return Err(tenferro_runtime::Error::from(
309 tenferro_tensor::Error::invalid_argument(
310 name,
311 "payload",
312 "the operation is not a bfloat16 contraction",
313 ),
314 ));
315 };
316 if inputs.len() != 2 {
317 return Err(tenferro_runtime::Error::from(
318 tenferro_tensor::Error::invalid_argument(
319 name,
320 "input",
321 "a bf16 contraction takes two operands",
322 ),
323 ));
324 }
325 let (lhs_labels, rhs_labels, out_labels) = contraction.labels();
326 let lhs_values = values_of(name, inputs[0])?;
327 let rhs_values = values_of(name, inputs[1])?;
328 let lhs_shape = inputs[0].shape().to_vec();
329 let rhs_shape = inputs[1].shape().to_vec();
330
331 let mut extents: Vec<(u32, usize)> = Vec::new();
332 for (labels, shape) in [(lhs_labels, &lhs_shape), (rhs_labels, &rhs_shape)] {
333 for (axis, label) in labels.iter().enumerate() {
334 match extents.iter().find(|(existing, _)| existing == label) {
335 Some((_, existing)) if *existing != shape[axis] => {
336 return Err(tenferro_runtime::Error::from(
337 tenferro_tensor::Error::invalid_argument(
338 name,
339 "inputs",
340 "the operands disagree on the extent of a shared label",
341 ),
342 ))
343 }
344 Some(_) => {}
345 None => extents.push((*label, shape[axis])),
346 }
347 }
348 }
349 let extent_of = |label: u32| {
350 extents
351 .iter()
352 .find(|(existing, _)| *existing == label)
353 .map(|(_, extent)| *extent)
354 .unwrap_or(1)
355 };
356 let out_shape: Vec<usize> = out_labels.iter().map(|label| extent_of(*label)).collect();
357 let mut summed_labels: Vec<u32> = Vec::new();
358 for label in lhs_labels.iter().chain(rhs_labels.iter()) {
359 if !out_labels.contains(label) && !summed_labels.contains(label) {
360 summed_labels.push(*label);
361 }
362 }
363 let summed_shape: Vec<usize> = summed_labels
364 .iter()
365 .map(|label| extent_of(*label))
366 .collect();
367
368 let out_count = element_count(&out_shape);
369 let summed_count = element_count(&summed_shape);
370 let mut accumulated = vec![0.0_f32; out_count];
371 let mut out_index = vec![0usize; out_shape.len()];
372 let mut summed_index = vec![0usize; summed_shape.len()];
373 for slot in accumulated.iter_mut() {
374 for value in summed_index.iter_mut() {
375 *value = 0;
376 }
377 let mut total = 0.0_f32;
378 for _ in 0..summed_count {
379 let lhs_offset = offset_for(
380 lhs_labels,
381 &lhs_shape,
382 out_labels,
383 &out_index,
384 &summed_labels,
385 &summed_index,
386 );
387 let rhs_offset = offset_for(
388 rhs_labels,
389 &rhs_shape,
390 out_labels,
391 &out_index,
392 &summed_labels,
393 &summed_index,
394 );
395 total += lhs_values[lhs_offset] * rhs_values[rhs_offset];
396 advance(&mut summed_index, &summed_shape);
397 }
398 *slot = total;
399 advance(&mut out_index, &out_shape);
400 }
401
402 let rounded: Vec<Bf16> = accumulated.into_iter().map(Bf16::from_f32).collect();
404 let tensor = TypedTensor::<_, DynRank, Host>::from_host_vec_col_major(out_shape, rounded)
405 .map_err(tenferro_runtime::Error::from)?;
406 Ok(vec![Tensor::external(
407 tenferro_tensor::ErasedHostTensor::new(tensor),
408 )])
409}
410
411#[derive(Debug)]
413struct Bf16EinsumEngine<B: std::fmt::Debug + Send + Sync> {
414 family_id: &'static str,
415 engine_id: EngineId,
416 _backend: PhantomData<B>,
417}
418
419impl<B: TensorBackend + std::fmt::Debug + Send + Sync + 'static> ExtensionEngine
420 for Bf16EinsumEngine<B>
421{
422 fn family_id(&self) -> &'static str {
423 self.family_id
424 }
425
426 fn engine_id(&self) -> &EngineId {
427 &self.engine_id
428 }
429
430 fn context_identity(&self) -> ExecutionContextIdentity {
431 ExecutionContextIdentity::of::<tenferro_cpu::CpuBackend>()
432 }
433
434 fn prepare(
435 &self,
436 request: ExtensionPrepareRequest<'_>,
437 ) -> Result<PrepareCapability, PrepareError> {
438 if request.operation().family_id() != self.family_id {
439 return Err(PrepareError::ProviderContract {
440 source: ProviderContractError::WrongOperationFamily {
441 expected: CoreCapabilityKind::Elementwise,
442 operation: self.family_id,
443 },
444 });
445 }
446 let prepared = Arc::new(Bf16EinsumPrepared::<B> {
447 binding: request.binding().clone(),
448 specialization: request.specialization().clone(),
449 op: request.operation().clone_arc(),
450 _backend: PhantomData,
451 });
452 Ok(PrepareCapability::Prepared(
453 PreparedOperationPlan::executable(prepared.clone(), prepared),
454 ))
455 }
456}
457
458#[derive(Debug)]
463struct Bf16EinsumPlanning {
464 family_id: &'static str,
465}
466
467impl ExtensionPlanningConfig for Bf16EinsumPlanning {
468 fn family_id(&self) -> &'static str {
469 self.family_id
470 }
471
472 fn as_any(&self) -> &dyn Any {
473 self
474 }
475
476 fn payload_hash(&self, _state: &mut dyn Hasher) {}
477
478 fn payload_eq(&self, other: &dyn ExtensionPlanningConfig) -> bool {
479 other.family_id() == self.family_id
480 }
481
482 fn retained_bytes(&self) -> usize {
483 0
484 }
485}
486
487#[derive(Debug)]
489struct Bf16EinsumPrepared<B: std::fmt::Debug + Send + Sync> {
490 binding: PreparedOperationBinding,
491 specialization: SpecializationProjection,
492 op: Arc<dyn ExtensionOp>,
493 _backend: PhantomData<B>,
494}
495
496impl<B: TensorBackend + std::fmt::Debug + Send + Sync + 'static> PreparedOperation
497 for Bf16EinsumPrepared<B>
498{
499 fn binding(&self) -> &PreparedOperationBinding {
500 &self.binding
501 }
502
503 fn specialization(&self) -> &SpecializationProjection {
504 &self.specialization
505 }
506
507 fn retained_bytes(&self) -> usize {
508 0
509 }
510}
511
512impl<B: TensorBackend + std::fmt::Debug + Send + Sync + 'static> PreparedOperationExecutor
513 for Bf16EinsumPrepared<B>
514{
515 fn execute(
516 &self,
517 context: &mut ErasedExecutionContext<'_>,
518 extension_caches: &mut ExtensionCacheStore,
519 inputs: &[TensorRead<'_>],
520 ) -> tenferro_runtime::Result<Vec<Tensor>> {
521 let backend = context
523 .downcast_mut::<B>(self.binding.context_identity())
524 .map_err(|source| {
525 tenferro_runtime::Error::runtime_state_source(
526 "extension",
527 ErrorPhase::Execution,
528 source,
529 )
530 })?;
531 let _ = extension_caches;
532 let materialized = backend
533 .with_backend_session(|exec| {
534 inputs
535 .iter()
536 .cloned()
537 .map(|input| exec.to_contiguous_read(input))
538 .collect::<tenferro_tensor::Result<Vec<Tensor>>>()
539 })?
540 .map_err(tenferro_runtime::Error::from)?;
541 let borrowed: Vec<&Tensor> = materialized.iter().collect();
542 contract(self.op.as_ref(), &borrowed)
543 }
544}
545
546#[derive(Debug)]
548struct Bf16EinsumModule<B: std::fmt::Debug + Send + Sync> {
549 module_id: ExtensionModuleId,
550 engine_id: EngineId,
551 _backend: PhantomData<B>,
552}
553
554impl<B: TensorBackend + std::fmt::Debug + Send + Sync + 'static> ExtensionModule
555 for Bf16EinsumModule<B>
556{
557 fn module_id(&self) -> &ExtensionModuleId {
558 &self.module_id
559 }
560
561 fn configure(
562 &self,
563 registrar: &mut ExtensionModuleRegistrar<'_>,
564 ) -> Result<(), tenferro_runtime::ExtensionModuleError> {
565 registrar.register_engine(Arc::new(Bf16EinsumEngine::<B> {
566 family_id: BF16_EINSUM_FAMILY,
567 engine_id: self.engine_id.clone(),
568 _backend: PhantomData,
569 }))?;
570 registrar.register_planning_config(
571 self.engine_id.clone(),
572 Arc::new(Bf16EinsumPlanning {
573 family_id: BF16_EINSUM_FAMILY,
574 }),
575 )?;
576 Ok(())
577 }
578}
579
580pub fn module() -> Result<Arc<dyn ExtensionModule>, RuntimeConfigError> {
595 Ok(Arc::new(Bf16EinsumModule::<tenferro_cpu::CpuBackend> {
596 module_id: ExtensionModuleId::new("tenferro-bf16-proof.module")?,
597 engine_id: tenferro_cpu::runtime_engine_id()?,
598 _backend: PhantomData,
599 }))
600}