Skip to main content

henad_explore/testing/
settings.rs

1//! Settings of a check run: the device, the run length, the thread counts, parameter overrides and exemptions.
2
3use henad_compute::entry::ModelEntry;
4use henad_compute::gpu::GpuContext;
5use henad_core::explore::value::parse_value;
6use henad_core::params::{ParamKind, ParamValue};
7
8use super::determinism::COARSE_CADENCE;
9use super::{ModelCheck, declared_defaults, error_text};
10
11/// Number of ticks that a run of a determinism check steps by default. The count is not a multiple of the seven-tick
12/// cadence, so the last tick is sampled in addition.
13const DEFAULT_TICKS: u64 = 20;
14
15/// Minimum number of ticks that [`CheckSettings::ticks`] accepts, one past the coarser sampling cadence.
16pub const MIN_TICKS: u64 = COARSE_CADENCE + 1;
17
18/// Small value of each size parameter that the engines prepend, used by every check that builds the model.
19const SMALL_SIZES: [(&str, u32); 5] = [
20    ("grid_width", 128),
21    ("grid_height", 128),
22    ("num_agents", 256),
23    ("world_width", 128),
24    ("world_height", 128),
25];
26
27/// Device, run length, thread counts, parameter overrides and exemptions of a check run.
28#[derive(Debug, Clone)]
29pub struct CheckSettings {
30    gpu: Option<GpuContext>,
31    ticks: u64,
32    thread_counts: (usize, usize),
33    /// Model id, parameter id and value text of each override.
34    texts: Vec<(String, String, String)>,
35    /// Model id, check and reason of each exemption.
36    exemptions: Vec<(String, ModelCheck, String)>,
37}
38
39/// No device, 20 ticks, 1 and 7 threads, no overrides and no exemptions.
40impl Default for CheckSettings {
41    fn default() -> Self {
42        Self {
43            gpu: None,
44            ticks: DEFAULT_TICKS,
45            thread_counts: (1, 7),
46            texts: Vec::new(),
47            exemptions: Vec::new(),
48        }
49    }
50}
51
52impl CheckSettings {
53    /// Runs the GPU models' checks on `device`. Without a device, they are skipped.
54    ///
55    /// Note that a check attributes every fault it finds in the device's sink to itself, and clones of a context share
56    /// the sink. A test that shares `device` with another test can have its faults dropped, or reported by a check.
57    /// Each test takes its own device from `headless_test_device`.
58    pub fn gpu(mut self, device: GpuContext) -> Self {
59        self.gpu = Some(device);
60        self
61    }
62
63    /// Steps each run of a determinism check `ticks` ticks.
64    ///
65    /// # Panics
66    ///
67    /// Panics when `ticks` is below [`MIN_TICKS`]. A shorter run takes no sample at the coarser cadence of
68    /// [`ModelCheck::SamplingCadence`] past tick 0.
69    pub fn ticks(mut self, ticks: u64) -> Self {
70        assert!(
71            ticks >= MIN_TICKS,
72            "a check runs at least {MIN_TICKS} ticks, not {ticks}"
73        );
74        self.ticks = ticks;
75        self
76    }
77
78    /// Compares runs at `low` and `high` worker threads in [`ModelCheck::ThreadCount`], which sizes the work to split
79    /// into twice `high` jobs.
80    ///
81    /// # Panics
82    ///
83    /// Panics when `low` is 0 or not below `high`.
84    pub fn thread_counts(mut self, low: usize, high: usize) -> Self {
85        assert!(
86            0 < low && low < high,
87            "thread counts {low} and {high} are not two counts in order"
88        );
89        self.thread_counts = (low, high);
90        self
91    }
92
93    /// Sets parameter `param_id` of model `model_id` from `text`, as `--set` reads it.
94    ///
95    /// The value holds in every check that builds the model, and no check changes it. A parameter that the model does
96    /// not declare, or a value that the parameter rejects, fails every such check, including a check that cannot run
97    /// for want of a device.
98    pub fn set_text(mut self, model_id: &str, param_id: &str, text: &str) -> Self {
99        self.texts
100            .push((model_id.to_owned(), param_id.to_owned(), text.to_owned()));
101        self
102    }
103
104    /// Skips `check` for model `model_id`, recording `reason` in its report.
105    pub fn exempt(mut self, model_id: &str, check: ModelCheck, reason: &str) -> Self {
106        self.exemptions.push((model_id.to_owned(), check, reason.to_owned()));
107        self
108    }
109
110    pub(super) fn device(&self) -> Option<&GpuContext> {
111        self.gpu.as_ref()
112    }
113
114    pub(super) fn run_ticks(&self) -> u64 {
115        self.ticks
116    }
117
118    pub(super) fn low_and_high_threads(&self) -> (usize, usize) {
119        self.thread_counts
120    }
121
122    /// Returns the reason that model `model_id` is exempt from `check`, or `None` when the model is not exempt.
123    pub(super) fn exemption(&self, model_id: &str, check: ModelCheck) -> Option<&str> {
124        self.exemptions
125            .iter()
126            .find(|(model, exempt, _)| model == model_id && *exempt == check)
127            .map(|(_, _, reason)| reason.as_str())
128    }
129
130    /// Returns every model id that an override or an exemption specifies, each once.
131    pub(super) fn named_models(&self) -> impl Iterator<Item = &str> {
132        let mut ids: Vec<&str> = self
133            .texts
134            .iter()
135            .map(|(model, _, _)| model.as_str())
136            .chain(self.exemptions.iter().map(|(model, _, _)| model.as_str()))
137            .collect();
138        ids.sort_unstable();
139        ids.dedup();
140        ids.into_iter()
141    }
142
143    /// Returns whether an override sets parameter `param_id` of `entry`.
144    pub(super) fn overrides(&self, entry: &ModelEntry, param_id: &str) -> bool {
145        self.texts
146            .iter()
147            .any(|(model, param, _)| model == entry.id() && param == param_id)
148    }
149
150    /// Returns the values that a check building `entry` uses: the declared defaults, the sizes made small, then the
151    /// overrides.
152    ///
153    /// # Errors
154    ///
155    /// Returns the failure message for an override that refers to no parameter of `entry`, or a value that its
156    /// parameter rejects.
157    pub(super) fn check_values(&self, entry: &ModelEntry) -> Result<Vec<ParamValue>, String> {
158        let mut values = declared_defaults(entry);
159        for (param_id, small) in SMALL_SIZES {
160            if let Some(index) = entry.param_index(param_id) {
161                values[index] = shrunk(&entry.param_descriptors()[index].kind, small);
162            }
163        }
164        self.apply_overrides(entry, values)
165    }
166
167    /// Returns the declared defaults of `entry` with the overrides applied, the values the GPU checks build at.
168    ///
169    /// # Errors
170    ///
171    /// As [`Self::check_values`].
172    pub(super) fn default_values(&self, entry: &ModelEntry) -> Result<Vec<ParamValue>, String> {
173        self.apply_overrides(entry, declared_defaults(entry))
174    }
175
176    fn apply_overrides(&self, entry: &ModelEntry, mut values: Vec<ParamValue>) -> Result<Vec<ParamValue>, String> {
177        for (_, param_id, text) in self.texts.iter().filter(|(model, _, _)| model == entry.id()) {
178            let index = entry
179                .param_index(param_id)
180                .ok_or_else(|| format!("The settings set parameter '{param_id}', which the model does not declare."))?;
181            values[index] = parse_value(&entry.param_descriptors()[index].kind, text).map_err(|error| {
182                format!(
183                    "The settings set parameter '{param_id}' to '{text}': {}.",
184                    error_text(&error)
185                )
186            })?;
187        }
188        Ok(values)
189    }
190}
191
192/// Returns `small` as a value of `kind`, at most the default and clamped to the bounds.
193///
194/// A kind other than a number keeps its default.
195fn shrunk(kind: &ParamKind, small: u32) -> ParamValue {
196    match *kind {
197        ParamKind::U32 { min, max, default } => ParamValue::U32(small.min(default).clamp(min, max)),
198        // Every small size is exact in an `f32`.
199        #[expect(clippy::cast_precision_loss, reason = "the small sizes are below 2^24")]
200        ParamKind::F32 { min, max, default, .. } => ParamValue::F32((small as f32).min(default).clamp(min, max)),
201        _ => kind.default_value(),
202    }
203}