1use 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
11const DEFAULT_TICKS: u64 = 20;
14
15pub const MIN_TICKS: u64 = COARSE_CADENCE + 1;
17
18const 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#[derive(Debug, Clone)]
29pub struct CheckSettings {
30 gpu: Option<GpuContext>,
31 ticks: u64,
32 thread_counts: (usize, usize),
33 texts: Vec<(String, String, String)>,
35 exemptions: Vec<(String, ModelCheck, String)>,
37}
38
39impl 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 pub fn gpu(mut self, device: GpuContext) -> Self {
59 self.gpu = Some(device);
60 self
61 }
62
63 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 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 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 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 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 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 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 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 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
192fn shrunk(kind: &ParamKind, small: u32) -> ParamValue {
196 match *kind {
197 ParamKind::U32 { min, max, default } => ParamValue::U32(small.min(default).clamp(min, max)),
198 #[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}