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#[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 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#[derive(Debug, Clone, PartialEq, Eq, Hash)]
45pub enum NodeKind<EF> {
46 Input(InputKey),
48 Constant(EF),
50 Add(NodeId, NodeId),
52 Sub(NodeId, NodeId),
54 Mul(NodeId, NodeId),
56 Neg(NodeId),
58}
59
60#[derive(Debug, Clone)]
68pub(crate) struct SparseTerm<EF> {
69 pub(crate) scaled_value: EF,
71 pub(crate) twiddles: Vec<EF>,
73}
74
75#[derive(Debug, Clone)]
80pub(crate) enum PeriodicColumn<EF> {
81 Dense(Vec<EF>),
83 Sparse {
86 period: usize,
87 terms: Vec<SparseTerm<EF>>,
88 },
89}
90
91impl<EF> PeriodicColumn<EF> {
92 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#[derive(Debug, Clone)]
103pub struct PeriodicColumnData<EF> {
104 columns: Vec<PeriodicColumn<EF>>,
106}
107
108impl<EF> PeriodicColumnData<EF> {
109 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 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 pub fn num_columns(&self) -> usize {
151 self.columns.len()
152 }
153
154 pub fn max_period(&self) -> usize {
156 self.columns.iter().map(PeriodicColumn::period).max().unwrap_or(0)
157 }
158
159 pub(crate) fn columns(&self) -> &[PeriodicColumn<EF>] {
161 &self.columns
162 }
163}
164
165fn 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#[derive(Debug)]
209pub struct AceDag<EF> {
210 dag_id: DagId,
211 pub nodes: Vec<NodeKind<EF>>,
213 pub root: NodeId,
215}
216
217#[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 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 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 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}