Skip to main content

miden_ace_codegen/
factored.rs

1//! Factored ACE emission: a per-proof-order shuffle section composed with an
2//! order-invariant common section.
3//!
4//! The multi-AIR circuit depends on the proof order only through its inputs: which READ
5//! slot feeds which canonical wire, and which power of the fold challenge multiplies which
6//! per-AIR accumulator. Factored emission lowers the canonical DAG once into a common
7//! operation section whose encoding is byte-identical for every proof order, and prefixes
8//! it with a short per-order shuffle section that routes proof-order READ slots (and fold
9//! coefficients) to the canonical wires the common section consumes.
10//!
11//! Shuffle-section op positions are fixed across orders (only operand ids vary):
12//! 1. one copy gate `Add(input[src], 0)` per shuffled READ slot, in canonical slot order;
13//! 2. fold-challenge powers `beta^2 .. beta^(n-1)` by chained `Mul` (order-invariant);
14//! 3. one fold-coefficient gate `Add(power(e_j), 0)` per AIR `j` in canonical order;
15//! 4. `Add(0, 0)` padding up to an `adv_pipe` block boundary.
16//!
17//! Because the constants section is also padded to a block boundary, the encoded stream
18//! splits into two block-aligned segments: `[constants | shuffle ops]` (per-order) and
19//! `[common ops | root padding]` (order-invariant), which the MASM loader hashes
20//! separately and the registry binds as `merge(H(prefix_i), H(common))`.
21
22use std::collections::HashMap;
23
24use miden_core::Felt;
25use miden_crypto::field::Field;
26
27use crate::{
28    AceError, EXT_DEGREE, InputLayout,
29    circuit::{AceCircuit, AceNode, AceOp, AceOpNode},
30    dag::{AceDag, NodeKind},
31    encode::{ADV_PIPE_BLOCK_FELTS, CONST_EF_ALIGN, StreamGeometry},
32    layout::InputKey,
33};
34
35/// Constants are EF-encoded, so a block boundary is this many EF nodes.
36const CONST_EF_BLOCK_ALIGN: usize = ADV_PIPE_BLOCK_FELTS / EXT_DEGREE;
37
38// The factored scheme hashes `[constants | shuffle]` and `[common | padding]` as separate
39// adv_pipe-aligned segments. That split only lands on a block boundary if the encoder's
40// READ-row rounding (`CONST_EF_ALIGN`) is a no-op on the block-padded constants; otherwise
41// `to_ace` would insert extra constant padding and shift the segment boundary mid-block.
42const _: () = assert!(
43    CONST_EF_BLOCK_ALIGN.is_multiple_of(CONST_EF_ALIGN),
44    "constant block alignment must refine the encoder's READ-row alignment, or the two-segment split drifts off a block boundary"
45);
46
47/// Index of the seeded zero constant used by copy gates and padding.
48const CONST_ZERO: usize = 0;
49/// Index of the seeded one constant used for the zero-exponent fold coefficient.
50const CONST_ONE: usize = 1;
51
52/// Reusable scratch for `FactoredMultiAirCircuit::encode_shuffle_section_for_order`.
53///
54/// Holds the per-order sources, exponents, validation marks, operations, and encoded felts across
55/// calls so a caller enumerating many orderings does not regrow them each time.
56#[derive(Clone, Debug, Default)]
57pub struct ShuffleEncodeBuffer {
58    srcs: Vec<usize>,
59    exponents: Vec<usize>,
60    seen_srcs: Vec<bool>,
61    seen_exponents: Vec<bool>,
62    ops: Vec<AceOpNode>,
63    felts: Vec<Felt>,
64}
65
66impl ShuffleEncodeBuffer {
67    /// Create an empty buffer.
68    pub fn new() -> Self {
69        Self::default()
70    }
71
72    /// Borrow the per-order shuffle-source and fold-exponent scratch.
73    pub(crate) fn order_scratch(&mut self) -> (&mut Vec<usize>, &mut Vec<usize>) {
74        (&mut self.srcs, &mut self.exponents)
75    }
76}
77
78/// Multi-AIR ACE circuit factored into a per-order shuffle section and a common section.
79#[derive(Debug, Clone)]
80pub struct FactoredAceCircuit<EF> {
81    layout: InputLayout,
82    /// Seeded `[0, 1]` followed by the canonical DAG constants, padded to a block boundary.
83    constants: Vec<EF>,
84    /// Canonical (destination) global input index of each shuffle copy gate.
85    shuffle_dsts: Vec<usize>,
86    /// Membership map for [`Self::shuffle_dsts`], indexed by global input index.
87    shuffle_dst_mask: Vec<bool>,
88    /// Number of fold coefficients (one per AIR).
89    num_fold_coeffs: usize,
90    /// Total shuffle-section ops: copies + power muls + coefficient gates + padding.
91    num_shuffle_ops: usize,
92    /// Common-section ops; operands reference absolute node positions.
93    common_ops: Vec<AceOpNode>,
94    /// Node-id bases of the assembled stream; identical for every proof order.
95    geometry: StreamGeometry,
96}
97
98impl<EF: Field> FactoredAceCircuit<EF> {
99    /// Return the input layout shared by every assembled circuit.
100    pub fn layout(&self) -> &InputLayout {
101        &self.layout
102    }
103
104    /// Number of shuffle-section ops (also the section length in stream felts).
105    pub fn num_shuffle_ops(&self) -> usize {
106        self.num_shuffle_ops
107    }
108
109    /// Emit the shuffle-section operations for one proof order, appending to `out`.
110    ///
111    /// Shared by [`Self::assemble`] and [`Self::encode_shuffle_section`] so the assembled
112    /// circuit and the encode-only registry path cannot drift apart.
113    ///
114    /// `AceNode::Operation` operands are absolute indices into the finished operation list,
115    /// so the power-chain base is taken relative to `out`'s current length rather than
116    /// assuming this is the first emission into it.
117    fn emit_shuffle_ops(
118        &self,
119        shuffle_srcs: &[usize],
120        coeff_exponents: &[usize],
121        beta: Option<usize>,
122        out: &mut Vec<AceOpNode>,
123    ) {
124        let start = out.len();
125        let zero = AceNode::Constant(CONST_ZERO);
126        let beta_node =
127            || AceNode::Input(beta.expect("fold challenge is required beyond a single fold slot"));
128        let powers_start = start + self.shuffle_dsts.len();
129        let power_node = |e: usize| match e {
130            0 => AceNode::Constant(CONST_ONE),
131            1 => beta_node(),
132            _ => AceNode::Operation(powers_start + (e - 2)),
133        };
134
135        for &src in shuffle_srcs {
136            out.push(AceOpNode {
137                op: AceOp::Add,
138                lhs: AceNode::Input(src),
139                rhs: zero,
140            });
141        }
142        for e in 2..self.num_fold_coeffs {
143            out.push(AceOpNode {
144                op: AceOp::Mul,
145                lhs: power_node(e - 1),
146                rhs: beta_node(),
147            });
148        }
149        for &e in coeff_exponents {
150            out.push(AceOpNode {
151                op: AceOp::Add,
152                lhs: power_node(e),
153                rhs: zero,
154            });
155        }
156        debug_assert!(
157            out.len() - start <= self.num_shuffle_ops,
158            "shuffle emission overran its section and would displace the common ops"
159        );
160        while out.len() - start < self.num_shuffle_ops {
161            out.push(AceOpNode { op: AceOp::Add, lhs: zero, rhs: zero });
162        }
163    }
164
165    /// Encode the shuffle section for the sources and exponents already staged in `buffer`
166    /// (see [`ShuffleEncodeBuffer::order_scratch`]), reusing its allocations.
167    ///
168    /// Equivalent to taking the shuffle slice of `assemble(..).to_ace()`'s instruction
169    /// stream, without building the circuit or encoding the order-invariant common section.
170    /// Registry construction visits every proof order, so it pays only the per-order bytes.
171    ///
172    /// Rejects the same layouts [`AceCircuit::to_ace`] does. This path never builds an
173    /// `AceCircuit`, so without that check it would return felts for a stream that the
174    /// encoder — and therefore the chiplet — would refuse, and a registry built over it
175    /// would commit to circuits that can never be evaluated.
176    pub(crate) fn encode_shuffle_section<'a>(
177        &self,
178        buffer: &'a mut ShuffleEncodeBuffer,
179    ) -> Result<&'a [Felt], AceError> {
180        self.geometry.validate()?;
181
182        let beta = self.validate_assembly_with_scratch(
183            &buffer.srcs,
184            &buffer.exponents,
185            &mut buffer.seen_srcs,
186            &mut buffer.seen_exponents,
187        )?;
188
189        let mut ops = core::mem::take(&mut buffer.ops);
190        ops.clear();
191        self.emit_shuffle_ops(&buffer.srcs, &buffer.exponents, beta, &mut ops);
192        buffer.ops = ops;
193
194        buffer.felts.clear();
195        buffer.felts.reserve(buffer.ops.len());
196        for op in &buffer.ops {
197            buffer.felts.push(self.geometry.encode_operation(op)?);
198        }
199        Ok(&buffer.felts)
200    }
201
202    /// Shared precondition check for both the assembly and encode-only paths.
203    ///
204    /// Checks that `shuffle_srcs` is a permutation of the destination slots and that
205    /// `coeff_exponents` is a permutation of `0..num_fold_coeffs`; a repeated exponent
206    /// would silently produce a degenerate fold rather than an error.
207    ///
208    /// Returns the fold-challenge input slot, which is absent only for a single-slot fold.
209    fn validate_assembly(
210        &self,
211        shuffle_srcs: &[usize],
212        coeff_exponents: &[usize],
213    ) -> Result<Option<usize>, AceError> {
214        let mut seen_srcs = Vec::new();
215        let mut seen_exponents = Vec::new();
216        self.validate_assembly_with_scratch(
217            shuffle_srcs,
218            coeff_exponents,
219            &mut seen_srcs,
220            &mut seen_exponents,
221        )
222    }
223
224    fn validate_assembly_with_scratch(
225        &self,
226        shuffle_srcs: &[usize],
227        coeff_exponents: &[usize],
228        seen_srcs: &mut Vec<bool>,
229        seen_exponents: &mut Vec<bool>,
230    ) -> Result<Option<usize>, AceError> {
231        if shuffle_srcs.len() != self.shuffle_dsts.len() {
232            return Err(AceError::InvalidInputLayout {
233                message: format!(
234                    "shuffle source count ({}) does not match destination count ({})",
235                    shuffle_srcs.len(),
236                    self.shuffle_dsts.len()
237                ),
238            });
239        }
240        if !is_exact_permutation(
241            shuffle_srcs,
242            self.shuffle_dsts.len(),
243            &self.shuffle_dst_mask,
244            seen_srcs,
245        ) {
246            return Err(AceError::InvalidInputLayout {
247                message: "shuffle sources must be a permutation of the shuffled slots".into(),
248            });
249        }
250        if coeff_exponents.len() != self.num_fold_coeffs {
251            return Err(AceError::InvalidInputLayout {
252                message: format!(
253                    "fold coefficient count ({}) does not match AIR count ({})",
254                    coeff_exponents.len(),
255                    self.num_fold_coeffs
256                ),
257            });
258        }
259        // The fold assigns each slot a distinct challenge power, so the exponents must be a
260        // permutation. A repeated exponent still encodes into a well-formed circuit, so this is
261        // the only place it can be caught.
262        seen_exponents.resize(self.num_fold_coeffs, false);
263        seen_exponents.fill(false);
264        for &exponent in coeff_exponents {
265            let seen =
266                seen_exponents.get_mut(exponent).ok_or_else(|| AceError::InvalidInputLayout {
267                    message: format!("fold coefficient exponent {exponent} out of range"),
268                })?;
269            if *seen {
270                return Err(AceError::InvalidInputLayout {
271                    message: format!("fold coefficient exponent {exponent} is used twice"),
272                });
273            }
274            *seen = true;
275        }
276
277        // A single-slot fold only ever uses the exponent 0, so the challenge itself is not
278        // referenced and need not be present in the layout.
279        let beta = match self.layout.index(InputKey::MultiAirFoldBeta) {
280            Some(beta) => Some(beta),
281            None if self.num_fold_coeffs == 1 => None,
282            None => {
283                return Err(AceError::InvalidInputLayout {
284                    message: "factored circuit requires a MultiAirFoldBeta input slot".into(),
285                });
286            },
287        };
288        Ok(beta)
289    }
290
291    /// Assemble the full circuit for one proof order.
292    ///
293    /// `shuffle_srcs[i]` is the proof-order (source) global input index feeding the `i`-th
294    /// shuffle copy gate; it must be a permutation of the destination slots.
295    /// `coeff_exponents[j]` is the fold-challenge exponent assigned to canonical AIR `j`.
296    pub fn assemble(
297        &self,
298        shuffle_srcs: &[usize],
299        coeff_exponents: &[usize],
300    ) -> Result<AceCircuit<EF>, AceError> {
301        let beta = self.validate_assembly(shuffle_srcs, coeff_exponents)?;
302
303        let mut operations = Vec::with_capacity(self.num_shuffle_ops + self.common_ops.len());
304        self.emit_shuffle_ops(shuffle_srcs, coeff_exponents, beta, &mut operations);
305        operations.extend_from_slice(&self.common_ops);
306
307        let root = AceNode::Operation(operations.len() - 1);
308        Ok(AceCircuit {
309            layout: self.layout.clone(),
310            constants: self.constants.clone(),
311            operations,
312            root,
313        })
314    }
315}
316
317/// Return whether `values` has `expected_len` distinct entries admitted by `membership`.
318fn is_exact_permutation(
319    values: &[usize],
320    expected_len: usize,
321    membership: &[bool],
322    seen: &mut Vec<bool>,
323) -> bool {
324    if values.len() != expected_len {
325        return false;
326    }
327    seen.resize(membership.len(), false);
328    seen.fill(false);
329    values.iter().all(|&value| {
330        let Some(true) = membership.get(value).copied() else {
331            return false;
332        };
333        !core::mem::replace(&mut seen[value], true)
334    })
335}
336
337/// Lower a canonical multi-AIR DAG into a factored circuit.
338///
339/// `shuffle_dsts` enumerates the canonical global input indices of every shuffled READ
340/// slot (the order fixes the copy-gate order shared by all proof orders). The DAG may
341/// reference shuffled slots only through those destinations, and fold coefficients only
342/// through [`InputKey::MultiAirFoldCoeff`] with index below `num_fold_coeffs`.
343pub fn emit_factored_circuit<EF>(
344    dag: &AceDag<EF>,
345    layout: InputLayout,
346    shuffle_dsts: Vec<usize>,
347    num_fold_coeffs: usize,
348) -> Result<FactoredAceCircuit<EF>, AceError>
349where
350    EF: Field,
351{
352    layout.validate();
353    if num_fold_coeffs == 0 {
354        return Err(AceError::InvalidInputLayout {
355            message: "factored circuit requires at least one fold coefficient".into(),
356        });
357    }
358
359    let mut copy_by_dst = HashMap::with_capacity(shuffle_dsts.len());
360    let mut shuffle_dst_mask = vec![false; layout.total_inputs];
361    for (copy_idx, &dst) in shuffle_dsts.iter().enumerate() {
362        if dst >= layout.total_inputs {
363            return Err(AceError::InvalidInputLayout {
364                message: format!("shuffle destination {dst} is outside the READ layout"),
365            });
366        }
367        if copy_by_dst.insert(dst, copy_idx).is_some() {
368            return Err(AceError::InvalidInputLayout {
369                message: format!("duplicate shuffle destination {dst}"),
370            });
371        }
372        shuffle_dst_mask[dst] = true;
373    }
374
375    let num_copies = shuffle_dsts.len();
376    let num_power_ops = num_fold_coeffs.saturating_sub(2);
377    let unpadded = num_copies + num_power_ops + num_fold_coeffs;
378    let num_shuffle_ops = unpadded.next_multiple_of(ADV_PIPE_BLOCK_FELTS);
379    let coeffs_start = num_copies + num_power_ops;
380
381    let mut constants = vec![EF::ZERO, EF::ONE];
382    let mut constant_map = HashMap::<EF, usize>::new();
383    constant_map.insert(EF::ZERO, CONST_ZERO);
384    constant_map.insert(EF::ONE, CONST_ONE);
385
386    let mut common_ops: Vec<AceOpNode> = Vec::new();
387    let mut node_map: Vec<Option<AceNode>> = vec![None; dag.nodes().len()];
388
389    let lookup = |map: &[Option<AceNode>], id: crate::dag::NodeId| -> AceNode {
390        map[id.index()].expect("ACE DAG nodes must be topologically ordered")
391    };
392
393    for (idx, node) in dag.nodes().iter().enumerate() {
394        let ace_node = match node {
395            NodeKind::Input(InputKey::MultiAirFoldCoeff(air)) => {
396                if *air >= num_fold_coeffs {
397                    return Err(AceError::InvalidInputLayout {
398                        message: format!("fold coefficient index {air} out of range"),
399                    });
400                }
401                AceNode::Operation(coeffs_start + air)
402            },
403            NodeKind::Input(key) => {
404                let input_idx = layout.index(*key).ok_or_else(|| AceError::InvalidInputLayout {
405                    message: format!("missing input key in layout: {key:?}"),
406                })?;
407                match copy_by_dst.get(&input_idx) {
408                    Some(&copy_idx) => AceNode::Operation(copy_idx),
409                    // Reading a slot directly is only correct for keys whose position does not
410                    // depend on the proof order. This match is exhaustive on purpose: a new
411                    // per-AIR input kind must fail to compile here rather than silently wire a
412                    // proof-order slot into the canonical section.
413                    None => match *key {
414                        InputKey::Public(_)
415                        | InputKey::AuxRandAlpha
416                        | InputKey::AuxRandBeta
417                        | InputKey::MultiAirFoldBeta
418                        | InputKey::Reserved
419                        | InputKey::Alpha
420                        | InputKey::ZPowN
421                        | InputKey::ZK
422                        | InputKey::IsFirst
423                        | InputKey::IsLast
424                        | InputKey::IsTransition
425                        | InputKey::IsFirstAir(_)
426                        | InputKey::IsLastAir(_)
427                        | InputKey::IsTransitionAir(_)
428                        | InputKey::Weight0
429                        | InputKey::F
430                        | InputKey::S0
431                        | InputKey::QuotientChunkCoord { .. } => AceNode::Input(input_idx),
432                        // Per-AIR regions are laid out in proof order, so they must be reached
433                        // through the shuffle. Preprocessed traces are committed per AIR in the
434                        // same height-sorted order as main traces (see lifted-stark's
435                        // `preprocessed_air_for_trace_index`), so they shuffle identically.
436                        InputKey::Preprocessed { .. }
437                        | InputKey::Main { .. }
438                        | InputKey::AuxCoord { .. }
439                        | InputKey::AuxBusBoundary(_) => {
440                            return Err(AceError::InvalidInputLayout {
441                                message: format!(
442                                    "shuffled input key {key:?} has no shuffle destination"
443                                ),
444                            });
445                        },
446                        // Resolved above, before the layout lookup.
447                        InputKey::MultiAirFoldCoeff(_) => unreachable!(),
448                    },
449                }
450            },
451            NodeKind::Constant(value) => {
452                let const_idx = *constant_map.entry(*value).or_insert_with(|| {
453                    constants.push(*value);
454                    constants.len() - 1
455                });
456                AceNode::Constant(const_idx)
457            },
458            NodeKind::Add(a, b) => {
459                let (lhs, rhs) = (lookup(&node_map, *a), lookup(&node_map, *b));
460                common_ops.push(AceOpNode { op: AceOp::Add, lhs, rhs });
461                AceNode::Operation(num_shuffle_ops + common_ops.len() - 1)
462            },
463            NodeKind::Sub(a, b) => {
464                let (lhs, rhs) = (lookup(&node_map, *a), lookup(&node_map, *b));
465                common_ops.push(AceOpNode { op: AceOp::Sub, lhs, rhs });
466                AceNode::Operation(num_shuffle_ops + common_ops.len() - 1)
467            },
468            NodeKind::Mul(a, b) => {
469                let (lhs, rhs) = (lookup(&node_map, *a), lookup(&node_map, *b));
470                common_ops.push(AceOpNode { op: AceOp::Mul, lhs, rhs });
471                AceNode::Operation(num_shuffle_ops + common_ops.len() - 1)
472            },
473            NodeKind::Neg(a) => {
474                let rhs = lookup(&node_map, *a);
475                common_ops.push(AceOpNode {
476                    op: AceOp::Sub,
477                    lhs: AceNode::Constant(CONST_ZERO),
478                    rhs,
479                });
480                AceNode::Operation(num_shuffle_ops + common_ops.len() - 1)
481            },
482        };
483        node_map[idx] = Some(ace_node);
484    }
485
486    match lookup(&node_map, dag.root()) {
487        AceNode::Operation(idx) if idx == num_shuffle_ops + common_ops.len() - 1 => {},
488        other => {
489            return Err(AceError::InvalidInputLayout {
490                message: format!("factored DAG root must be the last common op, got {other:?}"),
491            });
492        },
493    }
494
495    // Block-align the constants so both encoded segments start on adv_pipe boundaries.
496    let padded_len = constants.len().next_multiple_of(CONST_EF_BLOCK_ALIGN);
497    constants.resize(padded_len, EF::ZERO);
498
499    // The assembled stream has the same node counts for every proof order, so its node-id
500    // bases can be fixed here. `from_counts` applies the encoder's padding rules, so this
501    // geometry is the one `to_ace` derives for every circuit assembled from this factoring.
502    let num_ops = num_shuffle_ops + common_ops.len();
503    let geometry = StreamGeometry::from_counts(layout.total_inputs, constants.len(), num_ops);
504
505    Ok(FactoredAceCircuit {
506        layout,
507        constants,
508        shuffle_dsts,
509        shuffle_dst_mask,
510        num_fold_coeffs,
511        num_shuffle_ops,
512        common_ops,
513        geometry,
514    })
515}
516
517#[cfg(test)]
518mod tests {
519    use super::is_exact_permutation;
520
521    #[test]
522    fn exact_permutation_rejects_missing_duplicate_and_foreign_values() {
523        let membership = [false, true, false, true, true];
524        let mut seen = Vec::new();
525
526        assert!(is_exact_permutation(&[4, 1, 3], 3, &membership, &mut seen));
527        assert!(!is_exact_permutation(&[1, 3], 3, &membership, &mut seen));
528        assert!(!is_exact_permutation(&[1, 1, 4], 3, &membership, &mut seen));
529        assert!(!is_exact_permutation(&[1, 2, 4], 3, &membership, &mut seen));
530        assert!(!is_exact_permutation(&[1, 3, 5], 3, &membership, &mut seen));
531    }
532}