1use std::collections::{BTreeMap, BTreeSet};
4
5use sva_formula::{Body, Held, NodeId};
6
7use crate::render::Render;
8use crate::{refs, schedule};
9
10use sva_samples::collapse::plan;
11
12#[derive(Clone, Debug, PartialEq)]
15pub struct Row {
16 pub depth: usize,
17 pub node: String,
18 pub own: u128,
19 pub subtree: u128,
20 pub percent: f64,
21 pub route: &'static str,
22 pub shared: bool,
23}
24
25#[derive(Clone, Debug, PartialEq)]
26pub struct Tree {
27 pub total: u128,
28 pub budget: u128,
29 pub rows: Vec<Row>,
30}
31
32struct Node {
33 id: NodeId,
34 subtree: u128,
35 own: u128,
36 route: &'static str,
37 shared: bool,
38 children: Vec<Node>,
39}
40
41const FOLD_PERCENT: u128 = 1;
43
44const PASS_THROUGH_PERCENT: u128 = 99;
46
47pub fn tree(render: &Render) -> Tree {
48 tree_at(render, render.root)
49}
50
51pub fn total(render: &Render) -> u128 {
54 match render.schedule.materialize.as_slice() {
55 [] => costed(render, render.root).0,
56 held => held.iter().map(|id| costed(render, *id).0).sum(),
57 }
58}
59
60pub fn tree_at(render: &Render, from: NodeId) -> Tree {
61 let mut walked = BTreeMap::from([(from, None)]);
62 let root = grow(render, from, &[], &mut walked);
63 let total = root.subtree;
64 let mut rows = Vec::new();
65 emit(&root, render, 0, total, &mut rows);
66 Tree {
67 total,
68 budget: render.config.flop_budget,
69 rows,
70 }
71}
72
73pub fn dominating(tree: &Tree) -> Option<&Row> {
75 let root = tree.rows.first()?;
76 tree.rows
77 .iter()
78 .skip(1)
79 .filter(|row| row.node != root.node)
80 .max_by_key(|row| row.subtree)
81 .or(Some(root))
82}
83
84fn grow(render: &Render, base: NodeId, chain: &[NodeId], walked: &mut Walked) -> Node {
86 let id = chain.last().copied().unwrap_or(base);
87 let (subtree, route, shared) = priced(render, base, chain, walked);
88 let separate = schedule::materialized_operands(&render.tys, id);
89 let reached: Vec<NodeId> = schedule::read_operands(&render.tys, id)
90 .into_iter()
91 .filter(|child| walked.insert(*child, None).is_none())
92 .collect();
93 let children: Vec<Node> = reached
94 .into_iter()
95 .map(|child| match separate.contains(&child) {
96 true => grow(render, child, &[], walked),
97 false => grow(render, base, &[chain, &[child]].concat(), walked),
98 })
99 .collect();
100 let (inlined, beside): (Vec<&Node>, Vec<&Node>) =
102 children.iter().partition(|c| !separate.contains(&c.id));
103 let inlined: u128 = inlined.iter().map(|c| c.subtree).sum();
104 let beside: u128 = beside.iter().map(|c| c.subtree).sum();
105 let held = Node {
106 id,
107 subtree: subtree + beside,
108 own: subtree.saturating_sub(inlined),
109 route,
110 shared,
111 children,
112 };
113 walked.insert(id, Some(held.subtree));
114 held
115}
116
117type Walked = BTreeMap<NodeId, Option<u128>>;
119
120fn priced(
121 render: &Render,
122 base: NodeId,
123 chain: &[NodeId],
124 walked: &Walked,
125) -> (u128, &'static str, bool) {
126 match chain.last() {
127 None => {
128 let (cost, route) = costed(render, base);
129 (cost, route, false)
130 }
131 Some(id) => read_as(render, base, chain, walked).unwrap_or_else(|| {
132 let mut carried = Carried::beside(walked, *id);
133 let (cost, route) = costed_under(render, *id, &mut carried);
134 (cost, route, carried.shared)
135 }),
136 }
137}
138
139fn read_as(
141 render: &Render,
142 base: NodeId,
143 chain: &[NodeId],
144 walked: &Walked,
145) -> Option<(u128, &'static str, bool)> {
146 let crate::typing::Value::ClosedForm(form) = render.tys.value(base) else {
147 return None;
148 };
149 let len = render.config.horizon.len(render.config.rate).ok()?;
150 let (rate, horizon) = (render.config.rate, render.config.horizon);
151 let profile = &render.config.profile;
152 let body = refs::fold_constants(&render.tys, &along(render, &form.body, chain)?);
153 let (carried, shared) = carried_already(&body, walked, chain);
154 let plan = match refs::spectral_sum_of_body(&render.tys, base, &body, form.var) {
155 Ok(sum) => plan::of(&sum, rate, horizon, profile, len).ok()?,
156 Err(_) => {
157 let written = sva_formula::ClosedForm {
158 var: form.var,
159 body: refs::substituted_body(&render.tys, base, &body)?,
160 origin: form.origin,
161 };
162 plan::of_written(&written, rate, horizon, profile, len).ok()?
163 }
164 };
165 Some((
166 plan.flops(len).saturating_sub(carried),
167 plan.rule().as_str(),
168 shared,
169 ))
170}
171
172fn carried_already(body: &Body, walked: &Walked, chain: &[NodeId]) -> (u128, bool) {
173 if let Body::Node(id) = body
174 && !chain.contains(id)
175 {
176 return match walked.get(id) {
177 Some(Some(paid)) => (*paid, true),
178 _ => (0, false),
179 };
180 }
181 sva_formula::closed_form::children(body)
182 .iter()
183 .fold((0, false), |(sum, held), part| {
184 let (paid, found) = carried_already(&part.body, walked, chain);
185 (sum + paid, held || found)
186 })
187}
188
189fn along(render: &Render, body: &sva_formula::Body, chain: &[NodeId]) -> Option<Body> {
190 let Some((next, rest)) = chain.split_first() else {
191 return Some(body.clone());
192 };
193 match body {
194 Body::Node(id) if id == next => match render.tys.value(*id) {
195 crate::typing::Value::ClosedForm(form) => along(render, &form.body, rest),
196 _ => rest.is_empty().then(|| body.clone()),
197 },
198 Body::Node(_) => Some(Body::Const(sva_formula::C64::ZERO)),
199 other => {
200 let mut held = true;
201 let out = sva_formula::closed_form::map_children(other, |part| {
202 match along(render, &part.body, chain) {
203 Some(body) => sva_formula::Part::new(part.origin, body),
204 None => {
205 held = false;
206 part.clone()
207 }
208 }
209 });
210 held.then_some(out)
211 }
212 }
213}
214
215struct Carried {
217 paid: BTreeSet<NodeId>,
218 shared: bool,
219}
220
221impl Carried {
222 fn of(id: NodeId) -> Self {
223 Carried {
224 paid: BTreeSet::from([id]),
225 shared: false,
226 }
227 }
228
229 fn beside(walked: &Walked, id: NodeId) -> Self {
230 let mut held = Carried::of(id);
231 held.paid.extend(
232 walked
233 .iter()
234 .filter(|(_, row)| row.is_some())
235 .map(|(n, _)| *n),
236 );
237 held.paid.insert(id);
238 held
239 }
240
241 fn opens(&mut self, read: NodeId) -> bool {
242 let fresh = self.paid.insert(read);
243 self.shared |= !fresh;
244 fresh
245 }
246}
247
248fn costed(render: &Render, id: NodeId) -> (u128, &'static str) {
249 costed_under(render, id, &mut Carried::of(id))
250}
251
252fn costed_under(render: &Render, id: NodeId, paid: &mut Carried) -> (u128, &'static str) {
253 let Ok(len) = render.config.horizon.len(render.config.rate) else {
254 return (0, "no horizon");
255 };
256 costed_at(render, id, len, paid)
257}
258
259fn costed_at(render: &Render, id: NodeId, len: usize, paid: &mut Carried) -> (u128, &'static str) {
260 match render.tys.ty(id).held {
261 Held::Frames => (frames_flops(render, id, len), "short-time transform"),
262 Held::Sampled => (ops_of(render, id) as u128 * len as u128, "sampled program"),
263 _ => closed_form_flops(render, id, len, paid),
264 }
265}
266
267fn closed_form_flops(
268 render: &Render,
269 id: NodeId,
270 len: usize,
271 paid: &mut Carried,
272) -> (u128, &'static str) {
273 let var = render.tys.var(id);
274 let (rate, horizon) = (render.config.rate, render.config.horizon);
275 let profile = &render.config.profile;
276 let sum = refs::spectral_sum_of(&render.tys, id, var)
277 .ok()
278 .and_then(|sum| sva_samples::collapse::plan::of(&sum, rate, horizon, profile, len).ok());
279 let plan = sum.or_else(|| {
281 let form = written_closed_form(render, id)?;
282 sva_samples::collapse::plan::of_written(&form, rate, horizon, profile, len).ok()
283 });
284 match plan {
285 Some(plan) => (
286 plan.flops(len) + scored(render, id, &plan, len),
287 plan.rule().as_str(),
288 ),
289 None => (
290 pointwise_flops(render, id, len, paid),
291 sva_samples::Rule::PointSampled.as_str(),
292 ),
293 }
294}
295
296fn pointwise_flops(render: &Render, id: NodeId, len: usize, paid: &mut Carried) -> u128 {
299 let own = match render.tys.value(id) {
300 crate::typing::Value::ClosedForm(form) => plan::point_nodes(&form.body),
301 _ => ops_of(render, id),
302 };
303 schedule::read_operands(&render.tys, id).into_iter().fold(
304 own as u128 * len as u128,
305 |sum, read| match paid.opens(read) {
306 true => sum + costed_at(render, read, len, paid).0,
307 false => sum,
308 },
309 )
310}
311
312pub(crate) fn per_sample(render: &Render, id: NodeId) -> u128 {
315 match render.tys.ty(id).held {
316 Held::Sampled => ops_of(render, id) as u128,
317 _ => pointwise_flops(render, id, 1, &mut Carried::of(id)),
318 }
319}
320
321#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
325pub struct Work {
326 pub samples: u64,
327 pub proofs: u64,
328 pub priced_flops: u128,
329 pub waves: Option<u128>,
330}
331
332fn scored(render: &Render, id: NodeId, plan: &plan::Plan, len: usize) -> u128 {
334 match render.alias_oversample(id) {
335 Some(asked) => plan.alias_flops(len) + plan.alias_flops_at(len, asked as usize),
336 None => 0,
337 }
338}
339
340fn written_closed_form(render: &Render, id: NodeId) -> Option<sva_formula::ClosedForm> {
341 match refs::resolve(&render.tys, id, 0, render.tys.ty(id).held) {
342 Ok(refs::Read::Substitute(form)) => Some(*form),
343 _ => None,
344 }
345}
346
347fn frames_flops(render: &Render, id: NodeId, len: usize) -> u128 {
348 let crate::typing::Value::Cast(crate::cast::Cast::Stft { window, hop }, _) =
349 *render.tys.value(id)
350 else {
351 return 0;
352 };
353 let frames = len.div_ceil(hop.max(1)) as u128;
354 frames * sva_samples::collapse::transform_flops(window.max(1))
355}
356
357fn ops_of(render: &Render, id: NodeId) -> usize {
359 let mut seen = BTreeSet::new();
360 inner_ops(render, id, &mut seen)
361}
362
363fn inner_ops(render: &Render, id: NodeId, seen: &mut BTreeSet<NodeId>) -> usize {
364 if !seen.insert(id) {
365 return 1;
366 }
367 match render.tys.value(id) {
368 crate::typing::Value::Op { args, .. } => {
369 1 + args
370 .iter()
371 .map(|a| inner_ops(render, *a, seen))
372 .sum::<usize>()
373 }
374 crate::typing::Value::Filter { x, .. } => 1 + inner_ops(render, *x, seen),
375 _ => 1,
376 }
377}
378
379fn emit(node: &Node, render: &Render, depth: usize, total: u128, rows: &mut Vec<Row>) {
380 rows.push(Row {
381 depth,
382 node: render.tys.name(node.id).to_string(),
383 own: node.own,
384 subtree: node.subtree,
385 percent: percent(node.subtree, total),
386 route: node.route,
387 shared: node.shared,
388 });
389 children(node, render, depth + 1, total, rows);
390}
391
392fn children(node: &Node, render: &Render, depth: usize, total: u128, rows: &mut Vec<Row>) {
393 let mut held: Vec<&Node> = node.children.iter().collect();
394 held.sort_by_key(|c| std::cmp::Reverse(c.subtree));
395 let (shown, folded): (Vec<&Node>, Vec<&Node>) = held
396 .into_iter()
397 .partition(|c| c.shared || c.subtree * 100 >= total * FOLD_PERCENT);
398 for child in shown {
399 let restates = child.subtree * 100 >= node.subtree * PASS_THROUGH_PERCENT
400 && child.own * 100 < total * FOLD_PERCENT;
401 match restates && !child.children.is_empty() {
402 true => children(child, render, depth, total, rows),
403 false => emit(child, render, depth, total, rows),
404 }
405 }
406 if !folded.is_empty() {
407 let under: u128 = folded.iter().map(|c| c.subtree).sum();
408 rows.push(Row {
409 depth,
410 node: format!("{} others", folded.len()),
411 own: under,
412 subtree: under,
413 percent: percent(under, total),
414 route: "folded",
415 shared: false,
416 });
417 }
418}
419
420fn percent(part: u128, whole: u128) -> f64 {
421 match whole {
422 0 => 0.0,
423 _ => 100.0 * part as f64 / whole as f64,
424 }
425}