Skip to main content

henad_core/explore/
design.rs

1//! Designs that combine the factors of a block into configs.
2//!
3//! A factorial or zip design combines listed levels. A random or Latin hypercube design samples, drawing each level
4//! from a design seed, and a table design reads every config from a table of values.
5
6use 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
13/// Maximum number of configs in one plan, over all its blocks.
14pub const MAX_CONFIGS: usize = 1 << 24;
15
16/// Rule that combines the factors of a block.
17#[derive(Debug, Clone, Default, PartialEq, Eq)]
18pub enum DesignKind {
19    /// Every combination of levels. The first factor varies slowest and the last fastest.
20    #[default]
21    Factorial,
22    /// Levels paired by position, so config `i` takes level `i` of every factor.
23    Zip,
24    /// `samples` configs, each factor's value drawn independently from its whole domain.
25    Random {
26        /// Number of configs to draw.
27        samples: usize,
28    },
29    /// `samples` configs in a Latin hypercube, each factor's domain split into `samples` strata.
30    ///
31    /// Every stratum of a continuous factor holds one sample, at a random point inside it. A stratum of a factor
32    /// with `m` levels takes the lowest level it covers, and each level is then taken `samples / m` times, rounded
33    /// down or up. The levels such a factor takes therefore depend on `samples` alone, and the design seed decides
34    /// only which config takes each level.
35    LatinHypercube {
36        /// Number of configs, and of strata in each factor's domain.
37        samples: usize,
38    },
39    /// One config per row of a comma-separated table, whose header lists parameter ids or `action.<name>`.
40    ///
41    /// [`crate::explore::design_csv`] reads the table.
42    Table {
43        /// Text of the table, header included.
44        text: String,
45    },
46}
47
48impl DesignKind {
49    /// Returns the design's name in a spec file, as in `lhs`.
50    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    /// Returns whether the design draws its configs from a design seed.
61    pub fn is_sampled(&self) -> bool {
62        matches!(self, Self::Random { .. } | Self::LatinHypercube { .. })
63    }
64}
65
66/// Resolved factors combined under one design.
67#[derive(Debug, Clone, PartialEq)]
68pub struct Block {
69    /// Design that combines the factors into configs.
70    pub design: DesignKind,
71    /// Factors of the block, or the columns of its table.
72    pub factors: Vec<Factor>,
73    /// Seed of a sampled design's draws. Other designs draw nothing.
74    pub design_seed: u64,
75}
76
77/// A block whose design cannot combine its factors.
78#[derive(Debug, Clone, PartialEq, Eq)]
79pub enum DesignError {
80    /// A zip over factors with different numbers of levels.
81    UnequalLengths {
82        /// Number of levels of each factor, in factor order.
83        lengths: Vec<usize>,
84    },
85    /// A design whose configs would take the plan past [`MAX_CONFIGS`].
86    TooManyConfigs,
87    /// A random or Latin hypercube design of 0 samples.
88    NoSamples,
89    /// A random or Latin hypercube design with no factor to draw.
90    NoFactors,
91    /// Factor `factor_index` of a factorial or zip design, a range with no listed levels.
92    UnlistedLevels {
93        /// Index of the factor in its block, counting from 0.
94        factor_index: usize,
95    },
96    /// Factor `factor_index`, a list with no levels.
97    NoLevels {
98        /// Index of the factor in its block, counting from 0.
99        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
130/// Returns the configs of `block`, each a copy of `base` with one level of every factor written into it.
131///
132/// A factorial or zip block with no factors returns `base` alone. Note that when two factors write one slot, the later
133/// factor's level is the one kept.
134///
135/// # Errors
136///
137/// Returns [`DesignError`] for a factor with an empty list of levels, a zip over factors of unequal length, a
138/// sampled design with no samples or no factors, a whole range in a design that lists levels, or more than
139/// [`MAX_CONFIGS`] configs.
140pub fn generate(block: &Block, base: &Config) -> Result<Vec<Config>, DesignError> {
141    generate_within(block, base, MAX_CONFIGS)
142}
143
144/// Returns the configs of `block` as [`generate`] does, rejecting more than `limit` of them before any is built.
145///
146/// # Errors
147///
148/// Returns [`DesignError`] as [`generate`] does, with [`DesignError::TooManyConfigs`] past `limit`.
149pub(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
181/// Returns the listed levels of every factor.
182fn 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    // A zip of no factors yields `base` alone.
236    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
254/// Returns `samples` configs, drawing every factor's value independently.
255fn 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
274/// Returns `samples` configs, where each factor takes its strata in its own order.
275///
276/// The draws go factor by factor, a permutation of the strata first, then one position within each stratum for a
277/// continuous factor.
278fn 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
302/// Returns the level, out of `count` levels, that stratum `stratum` of `samples` falls on.
303///
304/// Each level gets `samples / count` strata, rounded down or up.
305fn balanced_level(stratum: usize, samples: usize, count: u128) -> u128 {
306    stratum as u128 * count / samples as u128
307}
308
309/// Returns the `f32` value at fraction `fraction` of the way from `min` to `max`, clamped to the rounded ends.
310pub(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
315/// Returns [`continuous_value`] as a level.
316pub(crate) fn continuous_level(min: f64, max: f64, fraction: f64) -> FactorLevel {
317    FactorLevel::Param(ParamValue::F32(continuous_value(min, max, fraction)))
318}
319
320/// Returns whole number `value` as a level of the factor writing `slot`.
321pub(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    /// Returns the configs of a block as rows of `u32` values.
380    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    /// Returns parameter `slot` of every config as an `f64`.
398    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    /// Returns the number of configs holding each value of parameter `slot`.
410    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    /// Every design rejects a factor with no levels. Without the rejection, a sampled design would index into the
607    /// empty list and panic, and a listed design would yield no configs.
608    #[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    /// A stratum of a discrete factor takes the lowest level it covers, whatever the design seed.
625    #[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}