Skip to main content

binius_math/
tensor_algebra.rs

1// Copyright 2024-2025 Irreducible Inc.
2
3use std::{
4	iter::Sum,
5	marker::PhantomData,
6	ops::{Add, AddAssign, Sub, SubAssign},
7};
8
9use binius_field::{ExtensionField, Field, field::FieldOps};
10
11use crate::inner_product::inner_product;
12
13/// An element of the tensor algebra defined as the tensor product of `FE` and `FE` as fields.
14///
15/// A tensor algebra element is a length $D$ vector of `FE` field elements, where $D$ is the degree
16/// of `FE` as an extension of `F`. The algebra has a "vertical subring" and a "horizontal subring",
17/// which are both isomorphic to `FE` as a field.
18///
19/// See [DP24] Section 2 for further details.
20///
21/// [DP24]: <https://eprint.iacr.org/2024/504>
22#[derive(Debug, Clone, PartialEq, Eq)]
23pub struct TensorAlgebra<F, FE> {
24	pub elems: Vec<FE>,
25	_marker: PhantomData<F>,
26}
27
28impl<F, FE> TensorAlgebra<F, FE>
29where
30	F: Field,
31	FE: FieldOps<Scalar: ExtensionField<F>>,
32{
33	/// Constructs an element from a vector of vertical subring elements.
34	///
35	/// ## Preconditions
36	///
37	/// * `elems` must have length equal to the extension degree, otherwise this will pad or
38	///   truncate.
39	pub fn new(mut elems: Vec<FE>) -> Self {
40		elems.resize(FE::Scalar::DEGREE, FE::zero());
41		Self {
42			elems,
43			_marker: PhantomData,
44		}
45	}
46
47	/// Returns $\kappa$, the base-2 logarithm of the extension degree.
48	pub const fn kappa() -> usize {
49		FE::Scalar::LOG_DEGREE
50	}
51
52	/// Returns the multiplicative identity element, one.
53	pub fn one() -> Self {
54		let mut elems = vec![FE::zero(); FE::Scalar::DEGREE];
55		elems[0] = FE::one();
56		Self {
57			elems,
58			_marker: PhantomData,
59		}
60	}
61
62	/// Constructs a [`TensorAlgebra`] in the vertical subring.
63	pub fn from_vertical(x: FE) -> Self {
64		let mut elems = vec![FE::zero(); FE::Scalar::DEGREE];
65		elems[0] = x;
66		Self {
67			elems,
68			_marker: PhantomData,
69		}
70	}
71
72	/// Multiply by an element from the vertical subring.
73	pub fn scale_vertical(mut self, scalar: &FE) -> Self {
74		for elem_i in &mut self.elems {
75			*elem_i *= scalar;
76		}
77		self
78	}
79
80	/// Multiply by an element from the horizontal subring.
81	///
82	/// Internally, this performs a transpose, vertical scaling, then transpose sequence. If
83	/// multiple horizontal scaling operations are required and performance is a concern, it may be
84	/// better for the caller to do the transposes directly and amortize their cost.
85	pub fn scale_horizontal(self, scalar: &FE) -> Self {
86		self.transpose().scale_vertical(scalar).transpose()
87	}
88
89	/// Transposes the algebra element.
90	///
91	/// A transpose flips the vertical and horizontal subring elements.
92	pub fn transpose(mut self) -> Self {
93		FE::square_transpose::<F>(&mut self.elems);
94		self
95	}
96
97	/// Fold the tensor algebra element into a field element by scaling the rows and accumulating.
98	///
99	/// ## Preconditions
100	///
101	/// * `coeffs` must have length $2^\kappa$
102	pub fn fold_vertical(self, coeffs: &[FE]) -> FE {
103		inner_product(self.transpose().elems, coeffs.iter().cloned())
104	}
105}
106
107impl<F, FE> Default for TensorAlgebra<F, FE>
108where
109	F: Field,
110	FE: FieldOps<Scalar: ExtensionField<F>>,
111{
112	fn default() -> Self {
113		Self {
114			elems: vec![FE::zero(); FE::Scalar::DEGREE],
115			_marker: PhantomData,
116		}
117	}
118}
119
120impl<F, FE> TensorAlgebra<F, FE>
121where
122	F: Field,
123	FE: ExtensionField<F>,
124{
125	/// Tensor product of a vertical subring element and a horizontal subring element.
126	pub fn tensor(vertical: FE, horizontal: FE) -> Self {
127		let elems = horizontal
128			.iter_bases()
129			.map(|base| vertical * base)
130			.collect();
131		Self {
132			elems,
133			_marker: PhantomData,
134		}
135	}
136
137	/// If the algebra element lives in the vertical subring, this returns it as a field element.
138	pub fn try_extract_vertical(&self) -> Option<FE> {
139		self.elems
140			.iter()
141			.skip(1)
142			.all(|&elem| elem == FE::ZERO)
143			.then_some(self.elems[0])
144	}
145}
146
147impl<F, FE> Add<&Self> for TensorAlgebra<F, FE>
148where
149	F: Field,
150	FE: FieldOps<Scalar: ExtensionField<F>>,
151{
152	type Output = Self;
153
154	fn add(mut self, rhs: &Self) -> Self {
155		self.add_assign(rhs);
156		self
157	}
158}
159
160impl<F, FE> Sub<&Self> for TensorAlgebra<F, FE>
161where
162	F: Field,
163	FE: FieldOps<Scalar: ExtensionField<F>>,
164{
165	type Output = Self;
166
167	fn sub(mut self, rhs: &Self) -> Self {
168		self.sub_assign(rhs);
169		self
170	}
171}
172
173impl<F, FE> AddAssign<&Self> for TensorAlgebra<F, FE>
174where
175	F: Field,
176	FE: FieldOps<Scalar: ExtensionField<F>>,
177{
178	fn add_assign(&mut self, rhs: &Self) {
179		for (self_i, rhs_i) in self.elems.iter_mut().zip(rhs.elems.iter()) {
180			*self_i += rhs_i;
181		}
182	}
183}
184
185impl<F, FE> SubAssign<&Self> for TensorAlgebra<F, FE>
186where
187	F: Field,
188	FE: FieldOps<Scalar: ExtensionField<F>>,
189{
190	fn sub_assign(&mut self, rhs: &Self) {
191		for (self_i, rhs_i) in self.elems.iter_mut().zip(rhs.elems.iter()) {
192			*self_i -= rhs_i;
193		}
194	}
195}
196
197impl<'a, F, FE> Sum<&'a Self> for TensorAlgebra<F, FE>
198where
199	F: Field,
200	FE: FieldOps<Scalar: ExtensionField<F>>,
201{
202	fn sum<I: Iterator<Item = &'a Self>>(iter: I) -> Self {
203		iter.fold(Self::default(), |sum, item| sum + item)
204	}
205}
206
207#[cfg(test)]
208mod tests {
209	use binius_field::{BinaryField1b as B1, Random, arithmetic_traits::InvertOrZero};
210	use rand::{SeedableRng, rngs::StdRng};
211
212	use super::*;
213	use crate::test_utils::B128;
214
215	#[test]
216	fn test_tensor_product() {
217		type F = B1;
218		type FE = B128;
219
220		let mut rng = StdRng::seed_from_u64(0);
221
222		let vert = FE::random(&mut rng);
223		let hztl = FE::random(&mut rng);
224
225		let expected = TensorAlgebra::<F, _>::from_vertical(vert).scale_horizontal(&hztl);
226		assert_eq!(TensorAlgebra::tensor(vert, hztl), expected);
227	}
228
229	#[test]
230	fn test_try_extract_vertical() {
231		type F = B1;
232		type FE = B128;
233
234		let mut rng = StdRng::seed_from_u64(0);
235
236		let vert = FE::random(&mut rng);
237		let elem = TensorAlgebra::<F, _>::from_vertical(vert);
238		assert_eq!(elem.try_extract_vertical(), Some(vert));
239
240		// Scale horizontally by an extension element, and we should no longer be vertical.
241		let hztl = FE::new(1111);
242		let elem = elem.scale_horizontal(&hztl);
243		assert_eq!(elem.try_extract_vertical(), None);
244
245		// Scale back by the inverse to get back to the vertical subring.
246		// Safety: `hztl` is the non-zero constant 1111.
247		let hztl_inv = unsafe { hztl.invert() };
248		let elem = elem.scale_horizontal(&hztl_inv);
249		assert_eq!(elem.try_extract_vertical(), Some(vert));
250
251		// If we scale horizontally by an F element, we should remain in the vertical subring.
252		let hztl_subfield = FE::from(F::ONE);
253		let elem = elem.scale_horizontal(&hztl_subfield);
254		assert_eq!(elem.try_extract_vertical(), Some(vert * hztl_subfield));
255	}
256}