1#![cfg_attr(docsrs, feature(doc_cfg))]
49#![warn(missing_docs)]
50#![expect(
51 clippy::print_stdout,
52 clippy::print_stderr,
53 reason = "henad-cli is a command line: stdout carries its results, stderr its progress log"
54)]
55
56use std::ffi::OsString;
57use std::fs::File;
58use std::io::{BufWriter, Write as _};
59use std::ops::ControlFlow;
60use std::path::{Path, PathBuf};
61use std::time::Duration;
62
63use anyhow::{Context as _, Result, anyhow, bail};
64use clap::{ArgGroup, CommandFactory as _, FromArgMatches as _, Parser};
65
66use henad_compute::entry::{ModelEntry, ModelLookupError, ModelSet};
67use henad_compute::fault::install_panic_hook;
68use henad_compute::gpu::GpuContext;
69use henad_compute::runtime_info::{GpuVerdict, HostInfo, RuntimeInfo, classify_adapter};
70use henad_compute::simulation::{RunSetup, Simulation};
71use henad_core::action::Schedule;
72use henad_core::explore::value::{ValueError, parse_overrides, resolve_params};
73use henad_core::export::StatsWriter;
74use henad_core::metadata::Backend;
75use henad_core::params::{ParamFormat, ParamKind};
76use henad_core::provenance::BuildInfo;
77use henad_explore::benchmark::{BenchmarkEvent, BenchmarkSettings, run_benchmark};
78use henad_explore::device::acquire_headless;
79use henad_explore::spec_file::{LoadedSpec, SpecFileError};
80use henad_explore::sweep::Provenance;
81
82use crate::explore::ExploreArgs;
83use numfmt::{Formatter, Scales};
84
85mod explore;
86mod json_report;
87
88#[derive(Debug, Clone)]
91pub struct CliOptions {
92 models: ModelSet,
93 host: BuildInfo,
95 command_name: String,
97 about: Option<String>,
99}
100
101impl CliOptions {
102 pub fn new(models: ModelSet, host: BuildInfo) -> Self {
107 Self {
108 command_name: host.package().to_owned(),
109 models,
110 host,
111 about: None,
112 }
113 }
114
115 pub fn command_name(mut self, name: impl Into<String>) -> Self {
119 self.command_name = name.into();
120 self
121 }
122
123 pub fn about(mut self, text: impl Into<String>) -> Self {
125 self.about = Some(text.into());
126 self
127 }
128}
129
130pub const SOME_RUNS_NOT_OK: u8 = 3;
133
134pub fn run(options: CliOptions, arguments: impl IntoIterator<Item = OsString>) -> u8 {
143 install_panic_hook();
144 let CliOptions {
145 models,
146 host,
147 command_name,
148 about,
149 } = options;
150 let arguments: Vec<OsString> = arguments.into_iter().collect();
151 let args = match parse_args(command_name, host.version(), about, &arguments) {
152 Ok(args) => args,
153 Err(error) => {
154 error.print().ok();
156 return u8::try_from(error.exit_code()).unwrap_or(1);
157 }
158 };
159 match run_args(&models, &host, &args, &arguments) {
160 Ok(code) => code,
161 Err(error) => {
162 eprintln!("Error: {error:?}");
163 1
164 }
165 }
166}
167
168fn parse_args(
174 name: String,
175 version: &'static str,
176 about: Option<String>,
177 arguments: &[OsString],
178) -> Result<Args, clap::Error> {
179 let mut command = Args::command().name(name.clone()).bin_name(name).version(version);
181 if let Some(about) = about {
182 command = command.about(about);
183 }
184 let mut matches = command.try_get_matches_from_mut(arguments)?;
185 Args::from_arg_matches_mut(&mut matches).map_err(|error| error.format(&mut command))
186}
187
188#[derive(Parser)]
190#[command(about)]
191#[command(group(
192 ArgGroup::new("explore")
193 .args(["out", "spec", "dry_run"])
194 .multiple(true)
195 .conflicts_with_all(["list", "export", "export_stats", "global_warmup"])
196))]
197#[command(group(ArgGroup::new("spec_use").args(["out", "dry_run", "params"]).multiple(true)))]
199struct Args {
200 #[arg(required_unless_present_any = ["list", "info", "spec", "merge"])]
202 model: Option<String>,
203
204 #[arg(long, default_value_t = 1000)]
206 steps: u64,
207
208 #[arg(long, default_value_t = 0)]
210 warmup: u64,
211
212 #[arg(long = "global-warmup", default_value_t = 0)]
215 global_warmup: u64,
216
217 #[arg(long)]
219 seed: Option<u64>,
220
221 #[arg(long, default_value_t = 1, value_parser = clap::value_parser!(u64).range(1..))]
224 reps: u64,
225
226 #[arg(long = "set", value_name = "ID=VALUE", value_parser = check_set)]
228 set: Vec<String>,
229
230 #[arg(long = "act", value_name = "ID@TICK", value_parser = check_act)]
235 act: Vec<String>,
236
237 #[arg(long, value_name = "PATH")]
239 export: Option<PathBuf>,
240
241 #[arg(long = "export-stats", value_name = "PATH")]
244 export_stats: Option<PathBuf>,
245
246 #[arg(long = "stats-every", default_value_t = 1, value_name = "N")]
248 stats_every: u64,
249
250 #[arg(long)]
252 list: bool,
253
254 #[arg(long, conflicts_with_all = ["out", "dry_run"])]
257 params: bool,
258
259 #[arg(long)]
262 info: bool,
263
264 #[arg(long)]
266 json: bool,
267
268 #[arg(long, default_value_t = 0, value_name = "N")]
270 threads: usize,
271
272 #[command(flatten)]
273 explore: ExploreArgs,
274}
275
276#[derive(Debug, Clone, Copy, PartialEq, Eq)]
278enum Mode {
279 List,
280 Merge,
282 InfoOnly,
284 Params,
285 Explore,
286 ExportStats,
287 ExportFinal,
288 Benchmark,
289}
290
291impl Mode {
292 fn of(args: &Args) -> Self {
294 if args.list {
295 Self::List
296 } else if !args.explore.merge.is_empty() {
297 Self::Merge
298 } else if args.info && args.model.is_none() && !args.explore.is_sweep() {
299 Self::InfoOnly
300 } else if args.params {
301 Self::Params
302 } else if args.explore.is_sweep() {
303 Self::Explore
304 } else if args.export_stats.is_some() {
305 Self::ExportStats
306 } else if args.export.is_some() {
307 Self::ExportFinal
308 } else {
309 Self::Benchmark
310 }
311 }
312}
313
314fn run_args(models: &ModelSet, host: &BuildInfo, args: &Args, arguments: &[OsString]) -> Result<u8> {
322 let mode = Mode::of(args);
323
324 if args.threads > 0 {
327 rayon::ThreadPoolBuilder::new()
328 .num_threads(args.threads)
329 .build_global()
330 .context("cannot size the worker pool")?;
331 }
332 if mode == Mode::Merge {
333 return explore::merge_shards(args);
334 }
335
336 let gpu_models = has_gpu_models(models);
339 let gpu_ctx = if gpu_models || args.info {
340 match acquire_headless(models.gpu_needs()) {
341 Ok(ctx) => Some(ctx),
342 Err(err) => {
343 if gpu_models {
344 eprintln!("note: no GPU available ({err}); GPU models disabled");
345 }
346 None
347 }
348 }
349 } else {
350 None
351 };
352 let runtime = gpu_ctx.as_ref().and_then(GpuContext::runtime_info);
353
354 if let Some(runtime) = runtime
357 && gpu_models
358 && classify_adapter(&runtime.adapter) == GpuVerdict::Absent
359 {
360 eprintln!(
361 "!!! warning: adapter '{}' is a software rasteriser, not a GPU; \
362 GPU-model results from this machine are not GPU results !!!",
363 runtime.adapter.name
364 );
365 }
366
367 if args.info {
368 if args.json {
369 json_report::runtime(runtime);
370 } else {
371 print_runtime_info(runtime, gpu_models);
372 }
373 }
374
375 match mode {
376 Mode::List => {
377 print_models(models.runnable(gpu_ctx.as_ref()));
378 return Ok(0);
379 }
380 Mode::InfoOnly => return Ok(0),
381 _ => {}
382 }
383
384 let spec = args.explore.spec.as_deref().map(load_spec).transpose()?;
385 let model_id = match (args.model.as_deref(), spec.as_ref()) {
386 (Some(id), Some(loaded)) if id != loaded.spec.model => {
387 bail!("model '{id}' does not match the spec's model '{}'", loaded.spec.model)
388 }
389 (Some(id), _) => id,
390 (None, Some(loaded)) => &loaded.spec.model,
391 (None, None) => bail!("a model id is required (try --list)"),
392 };
393 let entry = models.lookup(model_id, gpu_ctx.as_ref()).map_err(|error| match error {
394 ModelLookupError::NotInSet { .. } => anyhow!("{error} (try --list)"),
395 _ => anyhow!(error),
396 })?;
397
398 match mode {
399 Mode::Params if args.json => json_report::emit(&json_report::params(entry, gpu_ctx.as_ref())),
400 Mode::Params => print!("{}", params_text(entry)),
401 Mode::Explore => {
402 let provenance = provenance(host, arguments);
403 return explore::run(args, entry, gpu_ctx.as_ref(), spec, provenance);
404 }
405 _ => return run_single(entry, args, mode, gpu_ctx.as_ref(), runtime).map(|()| 0),
406 }
407 Ok(0)
408}
409
410fn has_gpu_models(models: &ModelSet) -> bool {
412 models.iter().any(|entry| entry.gpu_needs().is_some())
413}
414
415fn provenance(host: &BuildInfo, arguments: &[OsString]) -> Provenance {
417 let arguments = arguments.iter().map(|arg| arg.to_string_lossy().into_owned()).collect();
418 Provenance::new(*host, arguments)
419}
420
421fn load_spec(path: &Path) -> Result<LoadedSpec> {
427 LoadedSpec::read(path).map_err(|error| match error {
428 SpecFileError::Read { .. } => anyhow::Error::new(error),
429 other => anyhow::Error::new(other).context(format!("cannot read '{}'", path.display())),
430 })
431}
432
433fn run_single(
435 entry: &ModelEntry,
436 args: &Args,
437 mode: Mode,
438 gpu_ctx: Option<&GpuContext>,
439 runtime: Option<&RuntimeInfo>,
440) -> Result<()> {
441 let overrides = parse_overrides(&args.set)?;
442 let params = resolve_params(entry.param_descriptors(), &overrides).map_err(set_error)?;
443 let schedule = Schedule::parse(&args.act, entry.id(), entry.action_descriptors())?;
444 if let Some(last) = schedule.last_tick()
445 && last > args.warmup + args.steps
446 {
447 eprintln!(
448 "note: --act at tick {last} is past the {} this run reaches, so it never fires",
449 args.warmup + args.steps
450 );
451 }
452
453 if let Some(ctx) = gpu_ctx {
456 let shortfalls = entry.shortfalls(¶ms, &ctx.device.limits());
457 if !shortfalls.is_empty() {
458 bail!("'{}' does not fit this device: {}", entry.id(), shortfalls.join("; "));
459 }
460 }
461 let setup = RunSetup::from_parts(entry, ¶ms, args.seed, schedule)?;
462
463 match (mode, &args.export_stats, &args.export) {
464 (Mode::ExportStats, Some(path), _) => export_stats(&setup, args, path, gpu_ctx),
465 (Mode::ExportFinal, _, Some(path)) => export_final(&setup, args, path),
466 _ => {
467 let adapter = runtime.map(|r| r.adapter.name.as_str());
468 benchmark(setup, args, gpu_ctx, adapter)
469 }
470 }
471}
472
473fn print_runtime_info(runtime: Option<&RuntimeInfo>, gpu_models: bool) {
477 let collected;
478 let host = if let Some(runtime) = runtime {
479 &runtime.host
480 } else {
481 collected = HostInfo::collect();
482 &collected
483 };
484
485 let fmt_opt = |value: Option<usize>| value.map_or_else(|| "unknown".to_owned(), |n| n.to_string());
486
487 println!("runtime info:");
488 println!(" host:");
489 println!(" platform: {} ({})", host.os, host.arch);
490 println!(" logical cpus: {}", fmt_opt(host.logical_cpus));
491 println!(" worker threads: {}", fmt_opt(host.worker_threads));
492
493 match runtime {
494 None if gpu_models => println!(" gpu: none (GPU models disabled)"),
495 None => println!(" gpu: none"),
496 Some(runtime) => {
497 let adapter = &runtime.adapter;
498 println!(" gpu:");
499 println!(" adapter: {}", adapter.name);
500 println!(" type: {:?}", adapter.device_type);
501 println!(" backend: {}", adapter.backend);
502 if !adapter.driver_info.is_empty() {
503 println!(" driver: {}", adapter.driver_info);
504 }
505 let limits = &runtime.granted;
506 println!(
507 " max storage binding: {} bytes ({} u32 cells)",
508 limits.max_storage_buffer_binding_size,
509 limits.max_storage_buffer_binding_size / 4
510 );
511 println!(" max buffer size: {} bytes", limits.max_buffer_size);
512 println!(" max 2d texture: {}", limits.max_texture_dimension_2d);
513 println!(
514 " storage buffers: {} per shader stage",
515 limits.max_storage_buffers_per_shader_stage
516 );
517 println!(" display texture cap: {0}x{0}", runtime.display_cap());
518 }
519 }
520}
521
522fn print_models<'a>(entries: impl Iterator<Item = &'a ModelEntry>) {
524 println!("available models:");
525 for entry in entries {
526 let (id, name) = (entry.id(), entry.name());
527 println!(" {id:<18} {name}");
528 }
529}
530
531fn params_text(entry: &ModelEntry) -> String {
536 let mut text = format!("parameters for {} ({}):\n", entry.id(), entry.name());
537 for (index, desc) in entry.param_descriptors().iter().enumerate() {
538 let (id, label) = (desc.id, desc.label);
539 let apply = if desc.is_live() { "live" } else { "reload" };
540 let kind = match &desc.kind {
541 ParamKind::F32 { min, max, default, .. } => format!("kind=f32 default={default} min={min} max={max}"),
542 ParamKind::U32 { min, max, default } => format!("kind=u32 default={default} min={min} max={max}"),
543 ParamKind::Bool { default } => format!("kind=bool default={default}"),
544 ParamKind::Choice { options, default } => {
545 format!("kind=choice default={default} options={}", options.join("|"))
546 }
547 };
548 let format = match desc.format {
550 ParamFormat::Plain => "",
551 ParamFormat::Percent => " format=percent",
552 };
553 text.push_str(&format!(
554 " index={index} id={id} {kind} apply={apply}{format} label=\"{label}\"\n"
555 ));
556 }
557 text
558}
559
560fn benchmark(setup: RunSetup, args: &Args, gpu_ctx: Option<&GpuContext>, adapter: Option<&str>) -> Result<()> {
565 let entry = setup.entry().clone();
566 let mut settings = BenchmarkSettings::new(setup, args.steps);
567 settings.warmup = args.warmup;
568 settings.global_warmup = args.global_warmup;
569 settings.repetitions = args.reps;
570 let mut backend = Backend::Cpu;
571 let mut number_open = false;
573 let finished = run_benchmark(&settings, gpu_ctx, &mut |event| {
574 number_open = matches!(
575 event,
576 BenchmarkEvent::GlobalWarmupStarted | BenchmarkEvent::RepetitionStarted(_)
577 );
578 match event {
579 BenchmarkEvent::Started {
580 backend: started,
581 parallel_jobs,
582 } => {
583 backend = started;
584 let (variant, adapter, tag) = match started {
585 Backend::Gpu => ("gpu", adapter, " [GPU]"),
586 Backend::Cpu => ("cpu", None, ""),
587 };
588 if args.json {
589 json_report::info(
590 entry.id(),
591 variant,
592 rayon::current_num_threads(),
593 parallel_jobs,
594 adapter,
595 );
596 }
597 eprintln!(
598 "benchmarking {} ({}){tag}: {} steps x {} reps, {} warmup, {} global-warmup",
599 entry.name(),
600 entry.id(),
601 args.steps,
602 args.reps,
603 args.warmup,
604 args.global_warmup
605 );
606 if cfg!(debug_assertions) {
607 eprintln!("!!! warning: debug build; use --release for benchmarking !!!");
608 }
609 }
610 BenchmarkEvent::GlobalWarmupStarted => eprint!(" #{: >4}: ", 0),
612 BenchmarkEvent::GlobalWarmupFinished(elapsed) => {
613 eprintln!("{elapsed:>8.3?} ({} global warmup steps)", args.global_warmup);
614 }
615 BenchmarkEvent::RepetitionStarted(index) => eprint!(" #{: >4}: ", index + 1),
616 BenchmarkEvent::RepetitionFinished(repetition) => {
617 eprintln!("{:>8.3?}", repetition.elapsed);
618 let (population, after) = (repetition.population_after_warmup, repetition.population_after_steps);
619 let noise = (3.0 * (population as f64).sqrt()) as u64;
621 if backend == Backend::Cpu && after.abs_diff(population) > (population / 10).max(noise) {
622 eprintln!(
623 " note: the population went from {population} to {after} during the timed steps. \
624 Updates per second are computed from the population after warmup. \
625 A longer --warmup can reach a steady population first."
626 );
627 }
628 if args.json {
629 json_report::rep(
630 repetition.index,
631 repetition.seed,
632 args.steps,
633 args.warmup,
634 repetition.elapsed,
635 population,
636 repetition.heap_bytes,
637 );
638 }
639 }
640 _ => {}
642 }
643 });
644 if finished.is_err() && number_open {
647 eprintln!();
648 }
649 let report = finished?;
650
651 let samples: Vec<Duration> = report.repetitions.iter().map(|repetition| repetition.elapsed).collect();
652 let population = report
655 .repetitions
656 .last()
657 .map_or(0, |repetition| repetition.population_after_warmup);
658 let setup = &settings.setup;
659 if args.json {
660 json_report::summary(
661 &samples,
662 args.steps,
663 population,
664 report.grid_size,
665 entry.param_descriptors(),
666 setup.values(),
667 setup.schedule(),
668 );
669 } else {
670 print_report(&samples, args.steps, population, report.grid_size)?;
671 }
672 Ok(())
673}
674
675fn median_of(sorted: &[Duration]) -> Duration {
679 match sorted.len() {
680 0 => Duration::ZERO,
681 n if n % 2 == 1 => sorted[n / 2],
682 n => (sorted[n / 2 - 1] + sorted[n / 2]) / 2,
683 }
684}
685
686fn print_report(
692 samples: &[Duration],
693 steps_per_rep: u64,
694 population: u64,
695 grid_dims: Option<(u32, u32)>,
696) -> Result<()> {
697 println!("benchmark result:");
698 let min = samples.iter().min().copied().unwrap_or_default();
700 let max = samples.iter().max().copied().unwrap_or_default();
701 let mean = samples.iter().sum::<Duration>() / (samples.len() as u32);
702 let median = {
703 let mut sorted = samples.to_vec();
704 sorted.sort_unstable();
705 median_of(&sorted)
706 };
707 let std_dev = {
708 let mean_secs = mean.as_secs_f64();
709 let variance = samples
710 .iter()
711 .map(|s| {
712 let diff = s.as_secs_f64() - mean_secs;
713 diff * diff
714 })
715 .sum::<f64>()
716 / (samples.len() as f64);
717 Duration::from_secs_f64(variance.sqrt())
718 };
719
720 let mut f = Formatter::new()
721 .scales(Scales::none())
722 .separator(' ')?
723 .precision(numfmt::Precision::Decimals(3));
724
725 println!(" min: {min:>10.3?}");
726 println!(" median: {median:>10.3?}");
727 println!(" max: {max:>10.3?}");
728 println!(" mean: {mean:>10.3?}");
729 println!(" std dev: {std_dev:>10.3?}");
730 let mean_steps_per_sec = f.fmt2(steps_per_rep as f64 / mean.as_secs_f64());
731 println!(" > mean steps/sec: {mean_steps_per_sec:>20}");
732 let mean_updates_per_sec = f.fmt2((steps_per_rep as f64 * population as f64) / mean.as_secs_f64());
733 println!(" > mean updates/sec: {mean_updates_per_sec:>20}");
734 if let Some((w, h)) = grid_dims {
735 f = f.precision(numfmt::Precision::Decimals(0));
736 let grid_size = f.fmt2(w as u64 * h as u64);
737 println!(" > grid size: {grid_size:>16}");
738 }
739 Ok(())
740}
741
742fn export_final(setup: &RunSetup, args: &Args, path: &Path) -> Result<()> {
746 let entry = setup.entry();
747 if entry.gpu_needs().is_some() {
748 bail!("model '{}' is GPU-backed; this path is CPU-only", entry.id());
749 }
750 let topology = entry.topology_hint();
751 if !topology.grid && !topology.agents {
752 bail!("model exposes no CPU-side view to export");
753 }
754 let mut simulation = setup.build(None)?;
755 simulation.run_to(args.warmup + args.steps)?;
756 let file = File::create(path).with_context(|| format!("cannot create '{}'", path.display()))?;
757 let mut out = BufWriter::new(file);
758 simulation.write_state(&mut out)?;
759 out.flush()?;
760 eprintln!(
761 "exported final state (tick {}) to {}",
762 simulation.tick(),
763 path.display()
764 );
765 Ok(())
766}
767
768fn export_stats(setup: &RunSetup, args: &Args, path: &Path, gpu_ctx: Option<&GpuContext>) -> Result<()> {
774 if args.stats_every == 0 {
775 bail!("--stats-every must be at least 1");
776 }
777 if args.reps > 1 {
778 eprintln!("note: --reps is ignored by --export-stats");
779 }
780
781 let total = args.warmup + args.steps;
782 let file = File::create(path).with_context(|| format!("cannot create '{}'", path.display()))?;
783 let writer = StatsWriter::new(BufWriter::new(file));
784
785 let entry = setup.entry();
786 eprintln!(
787 "exporting stats for {} ({}): {} steps, sampling every {}",
788 entry.name(),
789 entry.id(),
790 total,
791 args.stats_every
792 );
793
794 let mut simulation = setup.build(gpu_ctx)?;
795 let rows = write_series(&mut simulation, total, args.stats_every, writer)?;
796 eprintln!("wrote {rows} rows to {}", path.display());
797 Ok(())
798}
799
800fn write_series<W: std::io::Write + Send>(
806 simulation: &mut Simulation,
807 total: u64,
808 every: u64,
809 mut writer: StatsWriter<W>,
810) -> Result<u64> {
811 let flow = simulation.run_sampled(total, every, |sample| {
812 match writer.push(sample.tick(), sample.entries()) {
813 Ok(()) => ControlFlow::Continue(()),
814 Err(error) => ControlFlow::Break(error),
815 }
816 })?;
817 if let ControlFlow::Break(error) = flow {
818 return Err(error.into());
819 }
820 Ok(writer.finish()?)
821}
822
823#[cfg(test)]
825fn test_provenance() -> Provenance {
826 Provenance::new(henad_core::build_info!(), Vec::new())
827}
828
829fn set_error(error: ValueError) -> anyhow::Error {
831 match error {
832 ValueError::Param { id, source } => anyhow::Error::new(*source).context(format!("--set {id}")),
833 other => other.into(),
834 }
835}
836
837fn check_set(raw: &str) -> Result<String, String> {
841 if raw.contains('=') {
842 Ok(raw.to_owned())
843 } else {
844 Err("expected ID=VALUE".to_owned())
845 }
846}
847
848fn check_act(raw: &str) -> Result<String, String> {
852 let (_, tick) = raw.rsplit_once('@').ok_or("expected ID@TICK")?;
853 match tick.parse::<u64>() {
854 Ok(_) => Ok(raw.to_owned()),
855 Err(_) => Err(format!("expected a tick from 0 to {}, got '{tick}'", u64::MAX)),
856 }
857}
858
859#[cfg(test)]
860mod tests {
861 use std::ffi::OsString;
862 use std::path::{Path, PathBuf};
863
864 use super::{
865 Args, CliOptions, Mode, SOME_RUNS_NOT_OK, has_gpu_models, params_text, parse_args, run, test_provenance,
866 write_series,
867 };
868 use crate::explore::{self, ExploreArgs};
869 use crate::json_report;
870 use clap::Parser as _;
871 use henad_compute::entry::{ModelEntry, ModelSet};
872 use henad_compute::gpu::GpuContext;
873 use henad_compute::simulation::{RunSetup, Simulation};
874 use henad_core::action::Schedule;
875 use henad_core::explore::seed::run_seed;
876 use henad_core::explore::value::{parse_overrides, resolve_params};
877 use henad_core::export::csv::parse_records;
878 use henad_core::export::stats_csv::StatsWriter;
879 use henad_core::params::ParamValue;
880 use henad_explore::testing::{TestDeviceRequest, headless_test_device};
881 use henad_models::example_models;
882 use serde_json::json;
883
884 struct ScratchDir {
886 path: PathBuf,
887 }
888
889 impl ScratchDir {
890 fn new(name: &str) -> Self {
891 let path = std::env::temp_dir().join(format!("henad-cli-{name}-{}", std::process::id()));
892 if path.exists() {
893 std::fs::remove_dir_all(&path).expect("an earlier run's directory can be removed");
894 }
895 Self { path }
896 }
897
898 fn arg(&self) -> &str {
899 self.path.to_str().expect("the temporary directory is UTF-8")
900 }
901
902 fn read(&self, file: &str) -> String {
903 std::fs::read_to_string(self.path.join(file)).expect("the sweep wrote the file")
904 }
905 }
906
907 impl Drop for ScratchDir {
908 fn drop(&mut self) {
909 std::fs::remove_dir_all(&self.path).ok();
910 }
911 }
912
913 fn cpu_entry(id: &str) -> ModelEntry {
914 example_models().get(id).cloned().expect("the model is registered")
915 }
916
917 fn series_without_run_id(dir: &ScratchDir) -> Vec<String> {
919 dir.read("series.csv")
920 .lines()
921 .map(|line| line.split_once(',').expect("a run_id column").1.to_owned())
922 .collect()
923 }
924
925 fn assert_final_reducers(dir: &ScratchDir, header: &str, last_row: &str) {
929 let runs = parse_records(&dir.read("runs.csv")).expect("runs.csv is CSV");
930 assert_eq!(runs.len(), 2, "a header and one run");
931 for (column, value) in header.split(',').zip(last_row.split(',')).skip(1) {
932 let final_column = runs[0]
933 .iter()
934 .position(|name| *name == format!("{column}:final"))
935 .expect("every stat column has a final reducer");
936 let reduced: f64 = runs[1][final_column].parse().expect("a number");
937 let exported: f64 = value.parse().expect("a number");
938 assert_eq!(reduced, exported, "{column}");
939 }
940 }
941
942 #[test]
946 fn a_single_point_sweep_matches_export_stats() {
947 let entry = cpu_entry("sir");
948 let dir = ScratchDir::new("single-point");
949 let args = Args::parse_from([
950 "henad-cli",
951 "sir",
952 "--set",
953 "grid_width=32",
954 "--set",
955 "grid_height=32",
956 "--steps",
957 "31",
958 "--stats-every",
959 "3",
960 "--seed",
961 "7",
962 "--out",
963 dir.arg(),
964 ]);
965 let status = explore::run(&args, &entry, None, None, test_provenance()).expect("the sweep runs");
966 assert_eq!(status, 0);
967
968 let overrides = parse_overrides(&args.set).expect("valid");
969 let params = resolve_params(entry.param_descriptors(), &overrides).expect("in range");
970 let mut simulation = cpu_simulation(&entry, ¶ms, run_seed(7, 0), &[]);
971 let mut exported_bytes = Vec::new();
972 write_series(&mut simulation, 31, 3, StatsWriter::new(&mut exported_bytes)).expect("writes");
973 let exported = String::from_utf8(exported_bytes).expect("utf-8");
974 let exported: Vec<&str> = exported.lines().collect();
975
976 assert_eq!(series_without_run_id(&dir), exported);
977 assert_eq!(exported.len(), 1 + 12, "a header, ticks 0 to 30 every 3, and tick 31");
978 assert_final_reducers(&dir, exported[0], exported[exported.len() - 1]);
979 }
980
981 #[test]
983 fn a_sweep_fires_its_actions_as_export_stats_does() {
984 let entry = cpu_entry("sir");
985 let dir = ScratchDir::new("single-point-actions");
986 let args = Args::parse_from([
987 "henad-cli",
988 "sir",
989 "--set",
990 "grid_width=32",
991 "--set",
992 "grid_height=32",
993 "--steps",
994 "20",
995 "--stats-every",
996 "2",
997 "--seed",
998 "7",
999 "--act",
1000 "seed_outbreak@0",
1001 "--act",
1002 "seed_outbreak@5",
1003 "--act",
1004 "seed_outbreak@20",
1005 "--out",
1006 dir.arg(),
1007 ]);
1008 let status = explore::run(&args, &entry, None, None, test_provenance()).expect("the sweep runs");
1009 assert_eq!(status, 0);
1010
1011 let overrides = parse_overrides(&args.set).expect("valid");
1012 let params = resolve_params(entry.param_descriptors(), &overrides).expect("in range");
1013 let mut simulation = cpu_simulation(&entry, ¶ms, run_seed(7, 0), &args.act);
1014 let mut exported_bytes = Vec::new();
1015 write_series(&mut simulation, 20, 2, StatsWriter::new(&mut exported_bytes)).expect("writes");
1016 let exported = String::from_utf8(exported_bytes).expect("utf-8");
1017 assert_eq!(series_without_run_id(&dir), exported.lines().collect::<Vec<_>>());
1018 }
1019
1020 fn runs_without_timing(dir: &Path) -> Vec<Vec<String>> {
1022 let text = std::fs::read_to_string(dir.join("runs.csv")).expect("the sweep wrote runs.csv");
1023 let mut records = parse_records(&text).expect("runs.csv is CSV");
1024 let timing: Vec<usize> = records[0]
1025 .iter()
1026 .enumerate()
1027 .filter(|(_, name)| matches!(name.as_str(), "build_ms" | "wall_ms" | "steps_per_s"))
1028 .map(|(column, _)| column)
1029 .collect();
1030 assert_eq!(timing.len(), 3, "runs.csv times each run in three columns");
1031 for record in &mut records {
1032 for &column in timing.iter().rev() {
1033 record.remove(column);
1034 }
1035 }
1036 records
1037 }
1038
1039 #[test]
1044 fn sharded_sweeps_merge_into_the_unsharded_files() {
1045 let entry = cpu_entry("sir");
1046 let dir = ScratchDir::new("shards");
1047 std::fs::create_dir_all(&dir.path).expect("the scratch directory can be made");
1048 let sweep = |dest: &Path, shard: Option<&str>| {
1049 let dest = dest.to_str().expect("the temporary directory is UTF-8");
1050 let mut line = vec![
1051 "henad-cli",
1052 "sir",
1053 "--set",
1054 "grid_width=16",
1055 "--set",
1056 "grid_height=16",
1057 "--set",
1058 "initial_infected_pct=0.02",
1059 "--set",
1060 "recovery_rate=0.3",
1061 "--steps",
1062 "24",
1063 "--stats-every",
1064 "3",
1065 "--series-every",
1066 "6",
1067 "--reps",
1068 "2",
1069 "--seed",
1070 "11",
1071 "--vary",
1072 "infection_rate=0.05:0.6",
1073 "--act",
1074 "seed_outbreak@6",
1075 "--vary",
1076 "action.seed_outbreak=2:20",
1077 "--sample",
1078 "lhs:3",
1079 "--stop",
1080 "Infected <= 0",
1081 "--reduce",
1082 "Infected:argmax",
1083 "--reduce",
1084 "Infected:first<=1",
1085 "--out",
1086 dest,
1087 ];
1088 line.extend(shard.map(|shard| ["--shard", shard]).into_iter().flatten());
1089 let args = Args::parse_from(line);
1090 explore::run(&args, &entry, None, None, test_provenance()).expect("the sweep runs")
1091 };
1092 let whole = dir.path.join("whole");
1093 assert_eq!(sweep(&whole, None), 0);
1094 let shards = [dir.path.join("shard-0"), dir.path.join("shard-1")];
1095 for (shard_dir, shard) in shards.iter().zip(["0/2", "1/2"]) {
1096 assert_eq!(sweep(shard_dir, Some(shard)), 0, "shard {shard}");
1097 assert_eq!(
1098 runs_without_timing(shard_dir).len(),
1099 1 + 3,
1100 "a header and every other run of 6"
1101 );
1102 }
1103
1104 let path = |dir: &Path| dir.to_str().expect("the temporary directory is UTF-8").to_owned();
1105 let merged = dir.path.join("merged");
1106 let args = Args::parse_from([
1107 "henad-cli".to_owned(),
1108 "--merge".to_owned(),
1109 path(&shards[1]),
1110 path(&shards[0]),
1111 "--out".to_owned(),
1112 path(&merged),
1113 ]);
1114 assert_eq!(Mode::of(&args), Mode::Merge);
1115 assert_eq!(explore::merge_shards(&args).expect("the shards merge"), 0);
1116 assert_eq!(runs_without_timing(&merged), runs_without_timing(&whole));
1117 for file in ["series.csv", "summary.csv"] {
1118 let read = |dir: &Path| std::fs::read_to_string(dir.join(file)).expect("the table is written");
1119 assert_eq!(read(&merged), read(&whole), "{file}");
1120 }
1121
1122 let partial = dir.path.join("partial");
1123 let args = Args::parse_from([
1124 "henad-cli".to_owned(),
1125 "--merge".to_owned(),
1126 path(&shards[0]),
1127 "--out".to_owned(),
1128 path(&partial),
1129 ]);
1130 let status = explore::merge_shards(&args).expect("one shard merges");
1131 assert_eq!(status, SOME_RUNS_NOT_OK, "half the runs are missing");
1132 }
1133
1134 #[test]
1136 fn run_returns_the_exit_code_of_each_failure() {
1137 let options = || CliOptions::new(ModelSet::new(henad_core::build_info!()), henad_core::build_info!());
1138 let line = |arguments: &[&str]| arguments.iter().map(OsString::from).collect::<Vec<_>>();
1139 assert_eq!(run(options(), line(&["henad-cli", "--no-such-flag"])), 2);
1140 assert_eq!(
1141 run(options(), line(&["henad-cli", "--merge", "a"])),
1142 2,
1143 "--merge needs --out"
1144 );
1145 assert_eq!(
1146 run(options(), line(&["henad-cli", "--spec", "s.toml"])),
1147 2,
1148 "a spec needs --out, --dry-run or --params"
1149 );
1150 assert_eq!(
1151 run(options(), line(&["henad-cli", "sir", "--set", "grid_width"])),
1152 2,
1153 "a malformed --set is refused before the model is looked up"
1154 );
1155 assert_eq!(
1156 run(options(), line(&["henad-cli", "sir", "--steps", "1"])),
1157 1,
1158 "an empty set holds no sir"
1159 );
1160 }
1161
1162 #[test]
1165 fn the_options_name_and_describe_the_command() {
1166 let render = |about: Option<&str>, arguments: &[&str]| {
1167 let arguments: Vec<OsString> = arguments.iter().map(OsString::from).collect();
1168 let about = about.map(str::to_owned);
1169 let Err(error) = parse_args("my-models".to_owned(), "1.2.3", about, &arguments) else {
1170 panic!("{arguments:?} ends the parse");
1171 };
1172 error.render().to_string()
1173 };
1174 let refused = render(None, &["other-name", "--no-such-flag"]);
1175 assert!(refused.contains("Usage: my-models"), "{refused}");
1176 let help = render(Some("Runs my models."), &["other-name", "--help"]);
1177 assert!(help.starts_with("Runs my models."), "{help}");
1178 assert!(help.contains("Usage: my-models"), "{help}");
1179 let default = render(None, &["other-name", "--help"]);
1180 assert!(default.starts_with(env!("CARGO_PKG_DESCRIPTION")), "{default}");
1181 assert_eq!(render(None, &["other-name", "--version"]), "my-models 1.2.3\n");
1182 }
1183
1184 #[test]
1186 fn a_cpu_only_set_has_no_gpu_models() {
1187 let mut cpu_only = ModelSet::new(henad_core::build_info!());
1188 cpu_only.insert(cpu_entry("sir")).expect("one id");
1189 assert!(!has_gpu_models(&cpu_only));
1190 assert!(has_gpu_models(&example_models()));
1191 }
1192
1193 #[test]
1195 fn params_text_output_is_unchanged() {
1196 let expected = "\
1197parameters for virus_network (Virus on a Network):
1198 index=0 id=num_agents kind=u32 default=10000 min=1 max=10000000 apply=reload label=\"Number of Nodes\"
1199 index=1 id=world_width kind=f32 default=1000 min=1 max=10000 apply=reload label=\"World Width\"
1200 index=2 id=world_height kind=f32 default=1000 min=1 max=10000 apply=reload label=\"World Height\"
1201 index=3 id=average_node_degree kind=u32 default=6 min=1 max=20 apply=reload label=\"Average Node Degree\"
1202 index=4 id=initial_outbreak_size kind=u32 default=3 min=1 max=10000 apply=reload label=\"Initial Outbreak Size\"
1203 index=5 id=virus_spread_chance kind=f32 default=0.025 min=0 max=1 apply=live format=percent label=\"Virus Spread Chance\"
1204 index=6 id=virus_check_frequency kind=u32 default=1 min=1 max=20 apply=live label=\"Virus Check Frequency\"
1205 index=7 id=recovery_chance kind=f32 default=0.05 min=0 max=1 apply=live format=percent label=\"Recovery Chance\"
1206 index=8 id=gain_resistance_chance kind=f32 default=0.05 min=0 max=1 apply=live format=percent label=\"Gain Resistance Chance\"
1207 index=9 id=directed kind=bool default=false apply=live label=\"Directed\"
1208 index=10 id=network kind=choice default=0 options=Random|Geometric apply=reload label=\"Network\"
1209 index=11 id=keep_rewiring kind=bool default=false apply=live label=\"Keep Rewiring\"
1210";
1211 assert_eq!(params_text(&cpu_entry("virus_network")), expected);
1212 }
1213
1214 #[test]
1215 fn params_json_lists_every_descriptor() {
1216 let models = example_models();
1217 for entry in models.runnable(None) {
1218 let line = json_report::params(entry, None);
1219 assert_eq!(line["kind"], json!("params"), "{}", entry.id());
1220 assert_eq!(line["model"], json!(entry.id()));
1221 let params = line["params"].as_array().expect("params is a list");
1222 let ids: Vec<&str> = params.iter().filter_map(|param| param["id"].as_str()).collect();
1223 let declared: Vec<&str> = entry.param_descriptors().iter().map(|d| d.id).collect();
1224 assert_eq!(ids, declared, "{}", entry.id());
1225 for param in params {
1226 assert!(param["kind"].is_string() && !param["default"].is_null(), "{param}");
1227 }
1228 let columns = line["stat_columns"]
1229 .as_array()
1230 .expect("the model builds at its defaults");
1231 assert!(columns.len() >= entry.stat_descriptors().len(), "{}", entry.id());
1232 let actions = line["actions"].as_array().expect("actions is a list");
1233 assert_eq!(actions.len(), entry.action_descriptors().len(), "{}", entry.id());
1234 }
1235 }
1236
1237 fn legacy_mode(args: &Args) -> Mode {
1239 if args.list {
1240 Mode::List
1241 } else if args.model.is_none() {
1242 Mode::InfoOnly
1243 } else if args.params {
1244 Mode::Params
1245 } else if args.export_stats.is_some() {
1246 Mode::ExportStats
1247 } else if args.export.is_some() {
1248 Mode::ExportFinal
1249 } else {
1250 Mode::Benchmark
1251 }
1252 }
1253
1254 fn command_lines(text: &str) -> Vec<Vec<String>> {
1259 let mut lines = Vec::new();
1260 let mut pending = String::new();
1261 for raw in text.lines() {
1262 let raw = raw.trim_start().trim_start_matches("//!").trim();
1263 if let Some(start) = raw.strip_suffix('\\') {
1264 pending.push_str(start);
1265 continue;
1266 }
1267 pending.push_str(raw);
1268 let line = std::mem::take(&mut pending);
1269 let rest = line
1270 .strip_prefix("cargo run --release -p henad-cli -- ")
1271 .or_else(|| line.strip_prefix("henad-cli "));
1272 if let Some(rest) = rest.filter(|rest| !rest.starts_with('[')) {
1273 lines.push(rest.split_whitespace().map(str::to_owned).collect());
1274 }
1275 }
1276 lines
1277 }
1278
1279 #[test]
1284 fn existing_invocations_keep_their_mode() {
1285 let reference = std::path::Path::new(env!("CARGO_MANIFEST_DIR")).join("../../docs/reference/cli.md");
1287 let Ok(reference) = std::fs::read_to_string(&reference) else {
1288 eprintln!("note: skipped, {} is absent", reference.display());
1289 return;
1290 };
1291 let mut lines = command_lines(&reference);
1292 lines.extend(command_lines(include_str!("lib.rs")));
1293 let documented = lines.len();
1294 let scripts = [
1295 "boids --json --steps 100 --warmup 10 --reps 5 --seed 42 --threads 1 --set num_agents=10000",
1296 "sir --json --steps 100 --warmup 10 --reps 5 --seed 42 --threads 0 --set grid_width=256 \
1297 --set grid_height=256",
1298 "gpu_boids --json --steps 100 --warmup 10 --reps 5 --seed 42 --threads 0 --global-warmup 1000 \
1299 --set num_agents=10000",
1300 "sir --params",
1301 "--list",
1302 "--info --json",
1303 "sir --set grid_width=100 --set grid_height=100 --set infection_rate=0.3 --set recovery_rate=0.05 \
1304 --set initial_infected_pct=0.01 --steps 200 --seed 1 --export-stats sir_henad_001.csv",
1305 "virus_network --set num_agents=150 --steps 200 --seed 1 --export-stats virus_henad_001.csv",
1306 ];
1307 lines.extend(
1308 scripts
1309 .iter()
1310 .map(|line| line.split_whitespace().map(str::to_owned).collect()),
1311 );
1312
1313 let mut sweeps = 0;
1314 for line in &lines {
1315 let argv = std::iter::once("henad-cli").chain(line.iter().map(String::as_str));
1316 let args = Args::try_parse_from(argv).unwrap_or_else(|error| panic!("{line:?} parses: {error}"));
1317 if args.explore == ExploreArgs::default() {
1318 assert_eq!(Mode::of(&args), legacy_mode(&args), "{line:?}");
1319 } else {
1320 let mode = if args.explore.merge.is_empty() {
1321 Mode::Explore
1322 } else {
1323 Mode::Merge
1324 };
1325 assert_eq!(Mode::of(&args), mode, "{line:?}");
1326 sweeps += 1;
1327 }
1328 }
1329 assert!(
1330 documented - sweeps >= 13,
1331 "found {documented} documented lines, {sweeps} of them sweeps"
1332 );
1333 }
1334
1335 #[test]
1341 fn exported_stats_are_prepared_like_a_publish() {
1342 let entry = example_models()
1343 .get("team_assembly")
1344 .cloned()
1345 .expect("team_assembly is registered");
1346 let overrides =
1347 parse_overrides(&["num_agents=4".to_owned(), "team_size=4".to_owned(), "p=0".to_owned()]).expect("valid");
1348 let params = resolve_params(entry.param_descriptors(), &overrides).expect("in range");
1349 let mut simulation = cpu_simulation(&entry, ¶ms, 1, &[]);
1350
1351 let mut out = Vec::new();
1352 let rows = write_series(&mut simulation, 3, 2, StatsWriter::new(&mut out)).expect("writes");
1353 assert_eq!(rows, 3);
1354 let text = String::from_utf8(out).expect("utf-8");
1355 let mut lines = text.lines();
1356 let header: Vec<&str> = lines.next().expect("a header").split(',').collect();
1357 let at = |name: &str| header.iter().position(|h| *h == name).expect("the column is exported");
1358 let (tick, share, size) = (at("tick"), at("Giant Component Share"), at("Mean Component Size"));
1359 for (line, cliques) in lines.zip([1.0, 3.0, 4.0]) {
1360 let row: Vec<f64> = line.split(',').map(|v| v.parse().expect("a number")).collect();
1361 assert!(
1362 (row[share] - 1.0 / cliques).abs() < 1e-9,
1363 "tick {}: giant share {}",
1364 row[tick],
1365 row[share]
1366 );
1367 assert_eq!(row[size], 4.0, "tick {}", row[tick]);
1368 }
1369 }
1370
1371 fn cpu_simulation(entry: &ModelEntry, params: &[ParamValue], seed: u64, actions: &[String]) -> Simulation {
1373 let schedule = Schedule::parse(actions, entry.id(), entry.action_descriptors()).expect("declared actions");
1374 RunSetup::from_parts(entry, params, Some(seed), schedule)
1375 .expect("the values fit the model")
1376 .build(None)
1377 .expect("the model builds")
1378 }
1379
1380 struct SmallGpuGrid {
1382 ctx: GpuContext,
1383 entry: ModelEntry,
1384 params: Vec<ParamValue>,
1385 }
1386
1387 impl SmallGpuGrid {
1388 fn new(id: &str) -> Option<Self> {
1394 let ctx = headless_test_device(&TestDeviceRequest::raised(example_models().gpu_needs()))?;
1395 let entry = example_models().get(id).cloned().expect("the model is registered");
1396 let overrides = parse_overrides(&["grid_width=64".to_owned(), "grid_height=64".to_owned()]).expect("valid");
1397 let params = resolve_params(entry.param_descriptors(), &overrides).expect("in range");
1398 Some(Self { ctx, entry, params })
1399 }
1400
1401 fn schedule(&self, raw: &[&str]) -> Schedule {
1402 let raw: Vec<String> = raw.iter().map(|&s| s.to_owned()).collect();
1403 Schedule::parse(&raw, self.entry.id(), self.entry.action_descriptors()).expect("declared actions")
1404 }
1405
1406 fn simulation(&self, seed: u64, schedule: &Schedule) -> Simulation {
1408 RunSetup::from_parts(&self.entry, &self.params, Some(seed), schedule.clone())
1409 .expect("the values fit the model")
1410 .build(Some(&self.ctx))
1411 .expect("the model builds")
1412 }
1413 }
1414
1415 fn rows(mut simulation: Simulation, every: u64, total: u64) -> Vec<String> {
1417 let mut out = Vec::new();
1418 write_series(&mut simulation, total, every, StatsWriter::new(&mut out)).expect("writes");
1419 let text = String::from_utf8(out).expect("utf-8");
1420 text.lines().skip(1).map(str::to_owned).collect()
1421 }
1422
1423 #[test]
1428 fn a_gpu_action_fires_once_whatever_the_sampling_interval() {
1429 let Some(life) = SmallGpuGrid::new("gpu_game_of_life") else {
1430 return;
1431 };
1432 let schedule = life.schedule(&["randomise@2"]);
1433 let every_tick = rows(life.simulation(1, &schedule), 1, 6);
1434 let every_third = rows(life.simulation(1, &schedule), 3, 6);
1435 assert_eq!(every_tick.len(), 7, "ticks 0 to 6");
1436 let shared = [0, 3, 6].map(|tick| every_tick[tick].clone());
1437 assert_eq!(every_third, shared, "rows at ticks 0, 3 and 6");
1438 }
1439
1440 #[test]
1442 fn a_gpu_single_point_sweep_matches_export_stats() {
1443 let Some(sir) = SmallGpuGrid::new("gpu_sir") else {
1444 return;
1445 };
1446 let dir = ScratchDir::new("gpu-single-point");
1447 let args = Args::parse_from([
1448 "henad-cli",
1449 "gpu_sir",
1450 "--set",
1451 "grid_width=64",
1452 "--set",
1453 "grid_height=64",
1454 "--steps",
1455 "13",
1456 "--stats-every",
1457 "4",
1458 "--seed",
1459 "7",
1460 "--out",
1461 dir.arg(),
1462 ]);
1463 let status = explore::run(&args, &sir.entry, Some(&sir.ctx), None, test_provenance()).expect("the sweep runs");
1464 assert_eq!(status, 0);
1465
1466 let exported = rows(sir.simulation(run_seed(7, 0), &sir.schedule(&[])), 4, 13);
1467 let series = series_without_run_id(&dir);
1468 assert_eq!(series[1..], exported, "rows at ticks 0, 4, 8, 12 and 13");
1469 assert_final_reducers(&dir, &series[0], &exported[exported.len() - 1]);
1470 }
1471}