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::{Radix2DFTSmallBatch, 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 dft = Radix2DFTSmallBatch::<F>::default();
121        let mut columns = Vec::with_capacity(periodic_columns.len());
122        for col in periodic_columns {
123            assert!(!col.is_empty(), "periodic column must not be empty");
124            assert!(col.len().is_power_of_two(), "periodic column length must be a power of two");
125
126            let period = col.len();
127            let log_len = period.ilog2() as usize;
128            let terms = sparse_terms::<F, EF>(&col);
129
130            // Dense Horner costs 2 ops per (nonzero-leading) coefficient. Sparse Lagrange
131            // costs `3 * log_len` ops per nonzero evaluation to build its doubling product,
132            // plus one combining op per term less one shared across the column. Keep
133            // whichever form yields the smaller circuit.
134            let dense_ops = 2 * period.saturating_sub(1);
135            let sparse_ops = terms.len() * (3 * log_len) + terms.len().saturating_sub(1);
136
137            let column = if terms.is_empty() || sparse_ops < dense_ops {
138                PeriodicColumn::Sparse { period, terms }
139            } else {
140                let coeffs = dft.idft(col).into_iter().map(EF::from).collect();
141                PeriodicColumn::Dense(coeffs)
142            };
143            columns.push(column);
144        }
145
146        Self { columns }
147    }
148
149    /// Number of periodic columns.
150    pub fn num_columns(&self) -> usize {
151        self.columns.len()
152    }
153
154    /// Maximum periodic column length (used to align powers).
155    pub fn max_period(&self) -> usize {
156        self.columns.iter().map(PeriodicColumn::period).max().unwrap_or(0)
157    }
158
159    /// Iterate over the per-column chosen representations.
160    pub(crate) fn columns(&self) -> &[PeriodicColumn<EF>] {
161        &self.columns
162    }
163}
164
165/// Build the sparse Lagrange-form terms for one periodic column's nonzero evaluations.
166///
167/// For a column of length `P = 2^m` with evaluation-domain generator `omega`, the
168/// coefficient-form value at a point `x` equals
169/// `(1/P) * sum_j v_j * D(x * omega^(-j))`, where `D(t) = sum_{k=0}^{P-1} t^k`. This
170/// is the same identity underlying the dense IDFT + Horner path, reordered so terms
171/// with `v_j == 0` drop out entirely and `D` is computed division-free via the
172/// doubling product `D(t) = Π_{i=0}^{m-1} (1 + t^(2^i))`.
173fn sparse_terms<F, EF>(col: &[F]) -> Vec<SparseTerm<EF>>
174where
175    F: TwoAdicField,
176    EF: From<F>,
177{
178    let log_len = col.len().ilog2();
179    let omega_inv = F::two_adic_generator(log_len as usize).inverse();
180
181    let mut domain_size = F::ZERO;
182    for _ in 0..col.len() {
183        domain_size += F::ONE;
184    }
185    let p_inv = domain_size.inverse();
186
187    let mut omega_inv_pow = F::ONE;
188    let mut terms = Vec::new();
189    for &v in col {
190        if v != F::ZERO {
191            let mut twiddles = Vec::with_capacity(log_len as usize);
192            let mut base = omega_inv_pow;
193            for _ in 0..log_len {
194                twiddles.push(EF::from(base));
195                base *= base;
196            }
197            terms.push(SparseTerm {
198                scaled_value: EF::from(v * p_inv),
199                twiddles,
200            });
201        }
202        omega_inv_pow *= omega_inv;
203    }
204    terms
205}
206
207/// A built DAG with a designated root.
208#[derive(Debug)]
209pub struct AceDag<EF> {
210    dag_id: DagId,
211    /// Topologically ordered nodes.
212    pub nodes: Vec<NodeKind<EF>>,
213    /// Root node of the verifier equation.
214    pub root: NodeId,
215}
216
217/// Exported DAG data that preserves the source DAG id across imports.
218#[derive(Debug, Clone)]
219pub struct DagSnapshot<EF> {
220    nodes: Vec<NodeKind<EF>>,
221    root: NodeId,
222    source_dag_id: DagId,
223}
224
225impl<EF> AceDag<EF> {
226    pub(crate) fn from_parts(dag_id: DagId, nodes: Vec<NodeKind<EF>>, root: NodeId) -> Self {
227        Self { dag_id, nodes, root }
228    }
229
230    pub(crate) fn nodes(&self) -> &[NodeKind<EF>] {
231        &self.nodes
232    }
233
234    pub(crate) fn into_nodes(self) -> Vec<NodeKind<EF>> {
235        self.nodes
236    }
237
238    pub(crate) fn dag_id(&self) -> DagId {
239        self.dag_id
240    }
241
242    pub fn root(&self) -> NodeId {
243        self.root
244    }
245
246    /// Consume the DAG and return an exported snapshot that can be re-imported later.
247    pub fn into_snapshot(self) -> DagSnapshot<EF> {
248        DagSnapshot {
249            nodes: self.nodes,
250            root: self.root,
251            source_dag_id: self.dag_id,
252        }
253    }
254}
255
256impl<EF: Clone> AceDag<EF> {
257    /// Remove nodes unreachable from `root` and compact the node vector.
258    ///
259    /// After compaction, `nodes` contains only nodes reachable from `root`, in the same
260    /// relative order. All `NodeId` references, including `root`, are remapped to the new
261    /// contiguous indices. When any node is removed, the DAG also takes a fresh `dag_id`,
262    /// so `NodeId`s issued before compaction fail provenance checks instead of silently
263    /// resolving to whichever node now occupies their old index.
264    pub fn compact(&mut self) {
265        let n = self.nodes.len();
266        if n == 0 {
267            return;
268        }
269
270        let mut reachable = vec![false; n];
271        let mut stack = vec![self.root.index()];
272        while let Some(idx) = stack.pop() {
273            if reachable[idx] {
274                continue;
275            }
276            reachable[idx] = true;
277            match &self.nodes[idx] {
278                NodeKind::Add(a, b) | NodeKind::Sub(a, b) | NodeKind::Mul(a, b) => {
279                    stack.push(a.index());
280                    stack.push(b.index());
281                },
282                NodeKind::Neg(a) => {
283                    stack.push(a.index());
284                },
285                NodeKind::Input(_) | NodeKind::Constant(_) => {},
286            }
287        }
288
289        let mut remap = vec![0usize; n];
290        let mut new_len = 0usize;
291        for i in 0..n {
292            if reachable[i] {
293                remap[i] = new_len;
294                new_len += 1;
295            }
296        }
297
298        if new_len == n {
299            return;
300        }
301
302        let dag_id = DagId::fresh();
303        let remap_id = |id: NodeId| NodeId::in_dag(remap[id.index()], dag_id);
304
305        let mut new_nodes = Vec::with_capacity(new_len);
306        for (i, node) in self.nodes.iter().enumerate() {
307            if !reachable[i] {
308                continue;
309            }
310            let remapped = match node {
311                NodeKind::Input(k) => NodeKind::Input(*k),
312                NodeKind::Constant(v) => NodeKind::Constant(v.clone()),
313                NodeKind::Add(a, b) => NodeKind::Add(remap_id(*a), remap_id(*b)),
314                NodeKind::Sub(a, b) => NodeKind::Sub(remap_id(*a), remap_id(*b)),
315                NodeKind::Mul(a, b) => NodeKind::Mul(remap_id(*a), remap_id(*b)),
316                NodeKind::Neg(a) => NodeKind::Neg(remap_id(*a)),
317            };
318            new_nodes.push(remapped);
319        }
320
321        self.nodes = new_nodes;
322        self.root = remap_id(self.root);
323        self.dag_id = dag_id;
324    }
325}
326
327impl<EF> DagSnapshot<EF> {
328    /// Root node of the verifier equation.
329    pub fn root(&self) -> NodeId {
330        self.root
331    }
332
333    pub(super) fn into_parts(self) -> (DagId, Vec<NodeKind<EF>>, NodeId) {
334        (self.source_dag_id, self.nodes, self.root)
335    }
336}