1#[cfg(not(target_arch = "wasm32"))]
7mod cpu;
8#[cfg(not(target_arch = "wasm32"))]
9mod gpu;
10
11use std::collections::{BTreeMap, BTreeSet};
12use std::fmt;
13use std::io;
14use std::num::{NonZeroUsize, ParseIntError};
15use std::str::FromStr;
16use std::sync::atomic::{AtomicU8, Ordering};
17use std::sync::{Arc, Condvar, Mutex, PoisonError};
18use std::time::Duration;
19
20use web_time::Instant;
21
22use henad_compute::entry::ModelEntry;
23use henad_compute::gpu::{Demand, GpuContext, MAX_STEPS_PER_SUBMISSION};
24use henad_compute::runner::CAN_SPAWN_THREADS;
25use henad_core::action::Schedule;
26use henad_core::explore::measure::MeasurePlan;
27use henad_core::explore::outcome::{PlannedRun, RunOutcome};
28use henad_core::explore::plan::Plan;
29use henad_core::metadata::Backend;
30use henad_core::params::ParamValue;
31
32use crate::cursor::{CursorState, RunCursor};
33use crate::probe::ProbeReport;
34
35pub const JOBS_PER_THREAD: usize = 4;
37
38pub const POPULATION_PER_JOB: u64 = 4096;
40
41const SLICE_TARGET_MS: f64 = 20.0;
43
44const MAX_SLICE_STEPS: u64 = 1 << 20;
46
47pub(crate) const MAX_AUTO_GPU_TRACKS: usize = 4;
49
50pub(crate) const LARGE_GPU_POPULATION: u64 = 1 << 20;
52
53#[derive(Debug, Clone, Copy, Default, PartialEq, Eq)]
55pub enum Concurrency {
56 #[default]
58 Auto,
59 Fixed(NonZeroUsize),
61}
62
63impl fmt::Display for Concurrency {
64 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
65 match self {
66 Self::Auto => f.write_str("auto"),
67 Self::Fixed(count) => write!(f, "{count}"),
68 }
69 }
70}
71
72impl FromStr for Concurrency {
73 type Err = ParseIntError;
74
75 fn from_str(raw: &str) -> Result<Self, Self::Err> {
77 if raw == "auto" {
78 Ok(Self::Auto)
79 } else {
80 raw.parse().map(Self::Fixed)
81 }
82 }
83}
84
85#[derive(Debug, Clone, Copy, PartialEq, Eq)]
87pub(crate) struct ExecutionBudget {
88 pub(crate) workers: usize,
90 pub(crate) memory_budget: Option<u64>,
93 pub(crate) gpu_memory_budget: Option<u64>,
96 pub(crate) can_spawn_threads: bool,
98}
99
100impl ExecutionBudget {
101 pub(crate) fn detect() -> Self {
103 Self {
104 workers: rayon::current_num_threads(),
105 memory_budget: None,
106 gpu_memory_budget: None,
107 can_spawn_threads: CAN_SPAWN_THREADS,
108 }
109 }
110}
111
112pub fn gpu_memory_budget(budget: Option<u64>, ctx: &GpuContext) -> u64 {
117 budget.unwrap_or_else(|| ctx.device.limits().max_buffer_size)
118}
119
120#[derive(Debug, Clone, Copy, PartialEq, Eq)]
124pub struct ExecutionLayout {
125 pub cpu_lanes: usize,
127 pub threads_per_lane: usize,
129 pub gpu_tracks: usize,
131}
132
133impl ExecutionLayout {
134 pub fn projected_bytes(&self, probe: &ProbeReport) -> u64 {
138 let device_bytes = probe.demand.as_ref().map_or(0, |demand| demand.bytes());
139 self.cpu_lanes as u64 * probe.heap_bytes + self.gpu_tracks as u64 * device_bytes
140 }
141}
142
143pub(crate) fn choose_layout(
157 concurrency: Concurrency,
158 resources: &ExecutionBudget,
159 backend: Backend,
160 probe: &ProbeReport,
161 runs: u64,
162) -> ExecutionLayout {
163 let runs = usize::try_from(runs).unwrap_or(usize::MAX);
164 if backend == Backend::Gpu {
165 let tracks = match concurrency {
166 Concurrency::Fixed(tracks) => tracks.get(),
167 Concurrency::Auto if probe.population >= LARGE_GPU_POPULATION => 1,
168 Concurrency::Auto => {
169 let demand = probe.demand.as_ref().map_or(0, Demand::bytes);
170 let tracks_by_memory = match resources.gpu_memory_budget {
171 Some(budget) if demand > 0 => usize::try_from(budget / demand).unwrap_or(usize::MAX),
172 _ => usize::MAX,
173 };
174 tracks_by_memory.min(MAX_AUTO_GPU_TRACKS)
175 }
176 };
177 return ExecutionLayout {
178 cpu_lanes: 0,
179 threads_per_lane: 0,
180 gpu_tracks: tracks.min(runs).max(1),
181 };
182 }
183 let workers = resources.workers.max(1);
184 let one_lane = ExecutionLayout {
185 cpu_lanes: 1,
186 threads_per_lane: workers,
187 gpu_tracks: 0,
188 };
189 if !resources.can_spawn_threads {
190 return one_lane;
191 }
192 let (lanes, threads_per_lane) = match concurrency {
193 Concurrency::Auto => {
194 let jobs = probe
195 .parallel_jobs
196 .unwrap_or_else(|| usize::try_from(probe.population / POPULATION_PER_JOB).unwrap_or(usize::MAX));
197 let threads = jobs.div_ceil(JOBS_PER_THREAD).clamp(1, workers);
198 (workers / threads, threads)
199 }
200 Concurrency::Fixed(lanes) => (lanes.get(), (workers / lanes).max(1)),
201 };
202 let lanes_by_memory = match resources.memory_budget {
203 Some(budget) if probe.heap_bytes > 0 => usize::try_from(budget / probe.heap_bytes).unwrap_or(usize::MAX),
204 _ => usize::MAX,
205 };
206 let lanes = lanes.min(runs).min(lanes_by_memory);
207 if lanes <= 1 {
208 one_lane
209 } else {
210 ExecutionLayout {
211 cpu_lanes: lanes,
212 threads_per_lane,
213 gpu_tracks: 0,
214 }
215 }
216}
217
218const RUNNING: u8 = 0;
219const PAUSED: u8 = 1;
220const ABORTED: u8 = 2;
221
222#[derive(Debug, Clone, Default)]
226pub struct SweepControl {
227 shared: Arc<ControlState>,
228}
229
230#[derive(Debug, Default)]
231struct ControlState {
232 mode: AtomicU8,
234 lock: Mutex<()>,
236 resumed: Condvar,
237}
238
239impl SweepControl {
240 pub fn new() -> Self {
242 Self::default()
243 }
244
245 pub fn pause(&self) {
247 let _guard = self.shared.lock.lock().unwrap_or_else(PoisonError::into_inner);
248 if self.shared.mode.load(Ordering::Acquire) == RUNNING {
249 self.shared.mode.store(PAUSED, Ordering::Release);
250 }
251 }
252
253 pub fn resume(&self) {
255 let _guard = self.shared.lock.lock().unwrap_or_else(PoisonError::into_inner);
256 if self.shared.mode.load(Ordering::Acquire) == PAUSED {
257 self.shared.mode.store(RUNNING, Ordering::Release);
258 self.shared.resumed.notify_all();
259 }
260 }
261
262 pub fn abort(&self) {
264 let _guard = self.shared.lock.lock().unwrap_or_else(PoisonError::into_inner);
265 self.shared.mode.store(ABORTED, Ordering::Release);
266 self.shared.resumed.notify_all();
267 }
268
269 pub fn is_paused(&self) -> bool {
271 self.shared.mode.load(Ordering::Acquire) == PAUSED
272 }
273
274 pub fn is_aborted(&self) -> bool {
276 self.shared.mode.load(Ordering::Acquire) == ABORTED
277 }
278
279 pub fn proceed(&self) -> bool {
288 match self.shared.mode.load(Ordering::Acquire) {
289 RUNNING => true,
290 PAUSED => self.wait_while_paused(),
291 _ => false,
292 }
293 }
294
295 fn wait_while_paused(&self) -> bool {
296 let mut guard = self.shared.lock.lock().unwrap_or_else(PoisonError::into_inner);
297 loop {
298 match self.shared.mode.load(Ordering::Acquire) {
299 RUNNING => return true,
300 PAUSED => {
301 guard = self.shared.resumed.wait(guard).unwrap_or_else(PoisonError::into_inner);
302 }
303 _ => return false,
304 }
305 }
306 }
307}
308
309#[derive(Debug, Clone, Default)]
311pub struct ActiveRuns {
312 shared: Arc<Mutex<ActiveRunsTable>>,
313}
314
315#[derive(Debug, Default)]
317struct ActiveRunsTable {
318 stepping: BTreeMap<u64, ActiveRun>,
320 waiting: BTreeSet<u64>,
322}
323
324#[derive(Debug, Clone, Copy, PartialEq, Eq)]
326pub struct ActiveRun {
327 pub run: PlannedRun,
329 pub tick: u64,
331 pub end_tick: u64,
333}
334
335impl ActiveRuns {
336 pub fn new() -> Self {
338 Self::default()
339 }
340
341 pub fn list(&self) -> Vec<ActiveRun> {
343 self.lock().stepping.values().copied().collect()
344 }
345
346 pub fn waiting_count(&self) -> usize {
350 self.lock().waiting.len()
351 }
352
353 pub fn watch(&self, run: PlannedRun, end_tick: u64) -> RunWatch {
355 self.lock()
356 .stepping
357 .insert(run.run_id, ActiveRun { run, tick: 0, end_tick });
358 RunWatch {
359 active_runs: self.clone(),
360 run_id: run.run_id,
361 }
362 }
363
364 pub fn mark_committed(&self, run_id: u64) {
366 self.lock().waiting.remove(&run_id);
367 }
368
369 fn lock(&self) -> std::sync::MutexGuard<'_, ActiveRunsTable> {
370 self.shared.lock().unwrap_or_else(PoisonError::into_inner)
371 }
372}
373
374#[derive(Debug)]
376pub struct RunWatch {
377 active_runs: ActiveRuns,
378 run_id: u64,
379}
380
381impl RunWatch {
382 pub fn reach(&self, tick: u64) {
384 if let Some(active) = self.active_runs.lock().stepping.get_mut(&self.run_id) {
385 active.tick = tick;
386 }
387 }
388
389 pub fn finish(self) {
391 self.active_runs.lock().waiting.insert(self.run_id);
392 }
393}
394
395impl Drop for RunWatch {
396 fn drop(&mut self) {
397 self.active_runs.lock().stepping.remove(&self.run_id);
398 }
399}
400
401#[derive(Debug, Clone)]
403pub struct RunRequest<'p> {
404 pub run: PlannedRun,
406 pub run_key: u64,
408 pub params: &'p [ParamValue],
410 pub schedule: Schedule,
412}
413
414impl<'p> RunRequest<'p> {
415 pub fn planned(plan: &'p Plan, run: PlannedRun) -> Self {
421 let config = plan
422 .config(run.config_id)
423 .expect("a planned run's config is in its plan");
424 Self {
425 run,
426 run_key: plan.run_key(&run),
427 params: &config.params,
428 schedule: plan.schedule(config),
429 }
430 }
431}
432
433pub trait RunSink {
435 fn commit(&mut self, outcome: RunOutcome) -> io::Result<()>;
441
442 fn finished(&mut self, _outcome: &RunOutcome) {}
445}
446
447#[derive(Debug, Clone, Copy, PartialEq, Eq)]
449pub enum BatchEnd {
450 Complete,
452 Aborted,
454 DeviceLost,
457}
458
459#[derive(Debug, Clone, Copy, PartialEq, Eq)]
461pub struct GpuTrackDepth {
462 pub submissions_per_track: usize,
464 pub steps_per_submission: u32,
466}
467
468impl Default for GpuTrackDepth {
469 fn default() -> Self {
471 Self {
472 submissions_per_track: 2,
473 steps_per_submission: MAX_STEPS_PER_SUBMISSION,
474 }
475 }
476}
477
478#[derive(Debug)]
480pub enum ExecutionError {
481 NoDevice,
483 Pool(rayon::ThreadPoolBuildError),
485 Spawn(io::Error),
487 Sink(io::Error),
489 LanePanicked,
491}
492
493impl fmt::Display for ExecutionError {
494 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
495 match self {
496 Self::NoDevice => f.write_str("a GPU model needs a GPU device"),
497 Self::Pool(_) => f.write_str("cannot build a lane's thread pool"),
498 Self::Spawn(_) => f.write_str("cannot start a lane's thread"),
499 Self::Sink(_) => f.write_str("cannot write a finished run"),
500 Self::LanePanicked => f.write_str("a lane's thread panicked outside any run"),
501 }
502 }
503}
504
505impl std::error::Error for ExecutionError {
506 fn source(&self) -> Option<&(dyn std::error::Error + 'static)> {
507 match self {
508 Self::Pool(error) => Some(error),
509 Self::Spawn(error) | Self::Sink(error) => Some(error),
510 Self::NoDevice | Self::LanePanicked => None,
511 }
512 }
513}
514
515#[derive(Debug)]
517pub struct Executor<'a> {
518 entry: &'a ModelEntry,
519 #[cfg_attr(target_arch = "wasm32", expect(dead_code, reason = "a browser steps no GPU track"))]
520 gpu: Option<&'a GpuContext>,
521 measure: Arc<MeasurePlan>,
522 layout: ExecutionLayout,
523 control: SweepControl,
524 timeout: Option<Duration>,
526 active_runs: Option<ActiveRuns>,
528 pools: Vec<rayon::ThreadPool>,
530 #[cfg_attr(target_arch = "wasm32", expect(dead_code, reason = "a browser steps no GPU track"))]
533 gpu_memory_budget: Option<u64>,
534 #[cfg_attr(target_arch = "wasm32", expect(dead_code, reason = "a browser steps no GPU track"))]
535 track_depth: GpuTrackDepth,
536}
537
538impl<'a> Executor<'a> {
539 pub fn new(
550 entry: &'a ModelEntry,
551 gpu: Option<&'a GpuContext>,
552 measure: Arc<MeasurePlan>,
553 layout: ExecutionLayout,
554 control: SweepControl,
555 ) -> Result<Self, ExecutionError> {
556 let mut layout = layout;
557 let mut pools = Vec::new();
558 match entry.metadata().backend {
559 Backend::Gpu if gpu.is_none() => return Err(ExecutionError::NoDevice),
560 Backend::Gpu => {}
561 Backend::Cpu if !CAN_SPAWN_THREADS => {
562 layout.cpu_lanes = 1;
563 layout.threads_per_lane = rayon::current_num_threads();
564 }
565 Backend::Cpu => {
566 layout.cpu_lanes = layout.cpu_lanes.max(1);
567 layout.threads_per_lane = layout.threads_per_lane.max(1);
568 pools = (0..layout.cpu_lanes)
569 .map(|lane| {
570 rayon::ThreadPoolBuilder::new()
571 .num_threads(layout.threads_per_lane)
572 .thread_name(move |worker| format!("henad-lane-{lane}-{worker}"))
573 .build()
574 })
575 .collect::<Result<_, _>>()
576 .map_err(ExecutionError::Pool)?;
577 }
578 }
579 Ok(Self {
580 entry,
581 gpu,
582 measure,
583 layout,
584 control,
585 timeout: None,
586 active_runs: None,
587 pools,
588 gpu_memory_budget: None,
589 track_depth: GpuTrackDepth::default(),
590 })
591 }
592
593 pub fn with_timeout(self, timeout: Option<Duration>) -> Self {
599 Self { timeout, ..self }
600 }
601
602 pub fn with_active_runs(self, active_runs: Option<ActiveRuns>) -> Self {
604 Self { active_runs, ..self }
605 }
606
607 pub fn with_gpu_memory_budget(self, budget: Option<u64>) -> Self {
613 Self {
614 gpu_memory_budget: budget,
615 ..self
616 }
617 }
618
619 pub fn with_track_depth(self, depth: GpuTrackDepth) -> Self {
625 assert!(
626 depth.submissions_per_track > 0 && (1..=MAX_STEPS_PER_SUBMISSION).contains(&depth.steps_per_submission),
627 "{depth:?} cannot step a track"
628 );
629 Self {
630 track_depth: depth,
631 ..self
632 }
633 }
634
635 pub fn layout(&self) -> ExecutionLayout {
637 self.layout
638 }
639
640 pub fn control(&self) -> &SweepControl {
642 &self.control
643 }
644
645 pub fn run_batch(&self, requests: &[RunRequest<'_>], sink: &mut dyn RunSink) -> Result<BatchEnd, ExecutionError> {
654 match (self.entry.metadata().backend, self.pools.as_slice()) {
655 #[cfg(not(target_arch = "wasm32"))]
656 (Backend::Gpu, _) => {
657 let ctx = self.gpu.ok_or(ExecutionError::NoDevice)?;
658 gpu::run_on_tracks(self, ctx, requests, sink)
659 }
660 #[cfg(target_arch = "wasm32")]
662 (Backend::Gpu, _) => self.run_in_order(requests, sink, Placement::CallingThread),
663 (Backend::Cpu, []) => self.run_in_order(requests, sink, Placement::CallingThread),
665 (Backend::Cpu, [pool]) => self.run_in_order(requests, sink, Placement::LanePool(pool)),
666 #[cfg(not(target_arch = "wasm32"))]
667 (Backend::Cpu, pools) => cpu::run_in_lanes(self, pools, requests, sink),
668 #[cfg(target_arch = "wasm32")]
669 (Backend::Cpu, _) => self.run_in_order(requests, sink, Placement::CallingThread),
670 }
671 }
672
673 fn run_in_order(
675 &self,
676 requests: &[RunRequest<'_>],
677 sink: &mut dyn RunSink,
678 placement: Placement<'_>,
679 ) -> Result<BatchEnd, ExecutionError> {
680 for request in requests {
681 let Some(outcome) = self.drive_in(placement, request) else {
682 return Ok(BatchEnd::Aborted);
683 };
684 sink.finished(&outcome);
685 self.commit(sink, outcome)?;
686 }
687 Ok(BatchEnd::Complete)
688 }
689
690 fn drive_in(&self, placement: Placement<'_>, request: &RunRequest<'_>) -> Option<RunOutcome> {
695 match placement {
696 Placement::CallingThread => self.drive(request),
697 Placement::LanePool(pool) => run_in_pool(pool, || self.drive(request)),
698 }
699 }
700
701 fn drive(&self, request: &RunRequest<'_>) -> Option<RunOutcome> {
703 if !self.control.proceed() {
704 return None;
705 }
706 let mut slice = SliceSize::default();
709 let watch = self
710 .active_runs
711 .as_ref()
712 .map(|active_runs| active_runs.watch(request.run, self.measure.total()));
713 let mut cursor = RunCursor::new(self.entry, &self.measure, request, self.timeout);
714 loop {
715 let started = Instant::now();
716 let state = cursor.advance(slice.steps());
717 slice.adapt(started.elapsed());
718 if let CursorState::Finished(outcome) = state {
719 if let Some(watch) = watch {
720 watch.finish();
721 }
722 return Some(outcome);
723 }
724 if let Some(watch) = &watch {
725 watch.reach(cursor.tick());
726 }
727 if !self.control.proceed() {
728 return None;
729 }
730 }
731 }
732
733 fn commit(&self, sink: &mut dyn RunSink, outcome: RunOutcome) -> Result<(), ExecutionError> {
735 if let Some(active_runs) = &self.active_runs {
736 active_runs.mark_committed(outcome.run.run_id);
737 }
738 sink.commit(outcome).map_err(|error| {
739 self.control.abort();
740 ExecutionError::Sink(error)
741 })
742 }
743}
744
745#[cfg(not(target_arch = "wasm32"))]
747pub(crate) fn run_in_pool<R: Send>(pool: &rayon::ThreadPool, task: impl FnOnce() -> R + Send) -> R {
748 pool.install(task)
749}
750
751#[cfg(target_arch = "wasm32")]
753pub(crate) fn run_in_pool<R>(_pool: &rayon::ThreadPool, task: impl FnOnce() -> R) -> R {
754 task()
755}
756
757#[derive(Clone, Copy)]
759enum Placement<'p> {
760 CallingThread,
763 LanePool(&'p rayon::ThreadPool),
765}
766
767#[derive(Debug, Clone, Copy)]
771pub(crate) struct SliceSize {
772 steps: u64,
773 target_ms: f64,
775}
776
777impl Default for SliceSize {
778 fn default() -> Self {
779 Self::aiming_at(SLICE_TARGET_MS)
780 }
781}
782
783impl SliceSize {
784 pub(crate) fn aiming_at(target_ms: f64) -> Self {
786 Self { steps: 1, target_ms }
787 }
788
789 pub(crate) fn steps(&self) -> u64 {
790 self.steps
791 }
792
793 pub(crate) fn adapt(&mut self, elapsed: Duration) {
796 let elapsed_ms = elapsed.as_secs_f64() * 1000.0;
797 if elapsed_ms < self.target_ms / 2.0 {
798 self.steps = (self.steps * 2).min(MAX_SLICE_STEPS);
799 } else if elapsed_ms > self.target_ms * 2.0 {
800 let scaled = self.steps as f64 * self.target_ms / elapsed_ms;
801 self.steps = (scaled as u64).max(1);
802 }
803 }
804}
805
806#[cfg(not(target_arch = "wasm32"))]
808#[derive(Debug)]
809struct ReorderBuffer<T> {
810 next: usize,
812 waiting: BTreeMap<usize, T>,
813}
814
815#[cfg(not(target_arch = "wasm32"))]
816impl<T> Default for ReorderBuffer<T> {
817 fn default() -> Self {
818 Self {
819 next: 0,
820 waiting: BTreeMap::new(),
821 }
822 }
823}
824
825#[cfg(not(target_arch = "wasm32"))]
826impl<T> ReorderBuffer<T> {
827 fn insert(&mut self, index: usize, item: T) {
828 debug_assert!(index >= self.next, "item {index} was already handed out");
829 self.waiting.insert(index, item);
830 }
831
832 fn pop_ready(&mut self) -> Option<T> {
834 let item = self.waiting.remove(&self.next)?;
835 self.next += 1;
836 Some(item)
837 }
838
839 fn handed_out(&self) -> usize {
841 self.next
842 }
843}
844
845#[cfg(test)]
846mod tests {
847 use std::num::NonZeroUsize;
848 use std::sync::Arc;
849 use std::sync::mpsc;
850 use std::time::{Duration, Instant};
851
852 use henad_compute::entry::{ModelEntry, ModelState};
853 use henad_compute::fault::Fault;
854 use henad_compute::gpu::Demand;
855 use henad_compute::gpu::GpuContext;
856 use henad_core::explore::design::DesignKind;
857 use henad_core::explore::factor::{FactorSpec, LevelSpec};
858 use henad_core::explore::measure::MeasurePlan;
859 use henad_core::explore::outcome::{PlannedRun, RunOutcome, RunStatus};
860 use henad_core::explore::plan::Plan;
861 use henad_core::explore::spec::{BlockSpec, SweepSpec};
862 use henad_core::export::StatColumns;
863 use henad_core::metadata::Backend;
864 use henad_core::params::ParamValue;
865 use henad_models::example_models;
866
867 use super::{
868 ActiveRuns, BatchEnd, Concurrency, ExecutionBudget, ExecutionError, ExecutionLayout, Executor,
869 LARGE_GPU_POPULATION, ReorderBuffer, RunRequest, RunSink, SliceSize, SweepControl, choose_layout,
870 };
871 use crate::probe::ProbeReport;
872 use crate::tests::support::{lanes, tracks};
873
874 fn probe(parallel_jobs: Option<usize>, population: u64, heap_bytes: u64) -> ProbeReport {
875 ProbeReport {
876 params: Vec::new(),
877 seed: None,
878 columns: StatColumns::plan(&[]),
879 heap_bytes,
880 population,
881 parallel_jobs,
882 demand: None,
883 }
884 }
885
886 fn resources(workers: usize) -> ExecutionBudget {
887 ExecutionBudget {
888 workers,
889 memory_budget: None,
890 gpu_memory_budget: None,
891 can_spawn_threads: true,
892 }
893 }
894
895 fn gpu_probe(population: u64, bytes: usize) -> ProbeReport {
897 let mut demand = Demand::default();
898 demand.push("cells".to_owned(), bytes / 4);
899 ProbeReport {
900 demand: Some(demand),
901 ..probe(None, population, 0)
902 }
903 }
904
905 fn fixed(count: usize) -> Concurrency {
906 Concurrency::Fixed(NonZeroUsize::new(count).expect("a test count is above 0"))
907 }
908
909 #[test]
910 fn a_small_model_gets_one_thread_per_lane() {
911 let small = probe(Some(3), 4096, 1 << 20);
912 let layout = choose_layout(Concurrency::Auto, &resources(8), Backend::Cpu, &small, 100);
913 assert_eq!(layout, lanes(8, 1));
914 }
915
916 #[test]
917 fn a_wide_model_gets_wide_lanes() {
918 let wide = probe(Some(9), 0, 0);
919 let layout = choose_layout(Concurrency::Auto, &resources(8), Backend::Cpu, &wide, 100);
920 assert_eq!(
921 layout,
922 lanes(2, 3),
923 "9 jobs take 3 threads, and 8 workers make 2 such lanes"
924 );
925 let widest = probe(Some(64), 0, 0);
926 let layout = choose_layout(Concurrency::Auto, &resources(8), Backend::Cpu, &widest, 100);
927 assert_eq!(layout, lanes(1, 8), "a model as wide as the machine runs alone");
928 }
929
930 #[test]
931 fn a_model_without_jobs_is_sized_by_population() {
932 let unsplit = probe(None, 16 * 4096, 0);
933 let layout = choose_layout(Concurrency::Auto, &resources(8), Backend::Cpu, &unsplit, 100);
934 assert_eq!(layout, lanes(2, 4));
935 }
936
937 #[test]
938 fn lanes_are_capped_by_runs_and_memory() {
939 let small = probe(Some(1), 0, 1000);
940 let few_runs = choose_layout(Concurrency::Auto, &resources(8), Backend::Cpu, &small, 3);
941 assert_eq!(few_runs, lanes(3, 1));
942 let one_run = choose_layout(Concurrency::Auto, &resources(8), Backend::Cpu, &small, 1);
943 assert_eq!(one_run, lanes(1, 8), "a single lane takes every worker");
944
945 let budget = ExecutionBudget {
946 memory_budget: Some(4500),
947 ..resources(8)
948 };
949 let layout = choose_layout(Concurrency::Auto, &budget, Backend::Cpu, &small, 100);
950 assert_eq!(layout, lanes(4, 1));
951 assert_eq!(layout.projected_bytes(&small), 4000);
952 let tight = ExecutionBudget {
953 memory_budget: Some(10),
954 ..resources(8)
955 };
956 let layout = choose_layout(Concurrency::Auto, &tight, Backend::Cpu, &small, 100);
957 assert_eq!(layout, lanes(1, 8), "one run goes ahead even past the budget");
958 }
959
960 #[test]
961 fn a_fixed_lane_count_splits_the_workers() {
962 let small = probe(Some(1), 0, 0);
963 assert_eq!(
964 choose_layout(fixed(3), &resources(8), Backend::Cpu, &small, 100),
965 lanes(3, 2)
966 );
967 assert_eq!(
968 choose_layout(fixed(16), &resources(8), Backend::Cpu, &small, 100),
969 lanes(16, 1)
970 );
971 assert_eq!(
972 choose_layout(fixed(4), &resources(8), Backend::Cpu, &small, 2),
973 lanes(2, 2)
974 );
975 assert_eq!(
976 choose_layout(fixed(1), &resources(8), Backend::Cpu, &small, 100),
977 lanes(1, 8)
978 );
979 }
980
981 #[test]
982 fn a_target_without_threads_runs_one_lane() {
983 let small = probe(Some(1), 0, 0);
984 let threadless = ExecutionBudget {
985 can_spawn_threads: false,
986 ..resources(8)
987 };
988 assert_eq!(
989 choose_layout(fixed(4), &threadless, Backend::Cpu, &small, 100),
990 lanes(1, 8)
991 );
992 }
993
994 #[test]
995 fn gpu_tracks_are_sized_by_the_memory_budget_and_the_population() {
996 let small = gpu_probe(64 * 64, 1000);
997 let gpu_budget = |bytes| ExecutionBudget {
998 gpu_memory_budget: Some(bytes),
999 ..resources(8)
1000 };
1001 let layout = |concurrency, budget: &ExecutionBudget, probe: &ProbeReport, runs| {
1002 choose_layout(concurrency, budget, Backend::Gpu, probe, runs)
1003 };
1004 assert_eq!(layout(Concurrency::Auto, &gpu_budget(1 << 30), &small, 100), tracks(4));
1005 assert_eq!(layout(Concurrency::Auto, &gpu_budget(2500), &small, 100), tracks(2));
1006 assert_eq!(
1007 layout(Concurrency::Auto, &gpu_budget(500), &small, 100),
1008 tracks(1),
1009 "a run larger than the budget runs alone"
1010 );
1011 assert_eq!(layout(Concurrency::Auto, &gpu_budget(1 << 30), &small, 3), tracks(3));
1012 assert_eq!(layout(Concurrency::Auto, &resources(8), &small, 100), tracks(4));
1013
1014 let large = gpu_probe(LARGE_GPU_POPULATION, 1000);
1015 assert_eq!(layout(Concurrency::Auto, &gpu_budget(1 << 30), &large, 100), tracks(1));
1016
1017 assert_eq!(
1018 layout(fixed(6), &gpu_budget(2500), &large, 100),
1019 tracks(6),
1020 "a fixed count is kept"
1021 );
1022 assert_eq!(layout(fixed(6), &gpu_budget(2500), &small, 2), tracks(2));
1023 }
1024
1025 #[test]
1026 fn concurrency_reads_auto_or_a_positive_count() {
1027 assert_eq!("auto".parse(), Ok(Concurrency::Auto));
1028 assert_eq!("3".parse(), Ok(fixed(3)));
1029 assert!("0".parse::<Concurrency>().is_err());
1030 assert!("many".parse::<Concurrency>().is_err());
1031 assert_eq!(fixed(3).to_string(), "3");
1032 assert_eq!(Concurrency::Auto.to_string(), "auto");
1033 }
1034
1035 #[test]
1036 fn the_slice_grows_on_fast_steps_and_shrinks_on_slow_ones() {
1037 let mut slice = SliceSize::default();
1038 for _ in 0..6 {
1039 slice.adapt(Duration::from_micros(10));
1040 }
1041 assert_eq!(slice.steps, 64);
1042 slice.adapt(Duration::from_millis(20));
1043 assert_eq!(slice.steps, 64, "a slice near the target keeps its size");
1044 slice.adapt(Duration::from_millis(160));
1045 assert_eq!(slice.steps, 8, "a slice eight times too long shrinks eightfold");
1046 slice.adapt(Duration::from_secs(1));
1047 assert_eq!(slice.steps, 1, "a slice never drops below one step");
1048 }
1049
1050 #[test]
1051 fn the_reorder_buffer_hands_items_out_in_index_order() {
1052 let mut buffer = ReorderBuffer::default();
1053 buffer.insert(2, 'c');
1054 buffer.insert(1, 'b');
1055 assert_eq!(buffer.pop_ready(), None, "item 0 has not arrived");
1056 buffer.insert(0, 'a');
1057 let mut handed_out = Vec::new();
1058 while let Some(item) = buffer.pop_ready() {
1059 handed_out.push(item);
1060 }
1061 assert_eq!(handed_out, ['a', 'b', 'c']);
1062 buffer.insert(4, 'e');
1063 assert_eq!(buffer.pop_ready(), None, "item 3 has not arrived");
1064 assert_eq!(buffer.handed_out(), 3);
1065 }
1066
1067 #[test]
1068 fn pause_blocks_until_resume() {
1069 let control = SweepControl::new();
1070 control.pause();
1071 assert!(control.is_paused());
1072 let (sender, receiver) = mpsc::channel();
1073 let waiter = control.clone();
1074 let thread = std::thread::spawn(move || sender.send(waiter.proceed()));
1075 assert!(
1076 receiver.recv_timeout(Duration::from_millis(100)).is_err(),
1077 "a paused control holds the run"
1078 );
1079 control.resume();
1080 assert_eq!(receiver.recv_timeout(Duration::from_secs(10)), Ok(true));
1081 assert!(thread.join().is_ok_and(|sent| sent.is_ok()), "the waiter reported");
1082 assert!(control.proceed(), "a resumed control lets runs through");
1083 }
1084
1085 #[test]
1086 fn abort_releases_a_paused_run_and_is_final() {
1087 let control = SweepControl::new();
1088 control.pause();
1089 let (sender, receiver) = mpsc::channel();
1090 let waiter = control.clone();
1091 let thread = std::thread::spawn(move || sender.send(waiter.proceed()));
1092 assert!(receiver.recv_timeout(Duration::from_millis(50)).is_err());
1093 control.abort();
1094 assert_eq!(receiver.recv_timeout(Duration::from_secs(10)), Ok(false));
1095 assert!(thread.join().is_ok_and(|sent| sent.is_ok()), "the waiter reported");
1096 control.pause();
1097 control.resume();
1098 assert!(control.is_aborted(), "neither a pause nor a resume undoes an abort");
1099 assert!(!control.proceed());
1100 }
1101
1102 #[test]
1103 fn a_finished_run_waits_until_it_is_committed() {
1104 let active_runs = ActiveRuns::new();
1105 let run = |run_id| PlannedRun {
1106 run_id,
1107 config_id: 0,
1108 rep: run_id,
1109 seed: 1,
1110 };
1111 let (first, second, third) = (
1112 active_runs.watch(run(0), 10),
1113 active_runs.watch(run(1), 10),
1114 active_runs.watch(run(2), 10),
1115 );
1116 second.reach(4);
1117 assert_eq!(active_runs.list().len(), 3);
1118 assert_eq!(active_runs.list()[1].tick, 4);
1119
1120 third.finish();
1121 assert_eq!(active_runs.list().len(), 2, "a finished run is no longer in progress");
1122 assert_eq!(active_runs.waiting_count(), 1, "run 2 waits for runs 0 and 1");
1123 drop(first);
1124 assert_eq!(active_runs.waiting_count(), 1, "an abandoned run never waits");
1125 active_runs.mark_committed(2);
1126 assert_eq!(active_runs.waiting_count(), 0);
1127 drop(second);
1128 assert!(active_runs.list().is_empty());
1129 }
1130
1131 struct KeepingSink {
1133 committed: Vec<RunOutcome>,
1134 finished: usize,
1135 abort_after: Option<(usize, SweepControl)>,
1136 }
1137
1138 impl KeepingSink {
1139 fn new() -> Self {
1140 Self {
1141 committed: Vec::new(),
1142 finished: 0,
1143 abort_after: None,
1144 }
1145 }
1146 }
1147
1148 impl RunSink for KeepingSink {
1149 fn commit(&mut self, outcome: RunOutcome) -> std::io::Result<()> {
1150 self.committed.push(outcome);
1151 Ok(())
1152 }
1153
1154 fn finished(&mut self, _outcome: &RunOutcome) {
1155 self.finished += 1;
1156 if let Some((count, control)) = &self.abort_after
1157 && self.finished >= *count
1158 {
1159 control.abort();
1160 }
1161 }
1162 }
1163
1164 fn entry(id: &str) -> ModelEntry {
1165 example_models().get(id).cloned().expect("the model is registered")
1166 }
1167
1168 fn plan(entry: &ModelEntry, steps: u64, replicates: u64) -> (Plan, Arc<MeasurePlan>) {
1170 let mut spec = SweepSpec::new(entry.id().to_owned());
1171 spec.run.steps = steps;
1172 spec.run.replicates = replicates;
1173 spec.measure.stats_every = 3;
1174 spec.measure.series_every = 6;
1175 spec.fixed = vec![("grid_height".to_owned(), "24".to_owned())];
1176 spec.blocks = vec![BlockSpec {
1177 design: DesignKind::Factorial,
1178 factors: vec![FactorSpec::param(
1179 "grid_width",
1180 LevelSpec::Values(vec!["16".to_owned(), "20".to_owned(), "24".to_owned()]),
1181 )],
1182 design_seed: None,
1183 }];
1184 let plan = spec.plan(&entry.schema()).expect("a valid spec");
1185 let probe = ProbeReport::for_plan(entry, None, &plan).expect("the probe builds");
1186 let measure =
1187 MeasurePlan::new(plan.run_settings(), plan.measure_settings(), probe.columns).expect("the columns bind");
1188 (plan, Arc::new(measure))
1189 }
1190
1191 fn requests(plan: &Plan) -> Vec<RunRequest<'_>> {
1192 plan.runs().map(|run| RunRequest::planned(plan, run)).collect()
1193 }
1194
1195 fn run(entry: &ModelEntry, plan: &Plan, measure: &Arc<MeasurePlan>, layout: ExecutionLayout) -> Vec<RunOutcome> {
1197 let executor =
1198 Executor::new(entry, None, Arc::clone(measure), layout, SweepControl::new()).expect("the lane pools build");
1199 let mut sink = KeepingSink::new();
1200 let end = executor.run_batch(&requests(plan), &mut sink).expect("the batch runs");
1201 assert_eq!(end, BatchEnd::Complete);
1202 assert_eq!(sink.finished, sink.committed.len());
1203 sink.committed
1204 .into_iter()
1205 .map(|outcome| RunOutcome {
1206 build_ms: 0.0,
1207 wall_ms: 0.0,
1208 ..outcome
1209 })
1210 .collect()
1211 }
1212
1213 #[test]
1214 fn a_batch_commits_the_same_outcomes_in_request_order_at_any_lane_count() {
1215 let entry = entry("sir");
1216 let (plan, measure) = plan(&entry, 40, 3);
1217 let global = rayon::current_num_threads();
1218 let alone = run(&entry, &plan, &measure, lanes(1, global));
1219 assert_eq!(alone.len(), 9);
1220 let ids: Vec<u64> = alone.iter().map(|outcome| outcome.run.run_id).collect();
1221 assert_eq!(ids, (0..9).collect::<Vec<_>>());
1222 assert!(
1223 alone
1224 .iter()
1225 .all(|outcome| outcome.status == RunStatus::Ok && outcome.ticks == 40)
1226 );
1227 assert_eq!(alone[0].series.ticks(), [0, 6, 12, 18, 24, 30, 36, 40]);
1228 for layout in [lanes(1, 2), lanes(3, 1), lanes(4, 2)] {
1229 assert_eq!(run(&entry, &plan, &measure, layout), alone, "{layout:?}");
1230 }
1231 }
1232
1233 #[test]
1234 fn an_abort_ends_every_lane() {
1235 let entry = entry("game_of_life");
1236 let (plan, measure) = plan(&entry, 1_000_000, 4);
1237 let requests = requests(&plan);
1238 for layout in [lanes(1, rayon::current_num_threads()), lanes(3, 1)] {
1239 let control = SweepControl::new();
1240 let executor = Executor::new(&entry, None, Arc::clone(&measure), layout, control.clone())
1241 .expect("the lane pools build");
1242 let mut sink = KeepingSink::new();
1243 let aborter = control.clone();
1244 let timer = std::thread::spawn(move || {
1245 std::thread::sleep(Duration::from_millis(200));
1246 aborter.abort();
1247 });
1248 let end = executor.run_batch(&requests, &mut sink).expect("the batch runs");
1249 assert!(timer.join().is_ok(), "the timer thread finished");
1250 assert_eq!(end, BatchEnd::Aborted, "{layout:?}");
1251 assert!(sink.committed.is_empty(), "no run reached a million steps: {layout:?}");
1252 }
1253 }
1254
1255 struct AbortsAfterFirstRun {
1257 control: SweepControl,
1258 committed: usize,
1259 aborter: Option<std::thread::JoinHandle<Instant>>,
1261 }
1262
1263 impl RunSink for AbortsAfterFirstRun {
1264 fn commit(&mut self, _outcome: RunOutcome) -> std::io::Result<()> {
1265 self.committed += 1;
1266 Ok(())
1267 }
1268
1269 fn finished(&mut self, _outcome: &RunOutcome) {
1270 if self.aborter.is_none() {
1271 let control = self.control.clone();
1272 self.aborter = Some(std::thread::spawn(move || {
1273 std::thread::sleep(Duration::from_millis(50));
1274 let aborted_at = Instant::now();
1275 control.abort();
1276 aborted_at
1277 }));
1278 }
1279 }
1280 }
1281
1282 #[test]
1283 fn an_abort_lands_within_a_slice_of_a_heavy_run_after_a_light_one() {
1284 let entry = entry("game_of_life");
1285 let mut spec = SweepSpec::new("game_of_life");
1286 spec.run.steps = 20_000;
1287 spec.measure.stats_every = 1000;
1288 spec.measure.series_every = 0;
1289 spec.blocks = vec![BlockSpec {
1290 design: DesignKind::Zip,
1291 factors: ["grid_width", "grid_height"]
1292 .map(|id| FactorSpec::param(id, LevelSpec::Values(vec!["8".to_owned(), "2048".to_owned()])))
1293 .to_vec(),
1294 design_seed: None,
1295 }];
1296 let plan = spec.plan(&entry.schema()).expect("a valid spec");
1297 let probe = ProbeReport::for_plan(&entry, None, &plan).expect("the probe builds");
1298 let measure = Arc::new(
1299 MeasurePlan::new(plan.run_settings(), plan.measure_settings(), probe.columns).expect("the columns bind"),
1300 );
1301 let control = SweepControl::new();
1302 let layout = lanes(1, rayon::current_num_threads());
1303 let executor = Executor::new(&entry, None, measure, layout, control.clone()).expect("the lane's pool builds");
1304 let mut sink = AbortsAfterFirstRun {
1305 control,
1306 committed: 0,
1307 aborter: None,
1308 };
1309 let end = executor.run_batch(&requests(&plan), &mut sink).expect("the batch runs");
1310 let aborted_at = sink
1311 .aborter
1312 .take()
1313 .expect("the light run finished")
1314 .join()
1315 .expect("the aborting thread finished");
1316 assert_eq!(end, BatchEnd::Aborted);
1317 assert_eq!(sink.committed, 1, "only the light run finished");
1318 let latency = aborted_at.elapsed();
1319 assert!(
1320 latency < Duration::from_secs(5),
1321 "the heavy run stopped {latency:?} after the abort"
1322 );
1323 }
1324
1325 fn first_tick(active_runs: &ActiveRuns) -> Option<u64> {
1327 active_runs.list().first().map(|run| run.tick)
1328 }
1329
1330 #[test]
1331 fn a_paused_single_lane_holds_no_worker_of_the_global_pool() {
1332 let entry = entry("game_of_life");
1333 let (plan, measure) = plan(&entry, 1_000_000, 1);
1334 let control = SweepControl::new();
1335 let active_runs = ActiveRuns::new();
1336 let layout = lanes(1, rayon::current_num_threads());
1337 let executor = Executor::new(&entry, None, measure, layout, control.clone())
1338 .expect("the lane's pool builds")
1339 .with_active_runs(Some(active_runs.clone()));
1340 let pauser = control.clone();
1341 let checker = std::thread::spawn(move || {
1342 let deadline = Instant::now() + Duration::from_secs(60);
1343 while first_tick(&active_runs).is_none_or(|tick| tick == 0) {
1344 if Instant::now() > deadline {
1345 pauser.abort();
1346 return Err("the run never stepped");
1347 }
1348 std::thread::sleep(Duration::from_millis(1));
1349 }
1350 pauser.pause();
1351 let mut last = first_tick(&active_runs);
1353 loop {
1354 std::thread::sleep(Duration::from_millis(200));
1355 let now = first_tick(&active_runs);
1356 if now == last {
1357 break;
1358 }
1359 last = now;
1360 }
1361 let broadcast = std::thread::spawn(|| rayon::broadcast(|_| ()));
1362 let deadline = Instant::now() + Duration::from_secs(20);
1363 while !broadcast.is_finished() && Instant::now() < deadline {
1364 std::thread::sleep(Duration::from_millis(10));
1365 }
1366 let reached_every_worker = broadcast.is_finished();
1367 pauser.abort();
1368 if reached_every_worker {
1369 Ok(())
1370 } else {
1371 Err("a paused run held a worker of the global pool")
1372 }
1373 });
1374 let end = executor
1375 .run_batch(&requests(&plan), &mut KeepingSink::new())
1376 .expect("the batch runs");
1377 assert_eq!(checker.join().expect("the checking thread finished"), Ok(()));
1378 assert_eq!(end, BatchEnd::Aborted);
1379 }
1380
1381 struct PanickingSink;
1383
1384 impl RunSink for PanickingSink {
1385 fn commit(&mut self, _outcome: RunOutcome) -> std::io::Result<()> {
1386 Ok(())
1387 }
1388
1389 fn finished(&mut self, _outcome: &RunOutcome) {
1390 panic!("the sink cannot take a run");
1391 }
1392 }
1393
1394 #[test]
1395 fn a_panicking_sink_aborts_every_lane() {
1396 let entry = entry("game_of_life");
1397 let (plan, measure) = plan(&entry, 200, 4);
1398 let requests = requests(&plan);
1399 let control = SweepControl::new();
1400 let executor = Executor::new(&entry, None, Arc::clone(&measure), lanes(3, 1), control.clone())
1401 .expect("the lane pools build");
1402 let batch = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
1403 executor.run_batch(&requests, &mut PanickingSink)
1404 }));
1405 assert!(batch.is_err(), "the sink's panic reaches the caller");
1406 assert!(control.is_aborted(), "the lanes were told to stop");
1407 }
1408
1409 #[test]
1410 fn a_panicking_lane_ends_the_batch_with_an_error() {
1411 let plain = entry("game_of_life");
1412 let (plan, measure) = plan(&plain, 20, 2);
1413 let panicking = plain.wrap_factory(|_| {
1415 Arc::new(
1416 |_params: &[ParamValue], _seed: Option<u64>, _gpu: Option<&GpuContext>| -> Result<ModelState, Fault> {
1417 panic!("the lane cannot build a run")
1418 },
1419 )
1420 });
1421 let control = SweepControl::new();
1422 let executor =
1423 Executor::new(&panicking, None, measure, lanes(3, 1), control.clone()).expect("the lane pools build");
1424 let batch = executor.run_batch(&requests(&plan), &mut KeepingSink::new());
1425 assert!(matches!(batch, Err(ExecutionError::LanePanicked)), "{batch:?}");
1426 assert!(control.is_aborted(), "the other lanes were told to stop");
1427 }
1428
1429 #[test]
1430 fn runs_committed_before_an_abort_are_a_prefix_of_the_requests() {
1431 let entry = entry("game_of_life");
1432 let (plan, measure) = plan(&entry, 2000, 4);
1433 let requests = requests(&plan);
1434 let control = SweepControl::new();
1435 let executor = Executor::new(&entry, None, Arc::clone(&measure), lanes(3, 1), control.clone())
1436 .expect("the lane pools build");
1437 let mut sink = KeepingSink::new();
1438 sink.abort_after = Some((2, control));
1439 let end = executor.run_batch(&requests, &mut sink).expect("the batch runs");
1440 assert_eq!(end, BatchEnd::Aborted);
1441 let ids: Vec<u64> = sink.committed.iter().map(|outcome| outcome.run.run_id).collect();
1442 assert_eq!(ids, (0..ids.len() as u64).collect::<Vec<_>>());
1443 assert!(ids.len() < requests.len());
1444 }
1445}