1use 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
18pub const MAX_LISTED_CONFIGS: usize = 5;
20
21pub const MAX_PROBED_CONFIGS: usize = 8;
23
24#[derive(Debug)]
26pub struct ProbeReport {
27 pub params: Vec<ParamValue>,
29 pub seed: Option<u64>,
31 pub columns: StatColumns,
33 pub heap_bytes: u64,
35 pub population: u64,
37 pub parallel_jobs: Option<usize>,
39 pub demand: Option<Demand>,
41}
42
43impl ProbeReport {
44 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 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 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 pub fn footprint(&self) -> u64 {
115 self.heap_bytes + self.demand.as_ref().map_or(0, Demand::bytes)
116 }
117
118 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#[derive(Debug)]
138pub(crate) struct PlanProbe {
139 next_config: u64,
141 faults: Vec<ConfigFault>,
143 started: Instant,
145 started_unix_ms: u64,
147}
148
149#[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 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 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#[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#[derive(Debug)]
250pub struct ConfigFault {
251 pub config_id: u64,
253 pub fault: Fault,
255}
256
257#[derive(Debug)]
259pub enum ProbeError {
260 NoDevice,
262 Fault(Fault),
264 EveryConfigFaulted(Vec<ConfigFault>),
266 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#[derive(Debug, Clone, PartialEq, Eq)]
303pub struct RefusedConfig {
304 pub config_id: u64,
306 pub reasons: Vec<String>,
308}
309
310#[derive(Debug, Clone, PartialEq, Eq)]
312pub struct CapacityError {
313 pub refused: Vec<RefusedConfig>,
315 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
337pub(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 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 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 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}