use core::sync::atomic::{AtomicUsize, Ordering};
use miden_crypto::{
field::TwoAdicField,
stark::dft::{Radix2DFTSmallBatch, TwoAdicSubgroupDft},
};
use crate::layout::InputKey;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, PartialOrd, Ord)]
pub(crate) struct DagId(usize);
impl DagId {
pub(crate) fn fresh() -> Self {
static NEXT_DAG_ID: AtomicUsize = AtomicUsize::new(0);
Self(NEXT_DAG_ID.fetch_add(1, Ordering::Relaxed))
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, PartialOrd, Ord)]
pub struct NodeId {
pub(super) dag_id: DagId,
pub(super) index: usize,
}
impl NodeId {
pub const fn index(self) -> usize {
self.index
}
pub(super) const fn in_dag(index: usize, dag_id: DagId) -> Self {
Self { dag_id, index }
}
}
#[derive(Debug, Clone, PartialEq, Eq, Hash)]
pub enum NodeKind<EF> {
Input(InputKey),
Constant(EF),
Add(NodeId, NodeId),
Sub(NodeId, NodeId),
Mul(NodeId, NodeId),
Neg(NodeId),
}
#[derive(Debug, Clone)]
pub(crate) struct SparseTerm<EF> {
pub(crate) scaled_value: EF,
pub(crate) twiddles: Vec<EF>,
}
#[derive(Debug, Clone)]
pub(crate) enum PeriodicColumn<EF> {
Dense(Vec<EF>),
Sparse {
period: usize,
terms: Vec<SparseTerm<EF>>,
},
}
impl<EF> PeriodicColumn<EF> {
pub(crate) fn period(&self) -> usize {
match self {
Self::Dense(coeffs) => coeffs.len(),
Self::Sparse { period, .. } => *period,
}
}
}
#[derive(Debug, Clone)]
pub struct PeriodicColumnData<EF> {
columns: Vec<PeriodicColumn<EF>>,
}
impl<EF> PeriodicColumnData<EF> {
pub fn from_periodic_columns<F>(periodic_columns: Vec<Vec<F>>) -> Self
where
F: TwoAdicField,
EF: From<F>,
{
let dft = Radix2DFTSmallBatch::<F>::default();
let mut columns = Vec::with_capacity(periodic_columns.len());
for col in periodic_columns {
assert!(!col.is_empty(), "periodic column must not be empty");
assert!(col.len().is_power_of_two(), "periodic column length must be a power of two");
let period = col.len();
let log_len = period.ilog2() as usize;
let terms = sparse_terms::<F, EF>(&col);
let dense_ops = 2 * period.saturating_sub(1);
let sparse_ops = terms.len() * (3 * log_len) + terms.len().saturating_sub(1);
let column = if terms.is_empty() || sparse_ops < dense_ops {
PeriodicColumn::Sparse { period, terms }
} else {
let coeffs = dft.idft(col).into_iter().map(EF::from).collect();
PeriodicColumn::Dense(coeffs)
};
columns.push(column);
}
Self { columns }
}
pub fn num_columns(&self) -> usize {
self.columns.len()
}
pub fn max_period(&self) -> usize {
self.columns.iter().map(PeriodicColumn::period).max().unwrap_or(0)
}
pub(crate) fn columns(&self) -> &[PeriodicColumn<EF>] {
&self.columns
}
}
fn sparse_terms<F, EF>(col: &[F]) -> Vec<SparseTerm<EF>>
where
F: TwoAdicField,
EF: From<F>,
{
let log_len = col.len().ilog2();
let omega_inv = F::two_adic_generator(log_len as usize).inverse();
let mut domain_size = F::ZERO;
for _ in 0..col.len() {
domain_size += F::ONE;
}
let p_inv = domain_size.inverse();
let mut omega_inv_pow = F::ONE;
let mut terms = Vec::new();
for &v in col {
if v != F::ZERO {
let mut twiddles = Vec::with_capacity(log_len as usize);
let mut base = omega_inv_pow;
for _ in 0..log_len {
twiddles.push(EF::from(base));
base *= base;
}
terms.push(SparseTerm {
scaled_value: EF::from(v * p_inv),
twiddles,
});
}
omega_inv_pow *= omega_inv;
}
terms
}
#[derive(Debug)]
pub struct AceDag<EF> {
dag_id: DagId,
pub nodes: Vec<NodeKind<EF>>,
pub root: NodeId,
}
#[derive(Debug, Clone)]
pub struct DagSnapshot<EF> {
nodes: Vec<NodeKind<EF>>,
root: NodeId,
source_dag_id: DagId,
}
impl<EF> AceDag<EF> {
pub(crate) fn from_parts(dag_id: DagId, nodes: Vec<NodeKind<EF>>, root: NodeId) -> Self {
Self { dag_id, nodes, root }
}
pub(crate) fn nodes(&self) -> &[NodeKind<EF>] {
&self.nodes
}
pub(crate) fn into_nodes(self) -> Vec<NodeKind<EF>> {
self.nodes
}
pub(crate) fn dag_id(&self) -> DagId {
self.dag_id
}
pub fn root(&self) -> NodeId {
self.root
}
pub fn into_snapshot(self) -> DagSnapshot<EF> {
DagSnapshot {
nodes: self.nodes,
root: self.root,
source_dag_id: self.dag_id,
}
}
}
impl<EF: Clone> AceDag<EF> {
pub fn compact(&mut self) {
let n = self.nodes.len();
if n == 0 {
return;
}
let mut reachable = vec![false; n];
let mut stack = vec![self.root.index()];
while let Some(idx) = stack.pop() {
if reachable[idx] {
continue;
}
reachable[idx] = true;
match &self.nodes[idx] {
NodeKind::Add(a, b) | NodeKind::Sub(a, b) | NodeKind::Mul(a, b) => {
stack.push(a.index());
stack.push(b.index());
},
NodeKind::Neg(a) => {
stack.push(a.index());
},
NodeKind::Input(_) | NodeKind::Constant(_) => {},
}
}
let mut remap = vec![0usize; n];
let mut new_len = 0usize;
for i in 0..n {
if reachable[i] {
remap[i] = new_len;
new_len += 1;
}
}
if new_len == n {
return;
}
let dag_id = DagId::fresh();
let remap_id = |id: NodeId| NodeId::in_dag(remap[id.index()], dag_id);
let mut new_nodes = Vec::with_capacity(new_len);
for (i, node) in self.nodes.iter().enumerate() {
if !reachable[i] {
continue;
}
let remapped = match node {
NodeKind::Input(k) => NodeKind::Input(*k),
NodeKind::Constant(v) => NodeKind::Constant(v.clone()),
NodeKind::Add(a, b) => NodeKind::Add(remap_id(*a), remap_id(*b)),
NodeKind::Sub(a, b) => NodeKind::Sub(remap_id(*a), remap_id(*b)),
NodeKind::Mul(a, b) => NodeKind::Mul(remap_id(*a), remap_id(*b)),
NodeKind::Neg(a) => NodeKind::Neg(remap_id(*a)),
};
new_nodes.push(remapped);
}
self.nodes = new_nodes;
self.root = remap_id(self.root);
self.dag_id = dag_id;
}
}
impl<EF> DagSnapshot<EF> {
pub fn root(&self) -> NodeId {
self.root
}
pub(super) fn into_parts(self) -> (DagId, Vec<NodeKind<EF>>, NodeId) {
(self.source_dag_id, self.nodes, self.root)
}
}