Skip to main content

henad_explore/
probe.rs

1//! Probe builds of a model, made before a sweep to fix its stat columns and measure the footprint of a run.
2
3use std::fmt;
4
5use web_time::Instant;
6
7use henad_compute::entry::{ModelEntry, ModelState};
8use henad_compute::fault::{Fault, STEPPING, catching};
9use henad_compute::gpu::{Demand, GpuContext};
10#[cfg(not(target_arch = "wasm32"))]
11use henad_compute::gpu::{GpuSimState, fault::catching_on, stepping};
12use henad_core::explore::plan::Plan;
13use henad_core::export::StatColumns;
14use henad_core::params::ParamValue;
15
16use crate::output::manifest::now_unix_ms;
17
18/// Maximum number of rejected configs that a [`CapacityError`] lists.
19pub const MAX_LISTED_CONFIGS: usize = 5;
20
21/// Maximum number of configs that [`ProbeReport::for_plan`] builds before it gives up.
22pub const MAX_PROBED_CONFIGS: usize = 8;
23
24/// Stat columns and footprint of one build of a model, sampled at tick 0.
25#[derive(Debug)]
26pub struct ProbeReport {
27    /// Values of the config the probe built, one per parameter.
28    pub params: Vec<ParamValue>,
29    /// Seed the probe built with.
30    pub seed: Option<u64>,
31    /// Columns of the sample at tick 0. Every later sample of the sweep must fit them.
32    pub columns: StatColumns,
33    /// Size in bytes of the state on the host.
34    pub heap_bytes: u64,
35    /// Population at tick 0.
36    pub population: u64,
37    /// Number of jobs that one step splits into, `None` when the backend does not say.
38    pub parallel_jobs: Option<usize>,
39    /// Device resources of the build, `None` for a CPU model.
40    pub demand: Option<Demand>,
41}
42
43impl ProbeReport {
44    /// Builds `entry` at `params` with `seed`, and samples tick 0 as a run of a sweep does.
45    ///
46    /// # Errors
47    ///
48    /// Returns [`ProbeError::NoDevice`] for a GPU model with no device in `gpu`, and [`ProbeError::Fault`] for a
49    /// build or a sample that faults.
50    pub fn build(
51        entry: &ModelEntry,
52        gpu: Option<&GpuContext>,
53        params: &[ParamValue],
54        seed: Option<u64>,
55    ) -> Result<Self, ProbeError> {
56        if entry.gpu_needs().is_some() && gpu.is_none() {
57            return Err(ProbeError::NoDevice);
58        }
59        match entry.build(params, seed, gpu).map_err(ProbeError::Fault)? {
60            ModelState::Cpu(mut state) => {
61                let stats = catching(STEPPING, || {
62                    state.prepare_view();
63                    state.stats()
64                })
65                .map_err(ProbeError::Fault)?;
66                Ok(Self {
67                    params: params.to_vec(),
68                    seed,
69                    columns: StatColumns::plan(&stats),
70                    heap_bytes: state.heap_bytes() as u64,
71                    population: state.population(),
72                    parallel_jobs: state.parallel_jobs(),
73                    demand: None,
74                })
75            }
76            ModelState::Gpu(state) => probe_gpu(entry, state, gpu, params, seed),
77        }
78    }
79
80    /// Builds the configs of `plan` in order, each with the seed of its first run, and returns the first report
81    /// without a fault.
82    ///
83    /// A config that faults is left for its runs to record. Note that only the first [`MAX_PROBED_CONFIGS`] configs
84    /// are tried.
85    ///
86    /// # Errors
87    ///
88    /// Returns [`ProbeError::NoDevice`] for a GPU model with no device in `gpu`, and
89    /// [`ProbeError::EveryConfigFaulted`] when every config tried faults.
90    pub fn for_plan(entry: &ModelEntry, gpu: Option<&GpuContext>, plan: &Plan) -> Result<Self, ProbeError> {
91        let mut probe = PlanProbe::new();
92        loop {
93            if let Some(timed) = probe.step(entry, gpu, plan)? {
94                return Ok(timed.report);
95            }
96        }
97    }
98
99    /// Builds the last config of `plan` with the seed of its first run.
100    ///
101    /// Returns `None` for a plan of one config or a last config `probed` already built. Also returns `None` for a last
102    /// config that faults, and leaves the fault for its runs to record.
103    pub fn for_last_config(entry: &ModelEntry, gpu: Option<&GpuContext>, plan: &Plan, probed: &Self) -> Option<Self> {
104        let config_id = plan.configs().len().checked_sub(1)? as u64;
105        let config = plan.config(config_id)?;
106        let run = plan.run(config_id * plan.replicates())?;
107        if config_id == 0 || (config.params == probed.params && Some(run.seed) == probed.seed) {
108            return None;
109        }
110        Self::build(entry, gpu, &config.params, Some(run.seed)).ok()
111    }
112
113    /// Memory in bytes that the build holds, on the host and on the device together.
114    pub fn footprint(&self) -> u64 {
115        self.heap_bytes + self.demand.as_ref().map_or(0, Demand::bytes)
116    }
117
118    /// Returns the report of the same CPU build made on a thread pool with `threads` workers.
119    ///
120    /// Note that a buffer sized to the pool, such as a scatter grid's scratch, can change the host bytes.
121    ///
122    /// # Errors
123    ///
124    /// Returns [`ProbeError::Pool`] when the pool cannot be built, and the error that [`Self::build`] returns.
125    pub fn rebuilt_on(&self, entry: &ModelEntry, threads: usize) -> Result<Self, ProbeError> {
126        let pool = rayon::ThreadPoolBuilder::new()
127            .num_threads(threads)
128            .build()
129            .map_err(ProbeError::Pool)?;
130        crate::exec::run_in_pool(&pool, || Self::build(entry, None, &self.params, self.seed))
131    }
132}
133
134/// Probe of a plan's configs in progress, one build per call to [`PlanProbe::step`].
135///
136/// The configs are tried as [`ProbeReport::for_plan`] tries them.
137#[derive(Debug)]
138pub(crate) struct PlanProbe {
139    /// Config the next step builds.
140    next_config: u64,
141    /// Faults of the configs tried so far.
142    faults: Vec<ConfigFault>,
143    /// Clock reading when the probe was created.
144    started: Instant,
145    /// Wall clock time when the probe was created, in milliseconds since the Unix epoch.
146    started_unix_ms: u64,
147}
148
149/// Report of the first config of a plan that builds, with the clock readings taken when its probe was created.
150#[derive(Debug)]
151pub(crate) struct TimedProbe {
152    pub(crate) report: ProbeReport,
153    pub(crate) started: Instant,
154    pub(crate) started_unix_ms: u64,
155}
156
157impl PlanProbe {
158    /// Returns a probe that has tried no config yet, timed from now.
159    pub(crate) fn new() -> Self {
160        Self {
161            next_config: 0,
162            faults: Vec::new(),
163            started: Instant::now(),
164            started_unix_ms: now_unix_ms(),
165        }
166    }
167
168    /// Builds the next config of `plan` with the seed of its first run, and returns its report with the probe's clock
169    /// readings when it builds without a fault.
170    ///
171    /// Returns `Ok(None)` for a config that faults while configs are left to try. The fault is left for the config's
172    /// runs to record.
173    ///
174    /// # Errors
175    ///
176    /// Returns [`ProbeError::NoDevice`] for a GPU model with no device in `gpu`, and
177    /// [`ProbeError::EveryConfigFaulted`] once every config tried faults.
178    pub(crate) fn step(
179        &mut self,
180        entry: &ModelEntry,
181        gpu: Option<&GpuContext>,
182        plan: &Plan,
183    ) -> Result<Option<TimedProbe>, ProbeError> {
184        let config_limit = plan.configs().len().min(MAX_PROBED_CONFIGS) as u64;
185        let config_id = self.next_config;
186        let Some(config) = plan.config(config_id).filter(|_| config_id < config_limit) else {
187            return Err(ProbeError::EveryConfigFaulted(std::mem::take(&mut self.faults)));
188        };
189        self.next_config += 1;
190        let run = plan
191            .run(config_id * plan.replicates())
192            .expect("every config has a first run");
193        match ProbeReport::build(entry, gpu, &config.params, Some(run.seed)) {
194            Ok(report) => Ok(Some(TimedProbe {
195                report,
196                started: self.started,
197                started_unix_ms: self.started_unix_ms,
198            })),
199            Err(ProbeError::Fault(fault)) => {
200                self.faults.push(ConfigFault { config_id, fault });
201                if self.next_config < config_limit {
202                    Ok(None)
203                } else {
204                    Err(ProbeError::EveryConfigFaulted(std::mem::take(&mut self.faults)))
205                }
206            }
207            Err(error) => Err(error),
208        }
209    }
210}
211
212#[cfg(not(target_arch = "wasm32"))]
213fn probe_gpu(
214    entry: &ModelEntry,
215    mut state: Box<dyn GpuSimState>,
216    gpu: Option<&GpuContext>,
217    params: &[ParamValue],
218    seed: Option<u64>,
219) -> Result<ProbeReport, ProbeError> {
220    let ctx = gpu.ok_or(ProbeError::NoDevice)?;
221    let stats = catching_on(ctx, STEPPING, || stepping::sample_stats(&mut *state, ctx))
222        .flatten()
223        .map_err(ProbeError::Fault)?;
224    stepping::wait(ctx).map_err(ProbeError::Fault)?;
225    Ok(ProbeReport {
226        params: params.to_vec(),
227        seed,
228        columns: StatColumns::plan(&stats),
229        heap_bytes: state.heap_bytes() as u64,
230        population: state.population(),
231        parallel_jobs: state.parallel_jobs(),
232        demand: entry.demand(params, &ctx.device.limits()),
233    })
234}
235
236/// Rejects a GPU model. A browser cannot block on the device, and a sample blocks.
237#[cfg(target_arch = "wasm32")]
238fn probe_gpu(
239    _entry: &ModelEntry,
240    _state: Box<dyn henad_compute::gpu::GpuSimState>,
241    _gpu: Option<&GpuContext>,
242    _params: &[ParamValue],
243    _seed: Option<u64>,
244) -> Result<ProbeReport, ProbeError> {
245    Err(ProbeError::NoDevice)
246}
247
248/// Fault of the probe build of one config.
249#[derive(Debug)]
250pub struct ConfigFault {
251    /// Id of the config in its plan.
252    pub config_id: u64,
253    /// Fault of the build or of its sample at tick 0.
254    pub fault: Fault,
255}
256
257/// A probe build that cannot run.
258#[derive(Debug)]
259pub enum ProbeError {
260    /// A GPU model with no device to sample it on. In a browser the probe never receives a device.
261    NoDevice,
262    /// The build or its sample at tick 0 faulted.
263    Fault(Fault),
264    /// Every config [`ProbeReport::for_plan`] tried faulted, each with its fault.
265    EveryConfigFaulted(Vec<ConfigFault>),
266    /// Building the thread pool for a probe build failed.
267    Pool(rayon::ThreadPoolBuildError),
268}
269
270impl fmt::Display for ProbeError {
271    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
272        match self {
273            Self::NoDevice => f.write_str("a GPU sweep needs a GPU device and a native build"),
274            Self::Fault(_) => f.write_str("the probe build failed"),
275            Self::EveryConfigFaulted(faults) => {
276                write!(
277                    f,
278                    "the probe build faulted on each of the {} configs tried",
279                    faults.len()
280                )?;
281                for config in faults {
282                    write!(f, "\n  config {}: {}", config.config_id, config.fault)?;
283                }
284                Ok(())
285            }
286            Self::Pool(_) => f.write_str("cannot build the thread pool of the probe build"),
287        }
288    }
289}
290
291impl std::error::Error for ProbeError {
292    fn source(&self) -> Option<&(dyn std::error::Error + 'static)> {
293        match self {
294            Self::NoDevice | Self::EveryConfigFaulted(_) => None,
295            Self::Fault(fault) => Some(fault),
296            Self::Pool(error) => Some(error),
297        }
298    }
299}
300
301/// A config the device cannot host.
302#[derive(Debug, Clone, PartialEq, Eq)]
303pub struct RefusedConfig {
304    /// Id of the config in its plan.
305    pub config_id: u64,
306    /// Reasons the device rejects the config, from [`ModelEntry::shortfalls`].
307    pub reasons: Vec<String>,
308}
309
310/// Configs of a plan the device cannot host, found before any run.
311#[derive(Debug, Clone, PartialEq, Eq)]
312pub struct CapacityError {
313    /// First rejected configs, up to [`MAX_LISTED_CONFIGS`] of them.
314    pub refused: Vec<RefusedConfig>,
315    /// Number of rejected configs.
316    pub count: u64,
317}
318
319impl fmt::Display for CapacityError {
320    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
321        if self.count == 1 {
322            f.write_str("1 config does not fit this GPU")?;
323        } else {
324            write!(f, "{} configs do not fit this GPU", self.count)?;
325        }
326        for config in &self.refused {
327            for reason in &config.reasons {
328                write!(f, "\n  config {}: {reason}", config.config_id)?;
329            }
330        }
331        Ok(())
332    }
333}
334
335impl std::error::Error for CapacityError {}
336
337/// Checks every config of `plan` against `limits` before any run. A CPU model always passes.
338///
339/// # Errors
340///
341/// Returns [`CapacityError`] when some config needs more than the device allows.
342pub(crate) fn check_capacity(entry: &ModelEntry, plan: &Plan, limits: &wgpu::Limits) -> Result<(), CapacityError> {
343    if entry.gpu_needs().is_none() {
344        return Ok(());
345    }
346    let mut error = CapacityError {
347        refused: Vec::new(),
348        count: 0,
349    };
350    for (config_id, config) in (0_u64..).zip(plan.configs()) {
351        let reasons = entry.shortfalls(&config.params, limits);
352        if reasons.is_empty() {
353            continue;
354        }
355        error.count += 1;
356        if error.refused.len() < MAX_LISTED_CONFIGS {
357            error.refused.push(RefusedConfig { config_id, reasons });
358        }
359    }
360    if error.count == 0 { Ok(()) } else { Err(error) }
361}
362
363#[cfg(test)]
364mod tests {
365    use henad_compute::entry::{ModelEntry, register_grid_model};
366    use henad_compute::fault::install_panic_hook;
367    use henad_core::explore::design::DesignKind;
368    use henad_core::explore::factor::{FactorSpec, LevelSpec};
369    use henad_core::explore::spec::{BlockSpec, SweepSpec};
370    use henad_core::params::ParamValue;
371    use henad_models::example_models;
372
373    use super::{MAX_LISTED_CONFIGS, MAX_PROBED_CONFIGS, ProbeError, ProbeReport, check_capacity};
374    use crate::tests::broken::DividesByParam;
375
376    fn entry(id: &str) -> ModelEntry {
377        example_models().get(id).cloned().expect("the model is registered")
378    }
379
380    #[test]
381    fn a_probe_reads_the_columns_and_footprint_of_config_zero() {
382        let entry = entry("boids");
383        let mut spec = SweepSpec::new("boids");
384        spec.fixed = vec![("num_agents".to_owned(), "300".to_owned())];
385        let plan = spec.plan(&entry.schema()).expect("a valid spec");
386        let probe = ProbeReport::for_plan(&entry, None, &plan).expect("boids builds");
387        assert_eq!(probe.population, 300);
388        assert!(probe.heap_bytes > 0);
389        assert!(probe.parallel_jobs.is_some());
390        assert!(probe.demand.is_none(), "boids runs on the CPU");
391        assert!(!probe.columns.is_empty());
392    }
393
394    /// Returns a spec over `model` whose configs are the square grids with the sides `sides`.
395    fn square_grids(model: &str, sides: &[&str]) -> SweepSpec {
396        let levels = LevelSpec::Values(sides.iter().map(|&side| side.to_owned()).collect());
397        let mut spec = SweepSpec::new(model);
398        spec.blocks = vec![BlockSpec {
399            design: DesignKind::Zip,
400            factors: ["grid_width", "grid_height"]
401                .map(|id| FactorSpec::param(id, levels.clone()))
402                .to_vec(),
403            design_seed: None,
404        }];
405        spec
406    }
407
408    #[cfg(not(target_arch = "wasm32"))]
409    #[test]
410    fn check_capacity_counts_every_config_past_the_limits_and_passes_a_cpu_model() {
411        let Some(ctx) = crate::tests::support::headless_device() else {
412            return;
413        };
414        let gpu_sir = crate::tests::support::entry("gpu_sir", Some(&ctx));
415        let sides = ["16", "32", "48", "64", "80", "96", "112"];
416        let plan = square_grids("gpu_sir", &sides)
417            .plan(&gpu_sir.schema())
418            .expect("a valid spec");
419        let fitting = plan.config(0).expect("the plan has config 0");
420        let demand = gpu_sir
421            .demand(&fitting.params, &ctx.device.limits())
422            .expect("a GPU model has a demand");
423        let largest = demand
424            .buffers
425            .iter()
426            .map(|alloc| alloc.bytes)
427            .max()
428            .expect("gpu_sir allocates buffers");
429        // The largest buffer of the smallest grid fills a binding, and every larger grid needs more.
430        let limits = wgpu::Limits {
431            max_storage_buffer_binding_size: largest,
432            ..wgpu::Limits::default()
433        };
434        let error = check_capacity(&gpu_sir, &plan, &limits).expect_err("the larger grids pass the binding size");
435        assert_eq!(error.count, 6);
436        let listed: Vec<u64> = error.refused.iter().map(|config| config.config_id).collect();
437        assert_eq!(listed, (1..=MAX_LISTED_CONFIGS as u64).collect::<Vec<_>>());
438        assert!(error.refused.iter().all(|config| !config.reasons.is_empty()));
439
440        let game_of_life = entry("game_of_life");
441        let plan = square_grids("game_of_life", &sides)
442            .plan(&game_of_life.schema())
443            .expect("a valid spec");
444        assert!(
445            check_capacity(&game_of_life, &plan, &limits).is_ok(),
446            "a CPU model passes limits that refuse a GPU model"
447        );
448    }
449
450    /// Returns a spec over `DividesByParam` on an 8 by 8 grid, varying `init_divisor` over `levels`.
451    fn init_divisors(levels: &[&str]) -> SweepSpec {
452        let mut spec = SweepSpec::new("divides_by_param");
453        spec.fixed = vec![
454            ("grid_width".to_owned(), "8".to_owned()),
455            ("grid_height".to_owned(), "8".to_owned()),
456        ];
457        spec.run.replicates = 2;
458        spec.blocks = vec![BlockSpec {
459            design: DesignKind::Factorial,
460            factors: vec![FactorSpec::param(
461                "init_divisor",
462                LevelSpec::Values(levels.iter().map(|&level| level.to_owned()).collect()),
463            )],
464            design_seed: None,
465        }];
466        spec
467    }
468
469    #[test]
470    fn a_config_that_faults_is_left_to_its_runs() {
471        install_panic_hook();
472        let entry = register_grid_model::<DividesByParam>();
473        let plan = init_divisors(&["0", "0", "1"])
474            .plan(&entry.schema())
475            .expect("a valid spec");
476        let probe = ProbeReport::for_plan(&entry, None, &plan).expect("config 2 builds");
477        assert_eq!(probe.params[3], ParamValue::U32(1), "init_divisor of config 2");
478        assert_eq!(
479            probe.seed,
480            plan.run(4).map(|run| run.seed),
481            "the seed of config 2's first run"
482        );
483
484        let zeros = vec!["0"; MAX_PROBED_CONFIGS + 1];
485        let plan = init_divisors(&zeros).plan(&entry.schema()).expect("a valid spec");
486        let error = ProbeReport::for_plan(&entry, None, &plan).expect_err("no config builds");
487        let ProbeError::EveryConfigFaulted(faults) = &error else {
488            panic!("{error:?}");
489        };
490        let tried: Vec<u64> = faults.iter().map(|config| config.config_id).collect();
491        assert_eq!(tried, (0..MAX_PROBED_CONFIGS as u64).collect::<Vec<_>>());
492        let text = error.to_string();
493        assert!(text.contains("\n  config 7: while building the model"), "{text}");
494    }
495
496    #[test]
497    fn a_rebuild_on_a_narrower_pool_holds_less_scratch() {
498        let entry = entry("ants");
499        let mut spec = SweepSpec::new("ants");
500        spec.fixed = vec![
501            ("num_agents".to_owned(), "300".to_owned()),
502            ("world_width".to_owned(), "64".to_owned()),
503            ("world_height".to_owned(), "64".to_owned()),
504        ];
505        let plan = spec.plan(&entry.schema()).expect("a valid spec");
506        let probe = ProbeReport::for_plan(&entry, None, &plan).expect("ants builds");
507        let [one, four] = [1, 4].map(|threads| probe.rebuilt_on(&entry, threads).expect("ants builds on a pool"));
508        assert!(
509            one.heap_bytes < four.heap_bytes,
510            "a scatter grid keeps a shadow grid per worker: {} and {}",
511            one.heap_bytes,
512            four.heap_bytes
513        );
514        assert_eq!((one.params, one.seed), (probe.params, probe.seed));
515        assert_eq!(one.columns, four.columns);
516    }
517}