1mod builder;
4mod expression;
5pub(crate) mod expression_ext;
6mod flatten;
7mod variable;
8
9use alloc::collections::BTreeMap;
10use alloc::sync::Arc;
11use core::iter::{Product, Sum};
12use core::ops;
13
14pub use builder::*;
15pub use expression::{BaseLeaf, SymbolicExpression};
16pub use expression_ext::{ExtLeaf, SymbolicExpressionExt};
17use p3_field::{Dup, ExtensionField, Field, PrimeCharacteristicRing};
18pub use variable::{BaseEntry, ExtEntry, SymbolicVariable, SymbolicVariableExt};
19
20pub trait SymLeaf: Clone + core::fmt::Debug {
27 type F: Field;
29
30 const ZERO: Self;
31 const ONE: Self;
32 const TWO: Self;
33 const NEG_ONE: Self;
34
35 fn degree_multiple(&self) -> usize;
37
38 fn degree_multiple_with_transition(&self, _transition_degree: usize) -> usize {
40 self.degree_multiple()
41 }
42
43 fn poly_degree(&self, trace_len: usize, periodic_periods: &[usize]) -> usize;
46
47 fn as_const(&self) -> Option<&Self::F>;
49
50 fn from_const(c: Self::F) -> Self;
52}
53
54#[derive(Clone, Debug)]
64pub enum SymbolicExpr<A> {
65 Leaf(A),
67
68 Add {
70 x: Arc<Self>,
71 y: Arc<Self>,
72 degree_multiple: usize,
73 },
74
75 Sub {
77 x: Arc<Self>,
78 y: Arc<Self>,
79 degree_multiple: usize,
80 },
81
82 Neg {
84 x: Arc<Self>,
85 degree_multiple: usize,
86 },
87
88 Mul {
90 x: Arc<Self>,
91 y: Arc<Self>,
92 degree_multiple: usize,
93 },
94}
95
96impl<A: SymLeaf> SymbolicExpr<A> {
97 pub fn degree_multiple(&self) -> usize {
99 match self {
100 Self::Leaf(a) => a.degree_multiple(),
101 Self::Add {
102 degree_multiple, ..
103 }
104 | Self::Sub {
105 degree_multiple, ..
106 }
107 | Self::Neg {
108 degree_multiple, ..
109 }
110 | Self::Mul {
111 degree_multiple, ..
112 } => *degree_multiple,
113 }
114 }
115
116 pub fn degree_multiple_with_transition(&self, transition_degree: usize) -> usize {
123 if transition_degree == 0 {
124 return self.degree_multiple();
125 }
126 self.degree_with(&|leaf| leaf.degree_multiple_with_transition(transition_degree))
127 }
128
129 pub fn poly_degree(&self, trace_len: usize, periodic_periods: &[usize]) -> usize {
137 self.degree_with(&|leaf| leaf.poly_degree(trace_len, periodic_periods))
138 }
139
140 fn degree_with(&self, leaf_degree: &impl Fn(&A) -> usize) -> usize {
141 let mut cache: BTreeMap<*const Self, usize> = BTreeMap::new();
145 self.degree_memo(leaf_degree, &mut cache)
146 }
147
148 fn degree_memo(
149 &self,
150 leaf_degree: &impl Fn(&A) -> usize,
151 cache: &mut BTreeMap<*const Self, usize>,
152 ) -> usize {
153 match self {
154 Self::Leaf(a) => leaf_degree(a),
155 Self::Add { x, y, .. } | Self::Sub { x, y, .. } => Self::child_degree(
156 x,
157 leaf_degree,
158 cache,
159 )
160 .max(Self::child_degree(y, leaf_degree, cache)),
161 Self::Neg { x, .. } => Self::child_degree(x, leaf_degree, cache),
162 Self::Mul { x, y, .. } => {
163 Self::child_degree(x, leaf_degree, cache)
164 + Self::child_degree(y, leaf_degree, cache)
165 }
166 }
167 }
168
169 fn child_degree(
172 node: &Arc<Self>,
173 leaf_degree: &impl Fn(&A) -> usize,
174 cache: &mut BTreeMap<*const Self, usize>,
175 ) -> usize {
176 let key = Arc::as_ptr(node);
177 if let Some(°ree) = cache.get(&key) {
178 return degree;
179 }
180 let degree = node.degree_memo(leaf_degree, cache);
181 cache.insert(key, degree);
182 degree
183 }
184
185 fn as_const(&self) -> Option<&A::F> {
187 match self {
188 Self::Leaf(a) => a.as_const(),
189 _ => None,
190 }
191 }
192
193 fn sym_add(self, rhs: Self) -> Self {
195 if let (Some(&a), Some(&b)) = (self.as_const(), rhs.as_const()) {
196 return Self::Leaf(A::from_const(a + b));
197 }
198 if self.as_const().is_some_and(|c| c.is_zero()) {
199 return rhs;
200 }
201 if rhs.as_const().is_some_and(|c| c.is_zero()) {
202 return self;
203 }
204 let dm = self.degree_multiple().max(rhs.degree_multiple());
205 Self::Add {
206 x: Arc::new(self),
207 y: Arc::new(rhs),
208 degree_multiple: dm,
209 }
210 }
211
212 fn sym_sub(self, rhs: Self) -> Self {
214 if let (Some(&a), Some(&b)) = (self.as_const(), rhs.as_const()) {
215 return Self::Leaf(A::from_const(a - b));
216 }
217 if self.as_const().is_some_and(|c| c.is_zero()) {
218 return rhs.sym_neg();
219 }
220 if rhs.as_const().is_some_and(|c| c.is_zero()) {
221 return self;
222 }
223 let dm = self.degree_multiple().max(rhs.degree_multiple());
224 Self::Sub {
225 x: Arc::new(self),
226 y: Arc::new(rhs),
227 degree_multiple: dm,
228 }
229 }
230
231 fn sym_neg(self) -> Self {
233 if let Some(&c) = self.as_const() {
234 return Self::Leaf(A::from_const(-c));
235 }
236 let dm = self.degree_multiple();
237 Self::Neg {
238 x: Arc::new(self),
239 degree_multiple: dm,
240 }
241 }
242
243 fn sym_mul(self, rhs: Self) -> Self {
245 if let (Some(&a), Some(&b)) = (self.as_const(), rhs.as_const()) {
246 return Self::Leaf(A::from_const(a * b));
247 }
248 if self.as_const().is_some_and(|c| c.is_zero())
249 || rhs.as_const().is_some_and(|c| c.is_zero())
250 {
251 return Self::Leaf(A::from_const(A::F::ZERO));
252 }
253 if self.as_const().is_some_and(|c| c.is_one()) {
254 return rhs;
255 }
256 if rhs.as_const().is_some_and(|c| c.is_one()) {
257 return self;
258 }
259 let dm = self.degree_multiple() + rhs.degree_multiple();
260 Self::Mul {
261 x: Arc::new(self),
262 y: Arc::new(rhs),
263 degree_multiple: dm,
264 }
265 }
266}
267
268impl<A: SymLeaf> PrimeCharacteristicRing for SymbolicExpr<A> {
269 type PrimeSubfield = <A::F as PrimeCharacteristicRing>::PrimeSubfield;
270
271 const ZERO: Self = Self::Leaf(A::ZERO);
272 const ONE: Self = Self::Leaf(A::ONE);
273 const TWO: Self = Self::Leaf(A::TWO);
274 const NEG_ONE: Self = Self::Leaf(A::NEG_ONE);
275
276 #[inline]
277 fn from_prime_subfield(f: Self::PrimeSubfield) -> Self {
278 Self::Leaf(A::from_const(A::F::from_prime_subfield(f)))
279 }
280}
281
282impl<A: SymLeaf> Dup for SymbolicExpr<A> {
283 #[inline(always)]
284 fn dup(&self) -> Self {
285 self.clone()
286 }
287}
288
289impl<A: SymLeaf> Default for SymbolicExpr<A> {
290 fn default() -> Self {
291 Self::ZERO
292 }
293}
294
295impl<A: SymLeaf, T: Into<Self>> ops::Add<T> for SymbolicExpr<A> {
296 type Output = Self;
297 fn add(self, rhs: T) -> Self {
298 self.sym_add(rhs.into())
299 }
300}
301
302impl<A: SymLeaf, T: Into<Self>> ops::Sub<T> for SymbolicExpr<A> {
303 type Output = Self;
304 fn sub(self, rhs: T) -> Self {
305 self.sym_sub(rhs.into())
306 }
307}
308
309impl<A: SymLeaf> ops::Neg for SymbolicExpr<A> {
310 type Output = Self;
311 fn neg(self) -> Self {
312 self.sym_neg()
313 }
314}
315
316impl<A: SymLeaf, T: Into<Self>> ops::Mul<T> for SymbolicExpr<A> {
317 type Output = Self;
318 fn mul(self, rhs: T) -> Self {
319 self.sym_mul(rhs.into())
320 }
321}
322
323impl<A: SymLeaf, T: Into<Self>> ops::AddAssign<T> for SymbolicExpr<A> {
324 fn add_assign(&mut self, rhs: T) {
325 *self = self.clone() + rhs.into();
326 }
327}
328
329impl<A: SymLeaf, T: Into<Self>> ops::SubAssign<T> for SymbolicExpr<A> {
330 fn sub_assign(&mut self, rhs: T) {
331 *self = self.clone() - rhs.into();
332 }
333}
334
335impl<A: SymLeaf, T: Into<Self>> ops::MulAssign<T> for SymbolicExpr<A> {
336 fn mul_assign(&mut self, rhs: T) {
337 *self = self.clone() * rhs.into();
338 }
339}
340
341impl<A: SymLeaf, T: Into<Self>> Sum<T> for SymbolicExpr<A> {
342 fn sum<I: Iterator<Item = T>>(iter: I) -> Self {
343 iter.map(Into::into)
344 .reduce(|a, b| a + b)
345 .unwrap_or(Self::ZERO)
346 }
347}
348
349impl<A: SymLeaf, T: Into<Self>> Product<T> for SymbolicExpr<A> {
350 fn product<I: Iterator<Item = T>>(iter: I) -> Self {
351 iter.map(Into::into)
352 .reduce(|a, b| a * b)
353 .unwrap_or(Self::ONE)
354 }
355}
356
357impl<F: Field, T: Into<SymbolicExpression<F>>> ops::Add<T> for SymbolicVariable<F> {
358 type Output = SymbolicExpression<F>;
359 fn add(self, rhs: T) -> Self::Output {
360 Self::Output::from(self) + rhs.into()
361 }
362}
363
364impl<F: Field, T: Into<SymbolicExpression<F>>> ops::Sub<T> for SymbolicVariable<F> {
365 type Output = SymbolicExpression<F>;
366 fn sub(self, rhs: T) -> Self::Output {
367 Self::Output::from(self) - rhs.into()
368 }
369}
370
371impl<F: Field, T: Into<SymbolicExpression<F>>> ops::Mul<T> for SymbolicVariable<F> {
372 type Output = SymbolicExpression<F>;
373 fn mul(self, rhs: T) -> Self::Output {
374 Self::Output::from(self) * rhs.into()
375 }
376}
377
378impl<F: Field, EF: ExtensionField<F>, T: Into<SymbolicExpressionExt<F, EF>>> ops::Add<T>
379 for SymbolicVariableExt<F, EF>
380{
381 type Output = SymbolicExpressionExt<F, EF>;
382 fn add(self, rhs: T) -> Self::Output {
383 Self::Output::from(self) + rhs.into()
384 }
385}
386
387impl<F: Field, EF: ExtensionField<F>, T: Into<SymbolicExpressionExt<F, EF>>> ops::Sub<T>
388 for SymbolicVariableExt<F, EF>
389{
390 type Output = SymbolicExpressionExt<F, EF>;
391 fn sub(self, rhs: T) -> Self::Output {
392 Self::Output::from(self) - rhs.into()
393 }
394}
395
396impl<F: Field, EF: ExtensionField<F>, T: Into<SymbolicExpressionExt<F, EF>>> ops::Mul<T>
397 for SymbolicVariableExt<F, EF>
398{
399 type Output = SymbolicExpressionExt<F, EF>;
400 fn mul(self, rhs: T) -> Self::Output {
401 Self::Output::from(self) * rhs.into()
402 }
403}
404
405#[cfg(test)]
406mod tests {
407 use p3_baby_bear::BabyBear;
408 use p3_field::extension::BinomialExtensionField;
409
410 use super::*;
411 use crate::symbolic::expression::BaseLeaf;
412 use crate::symbolic::expression_ext::ExtLeaf;
413 use crate::symbolic::variable::{BaseEntry, ExtEntry};
414
415 type F = BabyBear;
416 type EF = BinomialExtensionField<BabyBear, 4>;
417
418 #[test]
419 fn symbolic_variable_add_produces_add_node() {
420 let var = SymbolicVariable::<F>::new(BaseEntry::Main { offset: 0 }, 0);
422 let expr = SymbolicExpression::from(F::new(5));
423 let result = var + expr;
424 match result {
425 SymbolicExpr::Add {
426 x,
427 y,
428 degree_multiple,
429 } => {
430 assert_eq!(degree_multiple, 1);
431 assert!(matches!(
432 x.as_ref(),
433 SymbolicExpr::Leaf(BaseLeaf::Variable(v))
434 if v.index == 0 && v.entry == BaseEntry::Main { offset: 0 }
435 ));
436 assert!(matches!(
437 y.as_ref(),
438 SymbolicExpr::Leaf(BaseLeaf::Constant(c)) if *c == F::new(5)
439 ));
440 }
441 _ => panic!("Expected an Add node"),
442 }
443 }
444
445 #[test]
446 fn symbolic_variable_sub_produces_sub_node() {
447 let var = SymbolicVariable::<F>::new(BaseEntry::Main { offset: 0 }, 0);
449 let other = SymbolicExpression::Leaf(BaseLeaf::Variable(SymbolicVariable::new(
450 BaseEntry::Main { offset: 0 },
451 1,
452 )));
453 let result = var - other;
454 match result {
455 SymbolicExpr::Sub {
456 x,
457 y,
458 degree_multiple,
459 } => {
460 assert_eq!(degree_multiple, 1);
461 assert!(matches!(
462 x.as_ref(),
463 SymbolicExpr::Leaf(BaseLeaf::Variable(v))
464 if v.index == 0 && v.entry == BaseEntry::Main { offset: 0 }
465 ));
466 assert!(matches!(
467 y.as_ref(),
468 SymbolicExpr::Leaf(BaseLeaf::Variable(v))
469 if v.index == 1 && v.entry == BaseEntry::Main { offset: 0 }
470 ));
471 }
472 _ => panic!("Expected a Sub node"),
473 }
474 }
475
476 #[test]
477 fn symbolic_variable_mul_produces_mul_node() {
478 let var = SymbolicVariable::<F>::new(BaseEntry::Main { offset: 0 }, 0);
480 let other = SymbolicExpression::Leaf(BaseLeaf::Variable(SymbolicVariable::new(
481 BaseEntry::Main { offset: 0 },
482 1,
483 )));
484 let result = var * other;
485 match result {
486 SymbolicExpr::Mul {
487 x,
488 y,
489 degree_multiple,
490 } => {
491 assert_eq!(degree_multiple, 2);
492 assert!(matches!(
493 x.as_ref(),
494 SymbolicExpr::Leaf(BaseLeaf::Variable(v))
495 if v.index == 0 && v.entry == BaseEntry::Main { offset: 0 }
496 ));
497 assert!(matches!(
498 y.as_ref(),
499 SymbolicExpr::Leaf(BaseLeaf::Variable(v))
500 if v.index == 1 && v.entry == BaseEntry::Main { offset: 0 }
501 ));
502 }
503 _ => panic!("Expected a Mul node"),
504 }
505 }
506
507 #[test]
508 fn symbolic_variable_ext_add_produces_add_node() {
509 let var = SymbolicVariableExt::<F, EF>::new(ExtEntry::Permutation { offset: 0 }, 0);
511 let expr = SymbolicExpressionExt::<F, EF>::from(F::new(3));
512 let result = var + expr;
513 match result {
514 SymbolicExpr::Add {
515 x,
516 y,
517 degree_multiple,
518 } => {
519 assert_eq!(degree_multiple, 1);
520 assert!(matches!(
521 x.as_ref(),
522 SymbolicExpr::Leaf(ExtLeaf::ExtVariable(v))
523 if v.index == 0 && v.entry == ExtEntry::Permutation { offset: 0 }
524 ));
525 assert!(matches!(
526 y.as_ref(),
527 SymbolicExpr::Leaf(ExtLeaf::Base(SymbolicExpr::Leaf(BaseLeaf::Constant(c))))
528 if *c == F::new(3)
529 ));
530 }
531 _ => panic!("Expected an Add node"),
532 }
533 }
534
535 #[test]
536 fn symbolic_variable_ext_sub_produces_sub_node() {
537 let var = SymbolicVariableExt::<F, EF>::new(ExtEntry::Permutation { offset: 0 }, 0);
539 let other = SymbolicExpressionExt::<F, EF>::from(SymbolicVariableExt::<F, EF>::new(
540 ExtEntry::Permutation { offset: 0 },
541 1,
542 ));
543 let result = var - other;
544 match result {
545 SymbolicExpr::Sub {
546 x,
547 y,
548 degree_multiple,
549 } => {
550 assert_eq!(degree_multiple, 1);
551 assert!(matches!(
552 x.as_ref(),
553 SymbolicExpr::Leaf(ExtLeaf::ExtVariable(v))
554 if v.index == 0 && v.entry == ExtEntry::Permutation { offset: 0 }
555 ));
556 assert!(matches!(
557 y.as_ref(),
558 SymbolicExpr::Leaf(ExtLeaf::ExtVariable(v))
559 if v.index == 1 && v.entry == ExtEntry::Permutation { offset: 0 }
560 ));
561 }
562 _ => panic!("Expected a Sub node"),
563 }
564 }
565
566 #[test]
567 fn symbolic_variable_ext_mul_produces_mul_node() {
568 let var = SymbolicVariableExt::<F, EF>::new(ExtEntry::Permutation { offset: 0 }, 0);
570 let other = SymbolicExpressionExt::<F, EF>::from(SymbolicVariableExt::<F, EF>::new(
571 ExtEntry::Permutation { offset: 0 },
572 1,
573 ));
574 let result = var * other;
575 match result {
576 SymbolicExpr::Mul {
577 x,
578 y,
579 degree_multiple,
580 } => {
581 assert_eq!(degree_multiple, 2);
582 assert!(matches!(
583 x.as_ref(),
584 SymbolicExpr::Leaf(ExtLeaf::ExtVariable(v))
585 if v.index == 0 && v.entry == ExtEntry::Permutation { offset: 0 }
586 ));
587 assert!(matches!(
588 y.as_ref(),
589 SymbolicExpr::Leaf(ExtLeaf::ExtVariable(v))
590 if v.index == 1 && v.entry == ExtEntry::Permutation { offset: 0 }
591 ));
592 }
593 _ => panic!("Expected a Mul node"),
594 }
595 }
596}