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 match render.tys.ty(id).held {
257 Held::Frames => (frames_flops(render, id, len), "short-time transform"),
258 Held::Sampled => (ops_of(render, id) as u128 * len as u128, "sampled program"),
259 _ => closed_form_flops(render, id, len, paid),
260 }
261}
262
263fn closed_form_flops(
264 render: &Render,
265 id: NodeId,
266 len: usize,
267 paid: &mut Carried,
268) -> (u128, &'static str) {
269 let var = render.tys.var(id);
270 let (rate, horizon) = (render.config.rate, render.config.horizon);
271 let profile = &render.config.profile;
272 let sum = refs::spectral_sum_of(&render.tys, id, var)
273 .ok()
274 .and_then(|sum| sva_samples::collapse::plan::of(&sum, rate, horizon, profile, len).ok());
275 let plan = sum.or_else(|| {
277 let form = written_closed_form(render, id)?;
278 sva_samples::collapse::plan::of_written(&form, rate, horizon, profile, len).ok()
279 });
280 match plan {
281 Some(plan) => (
282 plan.flops(len) + scored(render, id, &plan, len),
283 plan.rule().as_str(),
284 ),
285 None => (
286 pointwise_flops(render, id, len, paid),
287 sva_samples::Rule::PointSampled.as_str(),
288 ),
289 }
290}
291
292fn pointwise_flops(render: &Render, id: NodeId, len: usize, paid: &mut Carried) -> u128 {
295 let own = match render.tys.value(id) {
296 crate::typing::Value::ClosedForm(form) => plan::point_nodes(&form.body),
297 _ => ops_of(render, id),
298 };
299 schedule::read_operands(&render.tys, id).into_iter().fold(
300 own as u128 * len as u128,
301 |sum, read| match paid.opens(read) {
302 true => sum + costed_under(render, read, paid).0,
303 false => sum,
304 },
305 )
306}
307
308fn scored(render: &Render, id: NodeId, plan: &plan::Plan, len: usize) -> u128 {
310 match render.alias_oversample(id) {
311 Some(asked) => plan.alias_flops(len) + plan.alias_flops_at(len, asked as usize),
312 None => 0,
313 }
314}
315
316fn written_closed_form(render: &Render, id: NodeId) -> Option<sva_formula::ClosedForm> {
317 match refs::resolve(&render.tys, id, 0, render.tys.ty(id).held) {
318 Ok(refs::Read::Substitute(form)) => Some(*form),
319 _ => None,
320 }
321}
322
323fn frames_flops(render: &Render, id: NodeId, len: usize) -> u128 {
324 let crate::typing::Value::Cast(crate::cast::Cast::Stft { window, hop }, _) =
325 *render.tys.value(id)
326 else {
327 return 0;
328 };
329 let frames = len.div_ceil(hop.max(1)) as u128;
330 frames * sva_samples::collapse::transform_flops(window.max(1))
331}
332
333fn ops_of(render: &Render, id: NodeId) -> usize {
335 let mut seen = BTreeSet::new();
336 inner_ops(render, id, &mut seen)
337}
338
339fn inner_ops(render: &Render, id: NodeId, seen: &mut BTreeSet<NodeId>) -> usize {
340 if !seen.insert(id) {
341 return 1;
342 }
343 match render.tys.value(id) {
344 crate::typing::Value::Op { args, .. } => {
345 1 + args
346 .iter()
347 .map(|a| inner_ops(render, *a, seen))
348 .sum::<usize>()
349 }
350 crate::typing::Value::Filter { x, .. } => 1 + inner_ops(render, *x, seen),
351 _ => 1,
352 }
353}
354
355fn emit(node: &Node, render: &Render, depth: usize, total: u128, rows: &mut Vec<Row>) {
356 rows.push(Row {
357 depth,
358 node: render.tys.name(node.id).to_string(),
359 own: node.own,
360 subtree: node.subtree,
361 percent: percent(node.subtree, total),
362 route: node.route,
363 shared: node.shared,
364 });
365 children(node, render, depth + 1, total, rows);
366}
367
368fn children(node: &Node, render: &Render, depth: usize, total: u128, rows: &mut Vec<Row>) {
369 let mut held: Vec<&Node> = node.children.iter().collect();
370 held.sort_by_key(|c| std::cmp::Reverse(c.subtree));
371 let (shown, folded): (Vec<&Node>, Vec<&Node>) = held
372 .into_iter()
373 .partition(|c| c.shared || c.subtree * 100 >= total * FOLD_PERCENT);
374 for child in shown {
375 let restates = child.subtree * 100 >= node.subtree * PASS_THROUGH_PERCENT
376 && child.own * 100 < total * FOLD_PERCENT;
377 match restates && !child.children.is_empty() {
378 true => children(child, render, depth, total, rows),
379 false => emit(child, render, depth, total, rows),
380 }
381 }
382 if !folded.is_empty() {
383 let under: u128 = folded.iter().map(|c| c.subtree).sum();
384 rows.push(Row {
385 depth,
386 node: format!("{} others", folded.len()),
387 own: under,
388 subtree: under,
389 percent: percent(under, total),
390 route: "folded",
391 shared: false,
392 });
393 }
394}
395
396fn percent(part: u128, whole: u128) -> f64 {
397 match whole {
398 0 => 0.0,
399 _ => 100.0 * part as f64 / whole as f64,
400 }
401}