binius_math/
tensor_algebra.rs1use 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#[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 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 pub const fn kappa() -> usize {
49 FE::Scalar::LOG_DEGREE
50 }
51
52 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 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 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 pub fn scale_horizontal(self, scalar: &FE) -> Self {
86 self.transpose().scale_vertical(scalar).transpose()
87 }
88
89 pub fn transpose(mut self) -> Self {
93 FE::square_transpose::<F>(&mut self.elems);
94 self
95 }
96
97 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 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 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 let hztl = FE::new(1111);
242 let elem = elem.scale_horizontal(&hztl);
243 assert_eq!(elem.try_extract_vertical(), None);
244
245 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 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}