Skip to main content

sva_samples/collapse/
plan.rs

1// Concern: which row of the collapse table a form takes, what it costs and its direct sums' bounds | Non-concern: running the row | IO: (&SpectralSum, rate, Horizon) -> Plan, flops, bounds
2
3use sva_formula::closed_form::{Part, map_children};
4use sva_formula::spectral_sum::atom::SpectralAtom;
5use sva_formula::{Body, ClosedForm, Lane, Line, SpectralSum, Var, normalize_closed_form};
6
7use super::truncate::Audible;
8use super::{Horizon, atoms, lines, point, truncate};
9use crate::error::CollapseError;
10use crate::label::Rule;
11use crate::profile::Profile;
12
13/// Both routes place a line exactly, so cost decides; the direct one hands a lone line back
14/// unchanged.
15pub fn direct_is_cheaper(lines: usize, n: usize, len: usize) -> bool {
16    (lines as f64) * (len as f64) <= (n as f64) * (n as f64).log2()
17}
18
19pub struct LinePlan {
20    pub placed: Vec<Vec<Line>>,
21    pub summed: Vec<Vec<Line>>,
22    pub dropped: Vec<Vec<Line>>,
23    pub bins: usize,
24    pub tail_db: Option<f64>,
25}
26
27/// FORMAT 9.2's window route.
28pub struct Group {
29    pub factor: SpectralAtom,
30    pub placed: Vec<Line>,
31    pub summed: Vec<Line>,
32}
33
34pub enum LanePlan {
35    Grouped { groups: Vec<Group>, bins: usize },
36    Sweep { atoms: usize, samples: usize },
37}
38
39pub struct Sampled {
40    pub sum: Box<SpectralSum>,
41    pub rule: Rule,
42    pub lanes: Vec<LanePlan>,
43}
44
45/// The row of FORMAT 9.1 that runs, decided once for the collapse and the count alike.
46pub enum Plan {
47    Spectrum(Box<SpectralSum>),
48    Lines(Box<LinePlan>),
49    Sampled(Box<Sampled>),
50    /// Row 4 over the written closed form: every node the truncation leaves, at every instant.
51    Point {
52        written: Box<ClosedForm>,
53        nodes: usize,
54        width: usize,
55    },
56    /// A sum whose addends name different rows, each taking its own.
57    Added(Vec<Plan>),
58}
59
60/// The rows in FORMAT 9.1's order; the first guard that holds decides.
61pub fn of(
62    sum: &SpectralSum,
63    rate: u32,
64    horizon: Horizon,
65    profile: &Profile,
66    len: usize,
67) -> Result<Plan, CollapseError> {
68    let ceiling = profile.ceiling(rate);
69    if sum.var == Var::F {
70        return Ok(Plan::Spectrum(Box::new(sum.clone())));
71    }
72    if let Some(found) = line_plan(sum, rate, horizon, profile, len, ceiling)? {
73        return Ok(Plan::Lines(Box::new(found)));
74    }
75    let truncated = truncate::spectral_sum(sum, Audible::of(profile, rate))?;
76    let rule = match () {
77        () if atoms::band_limited(&truncated, ceiling, profile) => Rule::BandLimited,
78        () if atoms::windowed(&truncated) => Rule::CroppedPair,
79        () => Rule::PointSampled,
80    };
81    let lanes = truncated
82        .lanes
83        .iter()
84        .map(|lane| lane_plan(lane, rate, horizon, len))
85        .collect();
86    Ok(Plan::Sampled(Box::new(Sampled {
87        sum: Box::new(truncated),
88        rule,
89        lanes,
90    })))
91}
92
93/// The row a closed form with no spectral sum takes: one per addend where they differ, else 4.
94pub fn of_written(
95    form: &ClosedForm,
96    rate: u32,
97    horizon: Horizon,
98    profile: &Profile,
99    len: usize,
100) -> Result<Plan, CollapseError> {
101    let Some(addends) = addends(form) else {
102        return point_plan(form, rate, profile);
103    };
104    let mut parts = Vec::with_capacity(addends.len());
105    for addend in &addends {
106        match of_term(addend, rate, horizon, profile, len)? {
107            Plan::Added(inner) => parts.extend(inner),
108            part => parts.push(part),
109        }
110    }
111    // Addends that all name row 4 are one row 4 over the whole sum, measured once.
112    match parts.iter().all(|part| matches!(part, Plan::Point { .. })) {
113        true => point_plan(form, rate, profile),
114        false => Ok(Plan::Added(parts)),
115    }
116}
117
118/// A row the addend alone cannot take refuses nothing; the written row still stands for it.
119/// A nesting past the bound is the one exception, refused wherever it is written.
120fn of_term(
121    form: &ClosedForm,
122    rate: u32,
123    horizon: Horizon,
124    profile: &Profile,
125    len: usize,
126) -> Result<Plan, CollapseError> {
127    let Ok(sum) = normalize_closed_form(form) else {
128        return of_written(form, rate, horizon, profile, len);
129    };
130    match of(&sum, rate, horizon, profile, len) {
131        Err(nested @ CollapseError::NestedSeries { .. }) => Err(nested),
132        Err(_) => of_written(form, rate, horizon, profile, len),
133        held => held,
134    }
135}
136
137pub(super) fn addends(form: &ClosedForm) -> Option<Vec<ClosedForm>> {
138    if let Body::Add(parts) = &form.body {
139        return Some(
140            parts
141                .iter()
142                .map(|part| ClosedForm {
143                    var: form.var,
144                    body: (*part.body).clone(),
145                    origin: part.origin,
146                })
147                .collect(),
148        );
149    }
150    let under = linear_over(&form.body)?;
151    let held = addends(&ClosedForm {
152        body: under,
153        ..form.clone()
154    })?;
155    Some(
156        held.into_iter()
157            .map(|addend| ClosedForm {
158                body: map_children(&form.body, |_| {
159                    Part::new(addend.origin, addend.body.clone())
160                }),
161                ..addend
162            })
163            .collect(),
164    )
165}
166
167/// Each of these is linear, so over a sum it is the sum of itself over every addend.
168fn linear_over(f: &Body) -> Option<Body> {
169    match f {
170        Body::Crop { of, .. }
171        | Body::Shift { of, .. }
172        | Body::Deriv { of, .. }
173        | Body::Channel(of, _) => Some((*of.body).clone()),
174        _ => None,
175    }
176}
177
178fn point_plan(form: &ClosedForm, rate: u32, profile: &Profile) -> Result<Plan, CollapseError> {
179    let written = ClosedForm {
180        body: truncate::written(&form.body, Audible::of(profile, rate))?,
181        ..form.clone()
182    };
183    let width = point::width_of(&written.body, &point::NoRefs).max(1);
184    Ok(Plan::Point {
185        nodes: point_nodes(&written.body),
186        written: Box::new(written),
187        width,
188    })
189}
190
191/// What one instant's evaluation walks over every component: a join walks that component's branch.
192pub fn point_nodes(f: &Body) -> usize {
193    let width = point::width_of(f, &point::NoRefs).max(1);
194    (0..width).map(|c| point_work(f, c).0).sum()
195}
196
197/// One component's `(nodes, waves)` at one instant: a run is priced by its Horner steps and
198/// turns its lines; every other node is one.
199pub fn point_work(f: &Body, component: usize) -> (usize, usize) {
200    let add = |a: (usize, usize), b: (usize, usize)| (a.0 + b.0, a.1 + b.1);
201    let branch = |part: &Part, c: usize| add((1, 0), point_work(&part.body, c));
202    match f {
203        Body::Run(run) => (super::run::steps(run), run.len()),
204        Body::Join(parts) => {
205            let widths: Vec<usize> = parts
206                .iter()
207                .map(|p| point::width_of(&p.body, &point::NoRefs))
208                .collect();
209            match point::lane_of(&widths, component) {
210                Some((at, inner)) => branch(&parts[at], inner),
211                None => (1, 0),
212            }
213        }
214        Body::Channel(of, k) => branch(of, usize::from(*k)),
215        other => sva_formula::closed_form::children(other)
216            .iter()
217            .map(|part| point_work(&part.body, component))
218            .fold((1, 0), add),
219    }
220}
221
222/// Each direct sum a row may take over this form's lines, as its rounding bound and the
223/// factor it is read under: none on a line row, the group's own on a windowed one.
224pub fn summed_bounds(
225    sum: &SpectralSum,
226    profile: &Profile,
227    rate: u32,
228) -> Result<Vec<(Option<SpectralAtom>, f64)>, CollapseError> {
229    if sum.var == Var::F {
230        return Ok(Vec::new());
231    }
232    if let Some(found) = kept_lines(sum, profile, profile.ceiling(rate))? {
233        let direct = found.kept.iter().filter_map(|kept| lines::Direct::of(kept));
234        return Ok(direct.map(|d| (None, d.bound())).collect());
235    }
236    let truncated = truncate::spectral_sum(sum, Audible::of(profile, rate))?;
237    let groups = truncated.lanes.iter().filter_map(lines::grouped).flatten();
238    Ok(groups
239        .filter_map(|(factor, held)| Some((Some(factor), lines::Direct::of(&held)?.bound())))
240        .collect())
241}
242
243/// Rows one and two: every atom a line, none of them windowed.
244fn line_plan(
245    sum: &SpectralSum,
246    rate: u32,
247    horizon: Horizon,
248    profile: &Profile,
249    len: usize,
250    ceiling: f64,
251) -> Result<Option<LinePlan>, CollapseError> {
252    let Some(found) = kept_lines(sum, profile, ceiling)? else {
253        return Ok(None);
254    };
255    let bins = lines::grid(&found.kept, &found.grids, rate, horizon.span(), len);
256    let split: Vec<(Vec<Line>, Vec<Line>)> = found
257        .kept
258        .iter()
259        .map(|k| match direct_is_cheaper(k.len(), bins, len) {
260            true => (Vec::new(), k.clone()),
261            false => lines::split(k, bins, rate),
262        })
263        .collect();
264    Ok(Some(LinePlan {
265        placed: split.iter().map(|(on, _)| on.clone()).collect(),
266        summed: split.into_iter().map(|(_, off)| off).collect(),
267        dropped: found.dropped,
268        bins,
269        tail_db: found.tail,
270    }))
271}
272
273/// Each lane's lines under the ceiling and over it, before any horizon places them.
274pub(super) struct Kept {
275    pub(super) kept: Vec<Vec<Line>>,
276    dropped: Vec<Vec<Line>>,
277    grids: Vec<f64>,
278    tail: Option<f64>,
279}
280
281pub(super) fn kept_lines(
282    sum: &SpectralSum,
283    profile: &Profile,
284    ceiling: f64,
285) -> Result<Option<Kept>, CollapseError> {
286    let mut per_lane = Vec::with_capacity(sum.lanes.len());
287    let mut grids: Vec<f64> = Vec::new();
288    let mut tail: Option<f64> = None;
289    for lane in &sum.lanes {
290        match lines::of_lane(lane, ceiling, profile.floor(ceiling), profile.half_lsb()) {
291            Some(found) => {
292                if let Some(left) = found.tail_db {
293                    tail = Some(tail.map_or(left, |held: f64| held.max(left)));
294                }
295                grids.extend(found.grids);
296                per_lane.push(found.lines);
297            }
298            None => return Ok(None),
299        }
300    }
301
302    let (mut kept, mut dropped): (Vec<Vec<Line>>, Vec<Vec<Line>>) = (Vec::new(), Vec::new());
303    for found in &per_lane {
304        let (here, gone): (Vec<Line>, Vec<Line>) = found.iter().partition(|l| l.hz.abs() < ceiling);
305        kept.push(here);
306        dropped.push(gone);
307    }
308    // A form with no line at all is silence, not a band every line sat above.
309    if kept.iter().all(Vec::is_empty) && !dropped.iter().all(Vec::is_empty) {
310        let lowest = dropped
311            .iter()
312            .flatten()
313            .map(|l| l.hz.abs())
314            .fold(f64::INFINITY, f64::min);
315        return Err(CollapseError::EmptyBand { ceiling, lowest });
316    }
317    Ok(Some(Kept {
318        kept,
319        dropped,
320        grids,
321        tail,
322    }))
323}
324
325fn lane_plan(lane: &Lane, rate: u32, horizon: Horizon, len: usize) -> LanePlan {
326    let bins = lines::bins(horizon.span(), rate);
327    let sweep = LanePlan::Sweep {
328        atoms: lane.atoms.len(),
329        samples: super::span::nonzero(lane, horizon, rate, len)
330            .iter()
331            .map(|(from, to)| to - from)
332            .sum(),
333    };
334    let Some(found) = lines::grouped(lane) else {
335        return sweep;
336    };
337    let groups: Vec<Group> = found
338        .into_iter()
339        .map(|(factor, held)| {
340            let (placed, summed) = match direct_is_cheaper(held.len(), bins, len) {
341                true => (Vec::new(), held),
342                false => lines::split(&held, bins, rate),
343            };
344            Group {
345                factor,
346                placed,
347                summed,
348            }
349        })
350        .collect();
351    let grouped = LanePlan::Grouped { groups, bins };
352    match grouped.flops(len) <= sweep.flops(len) {
353        true => grouped,
354        false => sweep,
355    }
356}
357
358/// `n*log2(n)` butterflies; a length that is no power of two is Bluestein's three
359/// transforms over the next one past `2n-1`.
360pub fn transform_flops(n: usize) -> u128 {
361    let stages = |m: usize| m as u128 * (m.max(2).trailing_zeros() as u128);
362    match n.is_power_of_two() {
363        true => stages(n),
364        false => 3 * stages((2 * n - 1).next_power_of_two()),
365    }
366}
367
368impl LinePlan {
369    pub fn flops(&self, len: usize) -> u128 {
370        let transforms: u128 = self
371            .placed
372            .iter()
373            .filter(|p| !p.is_empty())
374            .map(|_| transform_flops(self.bins))
375            .sum();
376        let direct: u128 = self
377            .summed
378            .iter()
379            .map(|s| s.len() as u128 * len as u128)
380            .sum();
381        transforms + direct
382    }
383
384    pub fn rule(&self) -> Rule {
385        match (lines::distinct(&self.placed), lines::distinct(&self.summed)) {
386            (_, 0) => Rule::LineSpectrumExact,
387            (0, _) => Rule::LineSpectrumSummed,
388            _ => Rule::LineSpectrumMixed,
389        }
390    }
391}
392
393impl LanePlan {
394    pub fn flops(&self, len: usize) -> u128 {
395        match self {
396            LanePlan::Grouped { groups, bins } => groups
397                .iter()
398                .map(|g| {
399                    let placed = match g.placed.is_empty() {
400                        true => 0,
401                        false => transform_flops(*bins),
402                    };
403                    placed + (g.summed.len() as u128 + 1) * len as u128
404                })
405                .sum(),
406            LanePlan::Sweep { atoms, samples } => *atoms as u128 * *samples as u128,
407        }
408    }
409
410    fn swept_flops(&self, len: u128) -> u128 {
411        let terms: u128 = match self {
412            LanePlan::Grouped { groups, .. } => groups
413                .iter()
414                .map(|g| (g.placed.len() + g.summed.len() + 1) as u128)
415                .sum(),
416            LanePlan::Sweep { atoms, .. } => *atoms as u128,
417        };
418        terms * len
419    }
420}
421
422impl Plan {
423    pub fn flops(&self, len: usize) -> u128 {
424        match self {
425            Plan::Spectrum(_) => transform_flops(len),
426            Plan::Lines(found) => found.flops(len),
427            Plan::Sampled(held) => held.lanes.iter().map(|lane| lane.flops(len)).sum(),
428            Plan::Point { nodes, .. } => *nodes as u128 * len as u128,
429            Plan::Added(parts) => parts.iter().map(|part| part.flops(len)).sum(),
430        }
431    }
432
433    pub fn alias_flops(&self, len: usize) -> u128 {
434        self.alias_flops_at(len, super::ALIAS_OVERSAMPLE)
435    }
436
437    /// The multiple a reading names need not be the label's `ALIAS_OVERSAMPLE`. Every component
438    /// is scored, so every component's reference is paid for.
439    pub fn alias_flops_at(&self, len: usize, oversample: usize) -> u128 {
440        let reference = (len * oversample) as u128;
441        match self {
442            Plan::Point { nodes, .. } => *nodes as u128 * reference,
443            Plan::Sampled(held) if held.rule == Rule::PointSampled => held
444                .lanes
445                .iter()
446                .map(|lane| lane.swept_flops(reference))
447                .sum(),
448            Plan::Added(parts) => parts
449                .iter()
450                .map(|part| part.alias_flops_at(len, oversample))
451                .sum(),
452            Plan::Spectrum(_) | Plan::Lines(_) | Plan::Sampled(_) => 0,
453        }
454    }
455
456    pub fn rule(&self) -> Rule {
457        match self {
458            Plan::Spectrum(_) => Rule::InverseSpectrum,
459            Plan::Lines(found) => found.rule(),
460            Plan::Sampled(held) => held.rule,
461            Plan::Point { .. } => Rule::PointSampled,
462            Plan::Added(_) => Rule::Added,
463        }
464    }
465}