Skip to main content

miden_ace_codegen/dag/
ir.rs

1use core::sync::atomic::{AtomicUsize, Ordering};
2
3use miden_crypto::{
4    field::TwoAdicField,
5    stark::dft::{NaiveDft, TwoAdicSubgroupDft},
6};
7
8use crate::layout::InputKey;
9
10#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, PartialOrd, Ord)]
11pub(crate) struct DagId(usize);
12
13impl DagId {
14    pub(crate) fn fresh() -> Self {
15        static NEXT_DAG_ID: AtomicUsize = AtomicUsize::new(0);
16
17        Self(NEXT_DAG_ID.fetch_add(1, Ordering::Relaxed))
18    }
19}
20
21/// Identifier for a node in the DAG.
22#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, PartialOrd, Ord)]
23pub struct NodeId {
24    pub(super) dag_id: DagId,
25    pub(super) index: usize,
26}
27
28impl NodeId {
29    /// Return the underlying node index.
30    pub const fn index(self) -> usize {
31        self.index
32    }
33
34    pub(super) const fn in_dag(index: usize, dag_id: DagId) -> Self {
35        Self { dag_id, index }
36    }
37}
38
39/// Node kinds in the DAG.
40///
41/// These nodes mirror the verifier expression tree after lowering:
42/// inputs are read via `InputKey`, constants are lifted into the DAG, and
43/// arithmetic nodes capture the evaluation order.
44#[derive(Debug, Clone, PartialEq, Eq, Hash)]
45pub enum NodeKind<EF> {
46    /// Layout-addressable input (public, OOD, aux, etc.).
47    Input(InputKey),
48    /// Constant extension-field value.
49    Constant(EF),
50    /// Addition node.
51    Add(NodeId, NodeId),
52    /// Subtraction node.
53    Sub(NodeId, NodeId),
54    /// Multiplication node.
55    Mul(NodeId, NodeId),
56    /// Negation node (modeled as 0 - x when emitting ops).
57    Neg(NodeId),
58}
59
60/// A nonzero evaluation-domain value of a periodic column, together with the
61/// doubling-basis twiddle powers needed to evaluate its Lagrange contribution
62/// at an arbitrary point `x` via `value * Π_i (1 + twiddle[i] * x^(2^i))`.
63///
64/// This is the sparse dual of the dense monomial-basis coefficients: an IDFT
65/// turns a sparse evaluation vector into dense coefficients, but the Lagrange
66/// form stays sparse in the number of nonzero evaluations.
67#[derive(Debug, Clone)]
68pub(crate) struct SparseTerm<EF> {
69    /// The evaluation-domain value, pre-scaled by the domain-size inverse.
70    pub(crate) scaled_value: EF,
71    /// `omega^(-j * 2^i)` for `i = 0..log2(period)`, where `j` is this term's domain index.
72    pub(crate) twiddles: Vec<EF>,
73}
74
75/// The in-circuit evaluation form chosen for a single periodic column.
76///
77/// The cheaper of the two representations is selected from the column values when
78/// the data is built; the lowering emits nodes for whichever form each column carries.
79#[derive(Debug, Clone)]
80pub(crate) enum PeriodicColumn<EF> {
81    /// Dense monomial-basis coefficients (highest-degree first) for Horner evaluation.
82    Dense(Vec<EF>),
83    /// Sparse Lagrange-form nonzero terms, tagged with the column period, for
84    /// division-free doubling-product evaluation.
85    Sparse {
86        period: usize,
87        terms: Vec<SparseTerm<EF>>,
88    },
89}
90
91impl<EF> PeriodicColumn<EF> {
92    /// The column period (its evaluation-domain length).
93    pub(crate) fn period(&self) -> usize {
94        match self {
95            Self::Dense(coeffs) => coeffs.len(),
96            Self::Sparse { period, .. } => *period,
97        }
98    }
99}
100
101/// Precomputed periodic column data for DAG construction.
102#[derive(Debug, Clone)]
103pub struct PeriodicColumnData<EF> {
104    /// The chosen evaluation form for each periodic column.
105    columns: Vec<PeriodicColumn<EF>>,
106}
107
108impl<EF> PeriodicColumnData<EF> {
109    /// Convert periodic columns (evaluations) into their cheaper in-circuit form.
110    ///
111    /// Each column is lowered to whichever of two representations yields the smaller
112    /// circuit: dense monomial-basis coefficients (via an inverse DFT) evaluated by
113    /// Horner, or a sparse Lagrange form over the column's nonzero evaluations. The
114    /// choice depends only on the column values, so it is fixed at construction.
115    pub fn from_periodic_columns<F>(periodic_columns: Vec<Vec<F>>) -> Self
116    where
117        F: TwoAdicField,
118        EF: From<F>,
119    {
120        let mut columns = Vec::with_capacity(periodic_columns.len());
121        for col in periodic_columns {
122            assert!(!col.is_empty(), "periodic column must not be empty");
123            assert!(col.len().is_power_of_two(), "periodic column length must be a power of two");
124
125            let period = col.len();
126            let log_len = period.ilog2() as usize;
127            let terms = sparse_terms::<F, EF>(&col);
128
129            // Dense Horner costs 2 ops per (nonzero-leading) coefficient. Sparse Lagrange
130            // costs `3 * log_len` ops per nonzero evaluation to build its doubling product,
131            // plus one combining op per term less one shared across the column. Keep
132            // whichever form yields the smaller circuit.
133            let dense_ops = 2 * period.saturating_sub(1);
134            let sparse_ops = terms.len() * (3 * log_len) + terms.len().saturating_sub(1);
135
136            let column = if terms.is_empty() || sparse_ops < dense_ops {
137                PeriodicColumn::Sparse { period, terms }
138            } else {
139                let coeffs = NaiveDft.idft(col).into_iter().map(EF::from).collect();
140                PeriodicColumn::Dense(coeffs)
141            };
142            columns.push(column);
143        }
144
145        Self { columns }
146    }
147
148    /// Number of periodic columns.
149    pub fn num_columns(&self) -> usize {
150        self.columns.len()
151    }
152
153    /// Maximum periodic column length (used to align powers).
154    pub fn max_period(&self) -> usize {
155        self.columns.iter().map(PeriodicColumn::period).max().unwrap_or(0)
156    }
157
158    /// Iterate over the per-column chosen representations.
159    pub(crate) fn columns(&self) -> &[PeriodicColumn<EF>] {
160        &self.columns
161    }
162}
163
164/// Build the sparse Lagrange-form terms for one periodic column's nonzero evaluations.
165///
166/// For a column of length `P = 2^m` with evaluation-domain generator `omega`, the
167/// coefficient-form value at a point `x` equals
168/// `(1/P) * sum_j v_j * D(x * omega^(-j))`, where `D(t) = sum_{k=0}^{P-1} t^k`. This
169/// is the same identity underlying the dense IDFT + Horner path, reordered so terms
170/// with `v_j == 0` drop out entirely and `D` is computed division-free via the
171/// doubling product `D(t) = Π_{i=0}^{m-1} (1 + t^(2^i))`.
172fn sparse_terms<F, EF>(col: &[F]) -> Vec<SparseTerm<EF>>
173where
174    F: TwoAdicField,
175    EF: From<F>,
176{
177    let log_len = col.len().ilog2();
178    let omega_inv = F::two_adic_generator(log_len as usize).inverse();
179
180    let mut domain_size = F::ZERO;
181    for _ in 0..col.len() {
182        domain_size += F::ONE;
183    }
184    let p_inv = domain_size.inverse();
185
186    let mut omega_inv_pow = F::ONE;
187    let mut terms = Vec::new();
188    for &v in col {
189        if v != F::ZERO {
190            let mut twiddles = Vec::with_capacity(log_len as usize);
191            let mut base = omega_inv_pow;
192            for _ in 0..log_len {
193                twiddles.push(EF::from(base));
194                base *= base;
195            }
196            terms.push(SparseTerm {
197                scaled_value: EF::from(v * p_inv),
198                twiddles,
199            });
200        }
201        omega_inv_pow *= omega_inv;
202    }
203    terms
204}
205
206/// A built DAG with a designated root.
207#[derive(Debug)]
208pub struct AceDag<EF> {
209    dag_id: DagId,
210    /// Topologically ordered nodes.
211    pub nodes: Vec<NodeKind<EF>>,
212    /// Root node of the verifier equation.
213    pub root: NodeId,
214}
215
216/// Exported DAG data that preserves the source DAG id across imports.
217#[derive(Debug, Clone)]
218pub struct DagSnapshot<EF> {
219    nodes: Vec<NodeKind<EF>>,
220    root: NodeId,
221    source_dag_id: DagId,
222}
223
224impl<EF> AceDag<EF> {
225    pub(crate) fn from_parts(dag_id: DagId, nodes: Vec<NodeKind<EF>>, root: NodeId) -> Self {
226        Self { dag_id, nodes, root }
227    }
228
229    pub(crate) fn nodes(&self) -> &[NodeKind<EF>] {
230        &self.nodes
231    }
232
233    pub(crate) fn into_nodes(self) -> Vec<NodeKind<EF>> {
234        self.nodes
235    }
236
237    pub(crate) fn dag_id(&self) -> DagId {
238        self.dag_id
239    }
240
241    pub fn root(&self) -> NodeId {
242        self.root
243    }
244
245    /// Consume the DAG and return an exported snapshot that can be re-imported later.
246    pub fn into_snapshot(self) -> DagSnapshot<EF> {
247        DagSnapshot {
248            nodes: self.nodes,
249            root: self.root,
250            source_dag_id: self.dag_id,
251        }
252    }
253}
254
255impl<EF: Clone> AceDag<EF> {
256    /// Remove nodes unreachable from `root` and compact the node vector.
257    ///
258    /// After compaction, `nodes` contains only nodes reachable from `root`, in the same
259    /// relative order. All `NodeId` references, including `root`, are remapped to the new
260    /// contiguous indices. When any node is removed, the DAG also takes a fresh `dag_id`,
261    /// so `NodeId`s issued before compaction fail provenance checks instead of silently
262    /// resolving to whichever node now occupies their old index.
263    pub fn compact(&mut self) {
264        let n = self.nodes.len();
265        if n == 0 {
266            return;
267        }
268
269        let mut reachable = vec![false; n];
270        let mut stack = vec![self.root.index()];
271        while let Some(idx) = stack.pop() {
272            if reachable[idx] {
273                continue;
274            }
275            reachable[idx] = true;
276            match &self.nodes[idx] {
277                NodeKind::Add(a, b) | NodeKind::Sub(a, b) | NodeKind::Mul(a, b) => {
278                    stack.push(a.index());
279                    stack.push(b.index());
280                },
281                NodeKind::Neg(a) => {
282                    stack.push(a.index());
283                },
284                NodeKind::Input(_) | NodeKind::Constant(_) => {},
285            }
286        }
287
288        let mut remap = vec![0usize; n];
289        let mut new_len = 0usize;
290        for i in 0..n {
291            if reachable[i] {
292                remap[i] = new_len;
293                new_len += 1;
294            }
295        }
296
297        if new_len == n {
298            return;
299        }
300
301        let dag_id = DagId::fresh();
302        let remap_id = |id: NodeId| NodeId::in_dag(remap[id.index()], dag_id);
303
304        let mut new_nodes = Vec::with_capacity(new_len);
305        for (i, node) in self.nodes.iter().enumerate() {
306            if !reachable[i] {
307                continue;
308            }
309            let remapped = match node {
310                NodeKind::Input(k) => NodeKind::Input(*k),
311                NodeKind::Constant(v) => NodeKind::Constant(v.clone()),
312                NodeKind::Add(a, b) => NodeKind::Add(remap_id(*a), remap_id(*b)),
313                NodeKind::Sub(a, b) => NodeKind::Sub(remap_id(*a), remap_id(*b)),
314                NodeKind::Mul(a, b) => NodeKind::Mul(remap_id(*a), remap_id(*b)),
315                NodeKind::Neg(a) => NodeKind::Neg(remap_id(*a)),
316            };
317            new_nodes.push(remapped);
318        }
319
320        self.nodes = new_nodes;
321        self.root = remap_id(self.root);
322        self.dag_id = dag_id;
323    }
324}
325
326impl<EF> DagSnapshot<EF> {
327    /// Root node of the verifier equation.
328    pub fn root(&self) -> NodeId {
329        self.root
330    }
331
332    pub(super) fn into_parts(self) -> (DagId, Vec<NodeKind<EF>>, NodeId) {
333        (self.source_dag_id, self.nodes, self.root)
334    }
335}