use std::sync::Arc;
use p3_field::Field;
use rustc_hash::FxHashMap;
use serde::{Deserialize, Serialize};
use super::SymbolicConstraints;
use crate::{
air_builders::symbolic::{
symbolic_expression::SymbolicExpression, symbolic_variable::SymbolicVariable,
},
interaction::{Interaction, SymbolicInteraction},
};
#[derive(Clone, Debug, Hash, Serialize, Deserialize, PartialEq, Eq)]
#[serde(bound(serialize = "F: Serialize", deserialize = "F: Deserialize<'de>"))]
#[repr(C)]
pub enum SymbolicExpressionNode<F> {
Variable(SymbolicVariable<F>),
IsFirstRow,
IsLastRow,
IsTransition,
Constant(F),
Add {
left_idx: usize,
right_idx: usize,
degree_multiple: usize,
},
Sub {
left_idx: usize,
right_idx: usize,
degree_multiple: usize,
},
Neg {
idx: usize,
degree_multiple: usize,
},
Mul {
left_idx: usize,
right_idx: usize,
degree_multiple: usize,
},
}
#[derive(Clone, Debug, Serialize, Deserialize, PartialEq, Eq)]
#[serde(bound(serialize = "F: Serialize", deserialize = "F: Deserialize<'de>"))]
#[repr(C)]
pub struct SymbolicExpressionDag<F> {
pub nodes: Vec<SymbolicExpressionNode<F>>,
pub constraint_idx: Vec<usize>,
}
impl<F> SymbolicExpressionDag<F> {
pub fn max_rotation(&self) -> usize {
let mut rotation = 0;
for node in &self.nodes {
if let SymbolicExpressionNode::Variable(var) = node {
rotation = rotation.max(var.entry.offset().unwrap_or(0));
}
}
rotation
}
pub fn num_constraints(&self) -> usize {
self.constraint_idx.len()
}
}
#[derive(Clone, Debug, Serialize, Deserialize)]
#[serde(bound(serialize = "F: Serialize", deserialize = "F: Deserialize<'de>"))]
#[repr(C)]
pub struct SymbolicConstraintsDag<F> {
pub constraints: SymbolicExpressionDag<F>,
pub interactions: Vec<Interaction<usize>>,
}
pub(crate) fn build_symbolic_constraints_dag<F: Field>(
constraints: &[SymbolicExpression<F>],
interactions: &[SymbolicInteraction<F>],
) -> SymbolicConstraintsDag<F> {
let mut builder = SymbolicDagBuilder::new();
let mut constraint_idx: Vec<usize> = constraints
.iter()
.map(|expr| builder.add_expr(expr))
.collect();
constraint_idx.sort();
constraint_idx.dedup();
let interactions: Vec<Interaction<usize>> = interactions
.iter()
.map(|interaction| {
let fields: Vec<usize> = interaction
.message
.iter()
.map(|field_expr| builder.add_expr(field_expr))
.collect();
let count = builder.add_expr(&interaction.count);
Interaction {
message: fields,
count,
bus_index: interaction.bus_index,
count_weight: interaction.count_weight,
}
})
.collect();
let constraints = SymbolicExpressionDag {
nodes: builder.nodes,
constraint_idx,
};
SymbolicConstraintsDag {
constraints,
interactions,
}
}
pub struct SymbolicDagBuilder<F: Field> {
pub expr_to_idx: FxHashMap<*const SymbolicExpression<F>, usize>,
pub node_to_idx: FxHashMap<SymbolicExpressionNode<F>, usize>,
pub nodes: Vec<SymbolicExpressionNode<F>>,
}
impl<F: Field> Default for SymbolicDagBuilder<F> {
fn default() -> Self {
Self::new()
}
}
impl<F: Field> SymbolicDagBuilder<F> {
pub fn new() -> Self {
Self {
expr_to_idx: FxHashMap::default(),
node_to_idx: FxHashMap::default(),
nodes: Vec::new(),
}
}
pub fn add_expr(&mut self, expr: &SymbolicExpression<F>) -> usize {
let ptr = expr as *const SymbolicExpression<F>;
if let Some(&idx) = self.expr_to_idx.get(&ptr) {
return idx;
}
let idx = match expr {
SymbolicExpression::Variable(var) => {
self.intern_node(SymbolicExpressionNode::Variable(*var))
}
SymbolicExpression::IsFirstRow => self.intern_node(SymbolicExpressionNode::IsFirstRow),
SymbolicExpression::IsLastRow => self.intern_node(SymbolicExpressionNode::IsLastRow),
SymbolicExpression::IsTransition => {
self.intern_node(SymbolicExpressionNode::IsTransition)
}
SymbolicExpression::Constant(cons) => {
self.intern_node(SymbolicExpressionNode::Constant(*cons))
}
SymbolicExpression::Add {
x,
y,
degree_multiple,
} => {
let left_idx = self.add_expr(x.as_ref());
let right_idx = self.add_expr(y.as_ref());
if let (Some(a), Some(b)) = (self.get_const(left_idx), self.get_const(right_idx)) {
self.intern_node(SymbolicExpressionNode::Constant(a + b))
}
else if self.is_const(left_idx, F::ZERO) {
right_idx
} else if self.is_const(right_idx, F::ZERO) {
left_idx
}
else if let Some(neg_child_idx) = self.get_neg_child(right_idx) {
self.intern_node(SymbolicExpressionNode::Sub {
left_idx,
right_idx: neg_child_idx,
degree_multiple: *degree_multiple,
})
} else {
self.intern_node(SymbolicExpressionNode::Add {
left_idx,
right_idx,
degree_multiple: *degree_multiple,
})
}
}
SymbolicExpression::Sub {
x,
y,
degree_multiple,
} => {
let left_idx = self.add_expr(x.as_ref());
let right_idx = self.add_expr(y.as_ref());
if let (Some(a), Some(b)) = (self.get_const(left_idx), self.get_const(right_idx)) {
self.intern_node(SymbolicExpressionNode::Constant(a - b))
}
else if self.is_const(right_idx, F::ZERO) {
left_idx
}
else if let Some(neg_child_idx) = self.get_neg_child(right_idx) {
self.intern_node(SymbolicExpressionNode::Add {
left_idx,
right_idx: neg_child_idx,
degree_multiple: *degree_multiple,
})
} else {
self.intern_node(SymbolicExpressionNode::Sub {
left_idx,
right_idx,
degree_multiple: *degree_multiple,
})
}
}
SymbolicExpression::Neg { x, degree_multiple } => {
let child_idx = self.add_expr(x.as_ref());
if let Some(c) = self.get_const(child_idx) {
self.intern_node(SymbolicExpressionNode::Constant(-c))
} else {
self.intern_node(SymbolicExpressionNode::Neg {
idx: child_idx,
degree_multiple: *degree_multiple,
})
}
}
SymbolicExpression::Mul {
x,
y,
degree_multiple,
} => {
let left_idx = self.add_expr(x.as_ref());
let right_idx = self.add_expr(y.as_ref());
if let (Some(a), Some(b)) = (self.get_const(left_idx), self.get_const(right_idx)) {
self.intern_node(SymbolicExpressionNode::Constant(a * b))
}
else if self.is_const(left_idx, F::ZERO) || self.is_const(right_idx, F::ONE) {
left_idx
} else if self.is_const(right_idx, F::ZERO) || self.is_const(left_idx, F::ONE) {
right_idx
} else {
self.intern_node(SymbolicExpressionNode::Mul {
left_idx,
right_idx,
degree_multiple: *degree_multiple,
})
}
}
};
self.expr_to_idx.insert(ptr, idx);
idx
}
fn intern_node(&mut self, node: SymbolicExpressionNode<F>) -> usize {
*self.node_to_idx.entry(node.clone()).or_insert_with(|| {
let idx = self.nodes.len();
self.nodes.push(node);
idx
})
}
fn is_const(&self, idx: usize, val: F) -> bool {
matches!(&self.nodes[idx], SymbolicExpressionNode::Constant(c) if *c == val)
}
fn get_const(&self, idx: usize) -> Option<F> {
match &self.nodes[idx] {
SymbolicExpressionNode::Constant(c) => Some(*c),
_ => None,
}
}
fn get_neg_child(&self, idx: usize) -> Option<usize> {
match &self.nodes[idx] {
SymbolicExpressionNode::Neg { idx, .. } => Some(*idx),
_ => None,
}
}
}
impl<F: Field> SymbolicExpressionDag<F> {
fn to_symbolic_expressions(&self) -> Vec<Arc<SymbolicExpression<F>>> {
let mut exprs: Vec<Arc<SymbolicExpression<_>>> = Vec::with_capacity(self.nodes.len());
for node in &self.nodes {
let expr = match *node {
SymbolicExpressionNode::Variable(var) => SymbolicExpression::Variable(var),
SymbolicExpressionNode::IsFirstRow => SymbolicExpression::IsFirstRow,
SymbolicExpressionNode::IsLastRow => SymbolicExpression::IsLastRow,
SymbolicExpressionNode::IsTransition => SymbolicExpression::IsTransition,
SymbolicExpressionNode::Constant(f) => SymbolicExpression::Constant(f),
SymbolicExpressionNode::Add {
left_idx,
right_idx,
degree_multiple,
} => SymbolicExpression::Add {
x: exprs[left_idx].clone(),
y: exprs[right_idx].clone(),
degree_multiple,
},
SymbolicExpressionNode::Sub {
left_idx,
right_idx,
degree_multiple,
} => SymbolicExpression::Sub {
x: exprs[left_idx].clone(),
y: exprs[right_idx].clone(),
degree_multiple,
},
SymbolicExpressionNode::Neg {
idx,
degree_multiple,
} => SymbolicExpression::Neg {
x: exprs[idx].clone(),
degree_multiple,
},
SymbolicExpressionNode::Mul {
left_idx,
right_idx,
degree_multiple,
} => SymbolicExpression::Mul {
x: exprs[left_idx].clone(),
y: exprs[right_idx].clone(),
degree_multiple,
},
};
exprs.push(Arc::new(expr));
}
exprs
}
}
impl<'a, F: Field> From<&'a SymbolicConstraintsDag<F>> for SymbolicConstraints<F> {
fn from(dag: &'a SymbolicConstraintsDag<F>) -> Self {
let exprs = dag.constraints.to_symbolic_expressions();
let constraints = dag
.constraints
.constraint_idx
.iter()
.map(|&idx| exprs[idx].as_ref().clone())
.collect::<Vec<_>>();
let interactions = dag
.interactions
.iter()
.map(|interaction| {
let fields = interaction
.message
.iter()
.map(|&idx| exprs[idx].as_ref().clone())
.collect();
let count = exprs[interaction.count].as_ref().clone();
Interaction {
message: fields,
count,
bus_index: interaction.bus_index,
count_weight: interaction.count_weight,
}
})
.collect::<Vec<_>>();
SymbolicConstraints {
constraints,
interactions,
}
}
}
impl<F: Field> From<SymbolicConstraintsDag<F>> for SymbolicConstraints<F> {
fn from(dag: SymbolicConstraintsDag<F>) -> Self {
(&dag).into()
}
}
impl<F: Field> From<SymbolicConstraints<F>> for SymbolicConstraintsDag<F> {
fn from(sc: SymbolicConstraints<F>) -> Self {
build_symbolic_constraints_dag(&sc.constraints, &sc.interactions)
}
}
#[cfg(test)]
mod tests {
use p3_baby_bear::BabyBear;
use p3_field::PrimeCharacteristicRing;
use crate::{
air_builders::symbolic::{
dag::{build_symbolic_constraints_dag, SymbolicExpressionDag, SymbolicExpressionNode},
symbolic_expression::SymbolicExpression,
symbolic_variable::{Entry, SymbolicVariable},
},
interaction::Interaction,
};
type F = BabyBear;
#[test]
fn test_duplicate_constraints_are_deduplicated() {
let expr: SymbolicExpression<F> = SymbolicExpression::Variable(SymbolicVariable::new(
Entry::Main {
part_index: 0,
offset: 0,
},
0,
));
let constraints = vec![expr.clone(), expr.clone()];
let interactions = vec![];
let dag = build_symbolic_constraints_dag(&constraints, &interactions);
assert_eq!(
dag.constraints.nodes.len(),
1,
"Nodes should be deduplicated"
);
assert_eq!(
dag.constraints.constraint_idx,
vec![0],
"constraint_idx should be deduplicated"
);
assert_eq!(
dag.constraints.num_constraints(),
1,
"Duplicate constraints should be deduplicated"
);
}
#[test]
fn test_structural_deduplication() {
let var = SymbolicVariable::<F>::new(
Entry::Main {
part_index: 0,
offset: 0,
},
0,
);
let expr1 = SymbolicExpression::from(var) - SymbolicExpression::Constant(F::ONE);
let expr2 = SymbolicExpression::from(var) - SymbolicExpression::Constant(F::ONE);
assert!(!std::ptr::eq(&expr1, &expr2));
let constraints = vec![expr1, expr2];
let dag = build_symbolic_constraints_dag(&constraints, &[]);
assert_eq!(dag.constraints.nodes.len(), 3);
assert_eq!(dag.constraints.constraint_idx, vec![2]);
}
#[test]
fn test_algebraic_simplifications() {
let var = SymbolicVariable::<F>::new(
Entry::Main {
part_index: 0,
offset: 0,
},
0,
);
let x = SymbolicExpression::from(var);
let zero = SymbolicExpression::Constant(F::ZERO);
let one = SymbolicExpression::Constant(F::ONE);
let expr_add_zero = x.clone() + zero.clone();
let dag = build_symbolic_constraints_dag(&[expr_add_zero], &[]);
assert_eq!(dag.constraints.nodes.len(), 2); assert!(matches!(
dag.constraints.nodes[dag.constraints.constraint_idx[0]],
SymbolicExpressionNode::Variable(_)
));
let expr_zero_add = zero.clone() + x.clone();
let dag = build_symbolic_constraints_dag(&[expr_zero_add], &[]);
assert!(matches!(
dag.constraints.nodes[dag.constraints.constraint_idx[0]],
SymbolicExpressionNode::Variable(_)
));
let expr_mul_one = x.clone() * one.clone();
let dag = build_symbolic_constraints_dag(&[expr_mul_one], &[]);
assert!(matches!(
dag.constraints.nodes[dag.constraints.constraint_idx[0]],
SymbolicExpressionNode::Variable(_)
));
let expr_one_mul = one.clone() * x.clone();
let dag = build_symbolic_constraints_dag(&[expr_one_mul], &[]);
assert!(matches!(
dag.constraints.nodes[dag.constraints.constraint_idx[0]],
SymbolicExpressionNode::Variable(_)
));
let expr_mul_zero = x.clone() * zero.clone();
let dag = build_symbolic_constraints_dag(&[expr_mul_zero], &[]);
assert!(matches!(
dag.constraints.nodes[dag.constraints.constraint_idx[0]],
SymbolicExpressionNode::Constant(c) if c == F::ZERO
));
let expr_sub_zero = x.clone() - zero.clone();
let dag = build_symbolic_constraints_dag(&[expr_sub_zero], &[]);
assert!(matches!(
dag.constraints.nodes[dag.constraints.constraint_idx[0]],
SymbolicExpressionNode::Variable(_)
));
let y = SymbolicExpression::from(SymbolicVariable::<F>::new(
Entry::Main {
part_index: 0,
offset: 0,
},
1,
));
let expr_add_neg = x.clone() + (-y.clone());
let expr_sub = x.clone() - y.clone();
let dag1 = build_symbolic_constraints_dag(&[expr_add_neg], &[]);
let dag2 = build_symbolic_constraints_dag(&[expr_sub], &[]);
assert!(matches!(
dag1.constraints.nodes[dag1.constraints.constraint_idx[0]],
SymbolicExpressionNode::Sub { .. }
));
assert_eq!(
dag1.constraints.nodes[dag1.constraints.constraint_idx[0]],
dag2.constraints.nodes[dag2.constraints.constraint_idx[0]]
);
let expr_sub_neg = x.clone() - (-y.clone());
let dag = build_symbolic_constraints_dag(&[expr_sub_neg], &[]);
assert!(matches!(
dag.constraints.nodes[dag.constraints.constraint_idx[0]],
SymbolicExpressionNode::Add { .. }
));
}
#[test]
fn test_constant_folding() {
let two = SymbolicExpression::<F>::Constant(F::TWO);
let three = SymbolicExpression::<F>::Constant(F::from_u32(3));
let five = F::from_u32(5);
let six = F::from_u32(6);
let neg_three = -F::from_u32(3);
let expr_add = two.clone() + three.clone();
let dag = build_symbolic_constraints_dag(&[expr_add], &[]);
assert!(matches!(
dag.constraints.nodes[dag.constraints.constraint_idx[0]],
SymbolicExpressionNode::Constant(c) if c == five
));
let expr_sub = three.clone() - two.clone();
let dag = build_symbolic_constraints_dag(&[expr_sub], &[]);
assert!(matches!(
dag.constraints.nodes[dag.constraints.constraint_idx[0]],
SymbolicExpressionNode::Constant(c) if c == F::ONE
));
let expr_mul = two.clone() * three.clone();
let dag = build_symbolic_constraints_dag(&[expr_mul], &[]);
assert!(matches!(
dag.constraints.nodes[dag.constraints.constraint_idx[0]],
SymbolicExpressionNode::Constant(c) if c == six
));
let expr_neg = -three.clone();
let dag = build_symbolic_constraints_dag(&[expr_neg], &[]);
assert!(matches!(
dag.constraints.nodes[dag.constraints.constraint_idx[0]],
SymbolicExpressionNode::Constant(c) if c == neg_three
));
let expr_chain = (two.clone() + three.clone()) * two.clone();
let dag = build_symbolic_constraints_dag(&[expr_chain], &[]);
assert!(matches!(
dag.constraints.nodes[dag.constraints.constraint_idx[0]],
SymbolicExpressionNode::Constant(c) if c == F::from_u32(10)
));
}
#[test]
fn test_symbolic_constraints_dag() {
let expr = SymbolicExpression::Constant(F::ONE)
* SymbolicVariable::new(
Entry::Main {
part_index: 1,
offset: 2,
},
3,
);
let constraints = vec![
SymbolicExpression::IsFirstRow * SymbolicExpression::IsLastRow
+ SymbolicExpression::Constant(F::ONE)
+ SymbolicExpression::IsFirstRow * SymbolicExpression::IsLastRow
+ expr.clone(),
expr.clone() * expr.clone(),
];
let interactions = vec![Interaction {
bus_index: 0,
message: vec![expr.clone(), SymbolicExpression::Constant(F::TWO)],
count: SymbolicExpression::Constant(F::ONE),
count_weight: 1,
}];
let dag = build_symbolic_constraints_dag(&constraints, &interactions);
assert_eq!(
dag.constraints,
SymbolicExpressionDag::<F> {
nodes: vec![
SymbolicExpressionNode::IsFirstRow,
SymbolicExpressionNode::IsLastRow,
SymbolicExpressionNode::Mul {
left_idx: 0,
right_idx: 1,
degree_multiple: 2
},
SymbolicExpressionNode::Constant(F::ONE),
SymbolicExpressionNode::Add {
left_idx: 2,
right_idx: 3,
degree_multiple: 2
},
SymbolicExpressionNode::Add {
left_idx: 4,
right_idx: 2,
degree_multiple: 2
},
SymbolicExpressionNode::Variable(SymbolicVariable::new(
Entry::Main {
part_index: 1,
offset: 2
},
3
)),
SymbolicExpressionNode::Add {
left_idx: 5,
right_idx: 6,
degree_multiple: 2
},
SymbolicExpressionNode::Mul {
left_idx: 6,
right_idx: 6,
degree_multiple: 2
},
SymbolicExpressionNode::Constant(F::TWO),
],
constraint_idx: vec![7, 8],
}
);
assert_eq!(
dag.interactions,
vec![Interaction {
bus_index: 0,
message: vec![6, 9],
count: 3,
count_weight: 1,
}]
);
}
}