1use std::fmt;
7
8use crate::explore::design_rng::DesignRng;
9use crate::explore::factor::{Factor, FactorDomain, FactorLevel, FactorSlot};
10use crate::explore::plan::Config;
11use crate::params::ParamValue;
12
13pub const MAX_CONFIGS: usize = 1 << 24;
15
16#[derive(Debug, Clone, Default, PartialEq, Eq)]
18pub enum DesignKind {
19 #[default]
21 Factorial,
22 Zip,
24 Random {
26 samples: usize,
28 },
29 LatinHypercube {
36 samples: usize,
38 },
39 Table {
43 text: String,
45 },
46}
47
48impl DesignKind {
49 pub fn as_str(&self) -> &'static str {
51 match self {
52 Self::Factorial => "factorial",
53 Self::Zip => "zip",
54 Self::Random { .. } => "random",
55 Self::LatinHypercube { .. } => "lhs",
56 Self::Table { .. } => "table",
57 }
58 }
59
60 pub fn is_sampled(&self) -> bool {
62 matches!(self, Self::Random { .. } | Self::LatinHypercube { .. })
63 }
64}
65
66#[derive(Debug, Clone, PartialEq)]
68pub struct Block {
69 pub design: DesignKind,
71 pub factors: Vec<Factor>,
73 pub design_seed: u64,
75}
76
77#[derive(Debug, Clone, PartialEq, Eq)]
79pub enum DesignError {
80 UnequalLengths {
82 lengths: Vec<usize>,
84 },
85 TooManyConfigs,
87 NoSamples,
89 NoFactors,
91 UnlistedLevels {
93 factor_index: usize,
95 },
96 NoLevels {
98 factor_index: usize,
100 },
101}
102
103impl fmt::Display for DesignError {
104 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
105 match self {
106 Self::UnequalLengths { lengths } => {
107 let lengths: Vec<String> = lengths.iter().map(ToString::to_string).collect();
108 write!(
109 f,
110 "zip needs every factor to have the same number of levels, got {}",
111 lengths.join(", ")
112 )
113 }
114 Self::TooManyConfigs => write!(f, "design takes the plan past {MAX_CONFIGS} configs"),
115 Self::NoSamples => write!(f, "a random or Latin hypercube design needs at least 1 sample"),
116 Self::NoFactors => write!(f, "a random or Latin hypercube design needs at least 1 factor"),
117 Self::UnlistedLevels { factor_index } => {
118 write!(
119 f,
120 "range of factor {factor_index} needs a step, except in a random or Latin hypercube design"
121 )
122 }
123 Self::NoLevels { factor_index } => write!(f, "factor {factor_index} has no levels"),
124 }
125 }
126}
127
128impl std::error::Error for DesignError {}
129
130pub fn generate(block: &Block, base: &Config) -> Result<Vec<Config>, DesignError> {
141 generate_within(block, base, MAX_CONFIGS)
142}
143
144pub(crate) fn generate_within(block: &Block, base: &Config, limit: usize) -> Result<Vec<Config>, DesignError> {
150 if let Some(factor_index) = block
151 .factors
152 .iter()
153 .position(|factor| factor.levels().is_some_and(<[FactorLevel]>::is_empty))
154 {
155 return Err(DesignError::NoLevels { factor_index });
156 }
157 match block.design {
158 DesignKind::Factorial => factorial(&listed_levels(&block.factors)?, &block.factors, base, limit),
159 DesignKind::Zip | DesignKind::Table { .. } => zip(&listed_levels(&block.factors)?, &block.factors, base, limit),
160 DesignKind::Random { samples } => {
161 check_samples(samples, &block.factors, limit)?;
162 Ok(random(
163 &block.factors,
164 base,
165 samples,
166 &mut DesignRng::new(block.design_seed),
167 ))
168 }
169 DesignKind::LatinHypercube { samples } => {
170 check_samples(samples, &block.factors, limit)?;
171 Ok(latin_hypercube(
172 &block.factors,
173 base,
174 samples,
175 &mut DesignRng::new(block.design_seed),
176 ))
177 }
178 }
179}
180
181fn listed_levels(factors: &[Factor]) -> Result<Vec<&[FactorLevel]>, DesignError> {
183 factors
184 .iter()
185 .enumerate()
186 .map(|(factor_index, factor)| factor.levels().ok_or(DesignError::UnlistedLevels { factor_index }))
187 .collect()
188}
189
190fn check_samples(samples: usize, factors: &[Factor], limit: usize) -> Result<(), DesignError> {
191 if samples == 0 {
192 return Err(DesignError::NoSamples);
193 }
194 if samples > limit {
195 return Err(DesignError::TooManyConfigs);
196 }
197 if factors.is_empty() {
198 return Err(DesignError::NoFactors);
199 }
200 Ok(())
201}
202
203fn factorial(
204 levels: &[&[FactorLevel]],
205 factors: &[Factor],
206 base: &Config,
207 limit: usize,
208) -> Result<Vec<Config>, DesignError> {
209 let count = levels
210 .iter()
211 .try_fold(1_usize, |product, levels| product.checked_mul(levels.len()))
212 .filter(|&count| count <= limit)
213 .ok_or(DesignError::TooManyConfigs)?;
214 let mut configs = Vec::with_capacity(count);
215 let mut positions = vec![0; factors.len()];
216 for _ in 0..count {
217 let mut config = base.clone();
218 for ((factor, levels), &position) in factors.iter().zip(levels).zip(&positions) {
219 factor.apply(&levels[position], &mut config);
220 }
221 configs.push(config);
222 for (position, levels) in positions.iter_mut().zip(levels).rev() {
223 *position += 1;
224 if *position < levels.len() {
225 break;
226 }
227 *position = 0;
228 }
229 }
230 Ok(configs)
231}
232
233fn zip(levels: &[&[FactorLevel]], factors: &[Factor], base: &Config, limit: usize) -> Result<Vec<Config>, DesignError> {
234 let lengths: Vec<usize> = levels.iter().map(|levels| levels.len()).collect();
235 let count = lengths.first().copied().unwrap_or(1);
237 if lengths.iter().any(|&length| length != count) {
238 return Err(DesignError::UnequalLengths { lengths });
239 }
240 if count > limit {
241 return Err(DesignError::TooManyConfigs);
242 }
243 Ok((0..count)
244 .map(|position| {
245 let mut config = base.clone();
246 for (factor, levels) in factors.iter().zip(levels) {
247 factor.apply(&levels[position], &mut config);
248 }
249 config
250 })
251 .collect())
252}
253
254fn random(factors: &[Factor], base: &Config, samples: usize, rng: &mut DesignRng) -> Vec<Config> {
256 (0..samples)
257 .map(|_| {
258 let mut config = base.clone();
259 for factor in factors {
260 let level = match &factor.domain {
261 FactorDomain::Levels(levels) => levels[rng.index(levels.len() as u64) as usize].clone(),
262 &FactorDomain::Continuous { min, max } => continuous_level(min, max, rng.unit_f64()),
263 &FactorDomain::WholeNumbers { min, max } => {
264 whole_number_level(factor.slot, rng.whole_number(min, max))
265 }
266 };
267 factor.apply(&level, &mut config);
268 }
269 config
270 })
271 .collect()
272}
273
274fn latin_hypercube(factors: &[Factor], base: &Config, samples: usize, rng: &mut DesignRng) -> Vec<Config> {
279 let mut configs = vec![base.clone(); samples];
280 let strata = samples as f64;
281 for factor in factors {
282 let order = rng.permutation(samples);
283 for (config, &stratum) in configs.iter_mut().zip(&order) {
284 let level = match &factor.domain {
285 FactorDomain::Levels(levels) => {
286 levels[balanced_level(stratum, samples, levels.len() as u128) as usize].clone()
287 }
288 &FactorDomain::Continuous { min, max } => {
289 continuous_level(min, max, (stratum as f64 + rng.unit_f64()) / strata)
290 }
291 &FactorDomain::WholeNumbers { min, max } => {
292 let count = u128::from(max - min) + 1;
293 whole_number_level(factor.slot, min + balanced_level(stratum, samples, count) as u64)
294 }
295 };
296 factor.apply(&level, config);
297 }
298 }
299 configs
300}
301
302fn balanced_level(stratum: usize, samples: usize, count: u128) -> u128 {
306 stratum as u128 * count / samples as u128
307}
308
309pub(crate) fn continuous_value(min: f64, max: f64, fraction: f64) -> f32 {
311 let value = (min + fraction * (max - min)) as f32;
312 value.clamp(min as f32, max as f32)
313}
314
315pub(crate) fn continuous_level(min: f64, max: f64, fraction: f64) -> FactorLevel {
317 FactorLevel::Param(ParamValue::F32(continuous_value(min, max, fraction)))
318}
319
320pub(crate) fn whole_number_level(slot: FactorSlot, value: u64) -> FactorLevel {
322 match slot {
323 FactorSlot::Param(_) => FactorLevel::Param(ParamValue::U32(value as u32)),
324 FactorSlot::Action(_) => FactorLevel::Tick(value),
325 }
326}
327
328#[cfg(test)]
329mod tests {
330 use std::collections::BTreeMap;
331
332 use super::{Block, DesignError, DesignKind, generate};
333 use crate::explore::factor::{Factor, FactorDomain, FactorLevel, FactorSlot};
334 use crate::explore::plan::Config;
335 use crate::params::ParamValue;
336
337 fn factor(slot: usize, values: &[u32]) -> Factor {
338 Factor {
339 slot: FactorSlot::Param(slot),
340 domain: FactorDomain::Levels(
341 values
342 .iter()
343 .map(|&value| FactorLevel::Param(ParamValue::U32(value)))
344 .collect(),
345 ),
346 }
347 }
348
349 fn continuous(slot: usize, min: f64, max: f64) -> Factor {
350 Factor {
351 slot: FactorSlot::Param(slot),
352 domain: FactorDomain::Continuous { min, max },
353 }
354 }
355
356 fn whole_numbers(slot: FactorSlot, min: u64, max: u64) -> Factor {
357 Factor {
358 slot,
359 domain: FactorDomain::WholeNumbers { min, max },
360 }
361 }
362
363 fn base() -> Config {
364 Config {
365 block: 0,
366 params: vec![ParamValue::U32(0); 3],
367 action_ticks: vec![0],
368 }
369 }
370
371 fn block(design: DesignKind, factors: Vec<Factor>) -> Block {
372 Block {
373 design,
374 factors,
375 design_seed: 11,
376 }
377 }
378
379 fn rows(design: DesignKind, factors: Vec<Factor>) -> Result<Vec<Vec<u32>>, DesignError> {
381 let configs = generate(&block(design, factors), &base())?;
382 Ok(configs
383 .into_iter()
384 .map(|config| {
385 config
386 .params
387 .into_iter()
388 .map(|value| match value {
389 ParamValue::U32(number) => number,
390 other => panic!("every test value is a u32, got {other:?}"),
391 })
392 .collect()
393 })
394 .collect())
395 }
396
397 fn column(configs: &[Config], slot: usize) -> Vec<f64> {
399 configs
400 .iter()
401 .map(|config| match &config.params[slot] {
402 &ParamValue::F32(value) => f64::from(value),
403 &ParamValue::U32(value) => f64::from(value),
404 other => panic!("every test value is a number, got {other:?}"),
405 })
406 .collect()
407 }
408
409 fn counts(configs: &[Config], slot: usize) -> BTreeMap<u32, usize> {
411 let mut counts = BTreeMap::new();
412 for value in column(configs, slot) {
413 *counts.entry(value as u32).or_default() += 1;
414 }
415 counts
416 }
417
418 fn lhs(samples: usize, factors: Vec<Factor>, design_seed: u64) -> Vec<Config> {
419 let block = Block {
420 design_seed,
421 ..block(DesignKind::LatinHypercube { samples }, factors)
422 };
423 generate(&block, &base()).expect("a valid design")
424 }
425
426 #[test]
427 fn a_factorial_varies_the_last_factor_fastest() {
428 let configs = rows(
429 DesignKind::Factorial,
430 vec![factor(0, &[1, 2]), factor(2, &[10, 20, 30])],
431 );
432 assert_eq!(
433 configs,
434 Ok(vec![
435 vec![1, 0, 10],
436 vec![1, 0, 20],
437 vec![1, 0, 30],
438 vec![2, 0, 10],
439 vec![2, 0, 20],
440 vec![2, 0, 30],
441 ])
442 );
443 }
444
445 #[test]
446 fn a_factorial_has_the_product_of_the_level_counts() {
447 let configs = rows(
448 DesignKind::Factorial,
449 vec![factor(0, &[1, 2]), factor(1, &[1, 2, 3]), factor(2, &[1, 2, 3, 4])],
450 )
451 .expect("24 configs");
452 assert_eq!(configs.len(), 24);
453 let mut distinct = configs.clone();
454 distinct.sort();
455 distinct.dedup();
456 assert_eq!(distinct.len(), 24, "every combination appears once");
457 }
458
459 #[test]
460 fn a_block_with_no_factors_is_its_base() {
461 assert_eq!(rows(DesignKind::Factorial, Vec::new()), Ok(vec![vec![0, 0, 0]]));
462 assert_eq!(rows(DesignKind::Zip, Vec::new()), Ok(vec![vec![0, 0, 0]]));
463 }
464
465 #[test]
466 fn a_zip_pairs_levels_by_position() {
467 let configs = rows(DesignKind::Zip, vec![factor(0, &[1, 2, 3]), factor(1, &[10, 20, 30])]);
468 assert_eq!(configs, Ok(vec![vec![1, 10, 0], vec![2, 20, 0], vec![3, 30, 0]]));
469 }
470
471 #[test]
472 fn a_zip_of_unequal_lengths_is_refused() {
473 let error = rows(DesignKind::Zip, vec![factor(0, &[1, 2, 3]), factor(1, &[10, 20])]);
474 assert_eq!(error, Err(DesignError::UnequalLengths { lengths: vec![3, 2] }));
475 }
476
477 #[test]
478 fn a_design_past_the_limit_is_refused() {
479 let wide: Vec<u32> = (0..4096).collect();
480 let error = rows(
481 DesignKind::Factorial,
482 vec![factor(0, &wide), factor(1, &wide), factor(2, &[1, 2])],
483 );
484 assert_eq!(error, Err(DesignError::TooManyConfigs));
485 }
486
487 #[test]
488 fn a_latin_hypercube_puts_one_sample_in_each_stratum() {
489 let samples = 50;
490 let configs = lhs(samples, vec![continuous(0, 0.0, 50.0), continuous(2, 10.0, 20.0)], 7);
491 assert_eq!(configs.len(), samples);
492 for (slot, min, max) in [(0, 0.0, 50.0), (2, 10.0, 20.0)] {
493 let mut strata: Vec<usize> = column(&configs, slot)
494 .into_iter()
495 .map(|value| ((value - min) / (max - min) * samples as f64).floor() as usize)
496 .collect();
497 strata.sort_unstable();
498 assert_eq!(strata, (0..samples).collect::<Vec<_>>(), "parameter {slot}");
499 }
500 assert!(
501 configs.iter().all(|config| config.params[1] == ParamValue::U32(0)),
502 "a parameter no factor varies keeps its base value"
503 );
504 }
505
506 #[test]
507 fn a_latin_hypercube_balances_discrete_levels() {
508 let options = Factor {
509 slot: FactorSlot::Param(0),
510 domain: FactorDomain::Levels(
511 (0..3)
512 .map(|option| FactorLevel::Param(ParamValue::Choice(option)))
513 .collect(),
514 ),
515 };
516 let configs = lhs(10, vec![options], 3);
517 let mut chosen = [0; 3];
518 for config in &configs {
519 let ParamValue::Choice(option) = config.params[0] else {
520 panic!("the factor writes a choice");
521 };
522 chosen[option] += 1;
523 }
524 chosen.sort_unstable();
525 assert_eq!(chosen, [3, 3, 4], "ten samples over three options");
526
527 let configs = lhs(8, vec![whole_numbers(FactorSlot::Param(1), 1, 4)], 3);
528 assert_eq!(counts(&configs, 1), BTreeMap::from([(1, 2), (2, 2), (3, 2), (4, 2)]));
529
530 let configs = lhs(8, vec![whole_numbers(FactorSlot::Action(0), 100, 103)], 3);
531 let mut ticks: Vec<u64> = configs.iter().map(|config| config.action_ticks[0]).collect();
532 ticks.sort_unstable();
533 assert_eq!(
534 ticks,
535 [100, 100, 101, 101, 102, 102, 103, 103],
536 "a tick range is discrete too"
537 );
538 }
539
540 #[test]
541 fn a_latin_hypercube_is_reproducible_from_its_seed() {
542 let factors = || vec![continuous(0, 0.0, 1.0), whole_numbers(FactorSlot::Param(1), 1, 1000)];
543 assert_eq!(lhs(20, factors(), 42), lhs(20, factors(), 42));
544 }
545
546 #[test]
547 fn different_design_seeds_give_different_designs() {
548 let factors = || vec![continuous(0, 0.0, 1.0), whole_numbers(FactorSlot::Param(1), 1, 1000)];
549 assert_ne!(lhs(20, factors(), 42), lhs(20, factors(), 43));
550 let random = |design_seed| {
551 let block = Block {
552 design_seed,
553 ..block(DesignKind::Random { samples: 20 }, factors())
554 };
555 generate(&block, &base()).expect("a valid design")
556 };
557 assert_eq!(random(42), random(42));
558 assert_ne!(random(42), random(43));
559 }
560
561 #[test]
562 fn random_samples_stay_within_bounds() {
563 let block = block(
564 DesignKind::Random { samples: 2000 },
565 vec![
566 continuous(0, 0.1, 0.2),
567 whole_numbers(FactorSlot::Param(1), 0, u64::from(u32::MAX)),
568 factor(2, &[5, 6]),
569 whole_numbers(FactorSlot::Action(0), 10, 20),
570 ],
571 );
572 let configs = generate(&block, &base()).expect("a valid design");
573 assert_eq!(configs.len(), 2000);
574 let lowest = f64::from(0.1_f32);
575 let highest = f64::from(0.2_f32);
576 assert!(
577 column(&configs, 0)
578 .iter()
579 .all(|&value| (lowest..=highest).contains(&value))
580 );
581 assert!(
582 column(&configs, 1).iter().any(|&value| value > f64::from(u32::MAX / 2)),
583 "a full span reaches its upper half"
584 );
585 assert_eq!(counts(&configs, 2).keys().copied().collect::<Vec<_>>(), [5, 6]);
586 assert!(configs.iter().all(|config| (10..=20).contains(&config.action_ticks[0])));
587 }
588
589 #[test]
590 fn a_design_that_cannot_draw_is_refused() {
591 let unlisted = || vec![continuous(0, 0.0, 1.0)];
592 assert_eq!(
593 generate(&block(DesignKind::LatinHypercube { samples: 0 }, unlisted()), &base()),
594 Err(DesignError::NoSamples)
595 );
596 assert_eq!(
597 generate(&block(DesignKind::Random { samples: 4 }, Vec::new()), &base()),
598 Err(DesignError::NoFactors)
599 );
600 assert_eq!(
601 generate(&block(DesignKind::Factorial, unlisted()), &base()),
602 Err(DesignError::UnlistedLevels { factor_index: 0 })
603 );
604 }
605
606 #[test]
609 fn a_factor_with_no_levels_is_refused_by_every_design() {
610 for design in [
611 DesignKind::Factorial,
612 DesignKind::Zip,
613 DesignKind::Random { samples: 4 },
614 DesignKind::LatinHypercube { samples: 4 },
615 ] {
616 assert_eq!(
617 rows(design.clone(), vec![factor(0, &[1, 2]), factor(1, &[])]),
618 Err(DesignError::NoLevels { factor_index: 1 }),
619 "{design:?}"
620 );
621 }
622 }
623
624 #[test]
626 fn a_latin_hypercube_takes_the_lowest_level_of_each_stratum() {
627 let levels = || vec![whole_numbers(FactorSlot::Param(1), 1, 100)];
628 let expected: BTreeMap<u32, usize> = (0..40).map(|stratum| (1 + stratum * 100 / 40, 1)).collect();
629 for design_seed in [3, 42] {
630 assert_eq!(counts(&lhs(40, levels(), design_seed), 1), expected);
631 }
632 }
633}