Skip to main content

tenferro_bf16_proof/
einsum.rs

1//! A bfloat16 contraction through the runtime's extension module.
2//!
3//! #1793's precision table lists a bf16 CPU einsum "with documented f32 internal
4//! calculation/accumulation", and asks for a contraction whose result distinguishes f32
5//! accumulation from repeated bfloat16 rounding, compared against an independent reference with the
6//! specified output rounding. This module supplies that: the operands are widened to `f32`, the
7//! contraction accumulates there, and the result is rounded to bfloat16 once.
8//!
9//! What it does not do is inherit a wider pattern surface. The operation is the pairwise
10//! contraction the table's example needs, so a trace, a repeated label, or more than two operands is
11//! refused with a typed error rather than folded, which keeps this module's body to the
12//! accumulation contract the row is about.
13
14use 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
32/// The family the bfloat16 contraction belongs to.
33pub const BF16_EINSUM_FAMILY: &str = "tenferro-bf16-proof.einsum.v1";
34
35/// The canonical identity of the bfloat16 scalar a program declares.
36pub const BF16_SCALAR_IDENTITY: &str = "tenferro-bf16-proof.bf16.v1";
37
38/// A pairwise contraction in bfloat16.
39///
40/// The labels are the ones an ordinary einsum subscript string names, so `"ik,kj->ij"` is
41/// `(&[0, 1], &[1, 2], &[0, 2])`. The operands are widened to `f32`, the contraction accumulates
42/// there, and the result is rounded to bfloat16 once.
43///
44/// # Examples
45///
46/// ```rust
47/// use tenferro_bf16_proof::einsum::Bf16Einsum;
48///
49/// let op = Bf16Einsum::new(&[0, 1], &[1, 2], &[0, 2]).expect("a contraction");
50/// assert_eq!(op.labels(), (&[0, 1][..], &[1, 2][..], &[0, 2][..]));
51/// ```
52#[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    /// Build the pairwise contraction `lhs,rhs->out`.
61    ///
62    /// # Errors
63    ///
64    /// Returns [`tenferro_tensor::Error::InvalidArgument`] when the pattern is not a pairwise contraction: an operand
65    /// carries no label, a label repeats inside one operand, the operands share no contracted
66    /// label, or an output label appears in no operand.
67    ///
68    /// # Examples
69    ///
70    /// ```rust
71    /// use tenferro_bf16_proof::einsum::Bf16Einsum;
72    ///
73    /// assert!(Bf16Einsum::new(&[0, 1], &[1, 2], &[0, 2]).is_ok());
74    /// assert!(Bf16Einsum::new(&[], &[1, 2], &[0, 2]).is_err());
75    /// ```
76    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    /// The pattern's three label lists.
112    ///
113    /// # Examples
114    ///
115    /// ```rust
116    /// use tenferro_bf16_proof::einsum::Bf16Einsum;
117    ///
118    /// let op = Bf16Einsum::new(&[0, 1], &[1, 2], &[0, 2]).expect("a contraction");
119    /// assert_eq!(op.labels(), (&[0, 1][..], &[1, 2][..], &[0, 2][..]));
120    /// ```
121    #[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
230/// The number of elements a shape describes.
231fn element_count(shape: &[usize]) -> usize {
232    shape.iter().product()
233}
234
235/// Advance a column-major odometer by one.
236fn 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
246/// The column-major offset an operand's labels select at one output and contracted index.
247fn 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
275/// The stored values of a bfloat16 tensor, widened to `f32`.
276fn 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
304/// Contract two operands, accumulating in `f32` and rounding once.
305fn 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    // One rounding, after the accumulation rather than at every step.
403    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/// The engine the bfloat16 contraction is registered under.
412#[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/// The planning config the runtime keys the bfloat16 family under.
459///
460/// The runtime keeps one planning config per engine identity, so the family states its identity
461/// here rather than leaving the engine without one.
462#[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/// The prepared bfloat16 contraction.
488#[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        // The body needs contiguous inputs, which a session on the binding's backend provides.
522        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/// The module an application installs to run the bfloat16 contraction.
547#[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
580/// The module for the CPU backend's engine identity.
581///
582/// # Errors
583///
584/// Returns [`RuntimeConfigError`] when the module identity cannot be built or the backend has no
585/// engine identity.
586///
587/// # Examples
588///
589/// ```rust
590/// use tenferro_bf16_proof::einsum::module;
591///
592/// assert!(module().is_ok());
593/// ```
594pub 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}