1mod built;
21mod declared;
22mod determinism;
23mod device;
24#[cfg(not(target_arch = "wasm32"))]
25mod gpu;
26mod report;
27mod settings;
28
29use std::error::Error;
30use std::fmt;
31
32use henad_compute::entry::{ModelEntry, ModelSet};
33use henad_compute::fault::{Fault, catching, install_panic_hook};
34use henad_compute::gpu::GpuContext;
35use henad_compute::gpu::fault::catching_on;
36use henad_core::metadata::Backend;
37use henad_core::params::ParamValue;
38use henad_core::view::{StatEntry, StatValue};
39
40pub use device::TestDeviceRequest;
41#[cfg(not(target_arch = "wasm32"))]
42pub use device::headless_test_device;
43pub use report::{CheckFailure, ModelReport, SetReport, SkipReason, SkippedCheck};
44pub use settings::{CheckSettings, MIN_TICKS};
45
46#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
48#[non_exhaustive]
49pub enum ModelCheck {
50 ModelId,
53 ParamIds,
56 StatLabels,
58 ActionIds,
61 Palette,
63 Metadata,
65 DefaultSetup,
67 DefaultsFit,
69 ApplyModes,
72 Views,
74 ParallelJobs,
76 Actions,
78 StatCount,
81 ThreadCount,
84 SameSeed,
86 SeedSensitivity,
88 SamplingCadence,
90 BaselineBuild,
93 FullSubmission,
96 SampledSlice,
98}
99
100impl ModelCheck {
101 pub const ALL: &'static [Self] = &[
103 Self::ModelId,
104 Self::ParamIds,
105 Self::StatLabels,
106 Self::ActionIds,
107 Self::Palette,
108 Self::Metadata,
109 Self::DefaultSetup,
110 Self::DefaultsFit,
111 Self::ApplyModes,
112 Self::Views,
113 Self::ParallelJobs,
114 Self::Actions,
115 Self::StatCount,
116 Self::ThreadCount,
117 Self::SameSeed,
118 Self::SeedSensitivity,
119 Self::SamplingCadence,
120 Self::BaselineBuild,
121 Self::FullSubmission,
122 Self::SampledSlice,
123 ];
124
125 fn builds(self) -> bool {
127 !matches!(
128 self,
129 Self::ModelId
130 | Self::ParamIds
131 | Self::StatLabels
132 | Self::ActionIds
133 | Self::Palette
134 | Self::Metadata
135 | Self::DefaultSetup
136 | Self::DefaultsFit
137 )
138 }
139
140 fn applies_to(self, backend: Backend) -> bool {
142 match self {
143 Self::ThreadCount => backend == Backend::Cpu,
144 Self::DefaultsFit | Self::BaselineBuild | Self::FullSubmission | Self::SampledSlice => {
145 backend == Backend::Gpu
146 }
147 _ => true,
148 }
149 }
150
151 fn compares_runs(self) -> bool {
153 matches!(self, Self::SameSeed | Self::SeedSensitivity | Self::SamplingCadence)
154 }
155}
156
157impl fmt::Display for ModelCheck {
158 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
159 fmt::Debug::fmt(self, f)
160 }
161}
162
163#[must_use]
169pub fn check_model(entry: &ModelEntry, settings: &CheckSettings) -> ModelReport {
170 check_model_requiring(entry, settings, device::gpu_required())
171}
172
173pub(crate) fn check_model_requiring(entry: &ModelEntry, settings: &CheckSettings, gpu_required: bool) -> ModelReport {
175 let mut report = ModelReport::new(entry.id());
176 let backend = entry.metadata().backend;
177 let gpu = if backend == Backend::Gpu {
178 settings.device()
179 } else {
180 None
181 };
182 let refused_override = settings.default_values(entry).err();
183 for &check in ModelCheck::ALL {
184 if let Some(reason) = skip_reason(entry, settings, check, gpu.is_some()) {
185 let exempt = settings.exemption(entry.id(), check).is_some();
186 match skipped_failure(check, &reason, exempt, refused_override.as_deref(), gpu_required) {
187 Some(message) => report.fail(check, message),
188 None => report.skip(check, reason),
189 }
190 continue;
191 }
192 let outcome = guarded(gpu, || run_check(entry, settings, check, gpu));
193 match outcome {
194 Ok(Ran::Passed) => {}
195 Ok(Ran::PassedAtJobs(jobs)) => report.set_thread_count_jobs(jobs),
196 Ok(Ran::Skipped(reason)) => report.skip(check, reason),
197 Err(message) => report.fail(check, message),
198 }
199 }
200 report
201}
202
203#[must_use]
208pub fn check_model_set(models: &ModelSet, settings: &CheckSettings) -> SetReport {
209 let reports = models.iter().map(|entry| check_model(entry, settings)).collect();
210 let unknown_models = settings
211 .named_models()
212 .filter(|id| models.get(id).is_none())
213 .map(str::to_owned)
214 .collect();
215 SetReport::new(reports, unknown_models)
216}
217
218pub fn assert_set_conforms(models: &ModelSet, settings: &CheckSettings) {
228 install_panic_hook();
229 let report = check_model_set(models, settings);
230 report.assert_passed();
231 #[expect(clippy::print_stdout, reason = "libtest captures a test's standard output")]
232 {
233 println!("{report}");
234 }
235}
236
237const SEED: u64 = 1;
239
240const CHECKING: &str = "checking the model";
242
243enum Ran {
245 Passed,
246 PassedAtJobs(usize),
248 Skipped(SkipReason),
250}
251
252fn skip_reason(
254 entry: &ModelEntry,
255 settings: &CheckSettings,
256 check: ModelCheck,
257 has_device: bool,
258) -> Option<SkipReason> {
259 let metadata = entry.metadata();
260 if !check.applies_to(metadata.backend) {
261 return Some(SkipReason::OtherBackend);
262 }
263 if check.compares_runs() && !metadata.replays_exactly {
264 return Some(SkipReason::InexactReplay);
265 }
266 if let Some(reason) = settings.exemption(entry.id(), check) {
267 return Some(SkipReason::Exempt(reason.to_owned()));
268 }
269 if cfg!(target_arch = "wasm32") && check == ModelCheck::ThreadCount {
270 return Some(SkipReason::NativeOnly);
271 }
272 if metadata.backend == Backend::Gpu && check.builds() {
273 if cfg!(target_arch = "wasm32") {
274 return Some(SkipReason::NativeOnly);
275 }
276 if !has_device {
277 return Some(SkipReason::NoDevice);
278 }
279 }
280 None
281}
282
283fn skipped_failure(
290 check: ModelCheck,
291 reason: &SkipReason,
292 exempt: bool,
293 refused_override: Option<&str>,
294 gpu_required: bool,
295) -> Option<String> {
296 match (reason, refused_override) {
297 (SkipReason::OtherBackend | SkipReason::InexactReplay, _) if exempt => Some(format!(
298 "The settings exempt this check, and it does not apply to the model: {reason}."
299 )),
300 (SkipReason::NoDevice | SkipReason::NativeOnly, Some(message)) if check.builds() => Some(message.to_owned()),
301 (SkipReason::NoDevice, _) if gpu_required => {
302 Some("HENAD_REQUIRE_GPU is set, but the settings give no GPU device.".to_owned())
303 }
304 _ => None,
305 }
306}
307
308fn run_check(
310 entry: &ModelEntry,
311 settings: &CheckSettings,
312 check: ModelCheck,
313 gpu: Option<&GpuContext>,
314) -> Result<Ran, String> {
315 match check {
316 ModelCheck::ModelId => declared::model_id(entry),
317 ModelCheck::ParamIds => declared::param_ids(entry),
318 ModelCheck::StatLabels => declared::stat_labels(entry),
319 ModelCheck::ActionIds => declared::action_ids(entry),
320 ModelCheck::Palette => declared::palette(entry),
321 ModelCheck::Metadata => declared::metadata(entry),
322 ModelCheck::DefaultSetup => declared::default_setup(entry),
323 ModelCheck::DefaultsFit => declared::defaults_fit(entry),
324 ModelCheck::ApplyModes => built::apply_modes(entry, &settings.check_values(entry)?, gpu),
325 ModelCheck::Views => built::views(entry, &settings.check_values(entry)?, gpu),
326 ModelCheck::ParallelJobs => built::parallel_jobs(entry, &settings.check_values(entry)?, gpu),
327 ModelCheck::Actions if gpu.is_some() => built::actions(entry, &settings.default_values(entry)?, gpu),
329 ModelCheck::Actions => built::actions(entry, &settings.check_values(entry)?, gpu),
330 ModelCheck::StatCount => built::stat_count(entry, &settings.check_values(entry)?, gpu),
331 ModelCheck::ThreadCount => return determinism::thread_count(entry, settings),
332 ModelCheck::SameSeed => determinism::same_seed(entry, settings, gpu),
333 ModelCheck::SeedSensitivity => determinism::seed_sensitivity(entry, settings, gpu),
334 ModelCheck::SamplingCadence => determinism::sampling_cadence(entry, settings, gpu),
335 #[cfg(not(target_arch = "wasm32"))]
336 ModelCheck::BaselineBuild => gpu::baseline_build(entry, settings, device(gpu)?),
337 #[cfg(not(target_arch = "wasm32"))]
338 ModelCheck::FullSubmission => gpu::full_submission(entry, settings, device(gpu)?),
339 #[cfg(not(target_arch = "wasm32"))]
340 ModelCheck::SampledSlice => gpu::sampled_slice(entry, settings, device(gpu)?),
341 #[cfg(target_arch = "wasm32")]
342 ModelCheck::BaselineBuild | ModelCheck::FullSubmission | ModelCheck::SampledSlice => {
343 return Ok(Ran::Skipped(SkipReason::NativeOnly));
344 }
345 }
346 .map(|()| Ran::Passed)
347}
348
349#[cfg(not(target_arch = "wasm32"))]
351fn device(gpu: Option<&GpuContext>) -> Result<&GpuContext, String> {
352 gpu.ok_or_else(|| "A GPU check ran with no device.".to_owned())
353}
354
355fn guarded<T>(gpu: Option<&GpuContext>, body: impl FnOnce() -> Result<T, String>) -> Result<T, String> {
361 let Some(ctx) = gpu else {
362 return catching(CHECKING, body).map_err(|fault| fault_text(&fault))?;
363 };
364 drop(ctx.faults.take());
366 let result = catching_on(ctx, CHECKING, body).map_err(|fault| fault_text(&fault))?;
367 #[cfg(not(target_arch = "wasm32"))]
368 henad_compute::gpu::stepping::wait(ctx).map_err(|fault| fault_text(&fault))?;
369 if let Some(fault) = ctx.faults.take() {
370 return Err(fault_text(&fault));
371 }
372 result
373}
374
375fn fault_text(fault: &Fault) -> String {
377 format!("Fault {fault}.")
378}
379
380fn error_text(error: &(dyn Error + 'static)) -> String {
384 let mut text = error.to_string();
385 let mut current = error;
386 while !current.is::<Fault>()
387 && let Some(source) = current.source()
388 {
389 text.push_str(": ");
390 text.push_str(&source.to_string());
391 current = source;
392 }
393 text
394}
395
396fn declared_defaults(entry: &ModelEntry) -> Vec<ParamValue> {
398 entry
399 .param_descriptors()
400 .iter()
401 .map(|descriptor| descriptor.kind.default_value())
402 .collect()
403}
404
405fn stat_bits(entries: &[StatEntry]) -> Vec<(&'static str, Vec<u64>)> {
407 entries
408 .iter()
409 .map(|entry| {
410 let bits = match &entry.value {
411 StatValue::Scalar(value) => vec![value.to_bits()],
412 StatValue::Vector2D { x, y } => vec![x.to_bits(), y.to_bits()],
413 StatValue::Histogram { edges, counts } => edges
414 .iter()
415 .map(|edge| edge.to_bits())
416 .chain(counts.iter().copied())
417 .collect(),
418 };
419 (entry.label, bits)
420 })
421 .collect()
422}
423
424fn first_stat_difference(first: &[StatEntry], second: &[StatEntry]) -> Option<String> {
426 let (first, second) = (stat_bits(first), stat_bits(second));
427 if first.len() != second.len() {
428 return Some(format!("{} stats against {}", first.len(), second.len()));
429 }
430 first
431 .iter()
432 .zip(&second)
433 .find(|(left, right)| left != right)
434 .map(|((label, _), _)| format!("stat '{label}'"))
435}