Skip to main content

henad_models/
sir.rs

1//! The SIR (susceptible, infected, recovered) epidemic as a [`GridModel`] on a torus.
2//!
3//! Each infected cell among a susceptible cell's eight neighbours infects it with probability `infection_rate`,
4//! independently of the others. An infected cell recovers with probability `recovery_rate` each tick.
5
6use henad_compute::cpu::primitives::chunked::{STATS_CHUNK, reduce_chunks};
7use henad_core::action::ActionDescriptor;
8use henad_core::authoring::model::grid_model::GridModel;
9use henad_core::authoring::primitives::rng::{below, next_bits, next_float};
10use henad_core::grid::Grid2D;
11use henad_core::helpers::{extract_f32, f32_param};
12use henad_core::params::{ParamDescriptor, ParamValue};
13use henad_core::topology::NeighborhoodKind;
14use henad_core::view::{StatDescriptor, StatValue};
15
16const S: u8 = 0;
17const I: u8 = 1;
18const R: u8 = 2;
19
20henad_core::params! {
21    const INFECTION_RATE = f32_param("infection_rate", "Infection Rate", 0.3, 0.0, 1.0, Some(0.01));
22    const RECOVERY_RATE = f32_param("recovery_rate", "Recovery Rate", 0.05, 0.0, 1.0, Some(0.01));
23    const INITIAL_INFECTED_PCT =
24        f32_param("initial_infected_pct", "Initial Infected", 0.01, 0.0, 1.0, Some(0.001))
25            .on_reload()
26            .percent();
27}
28
29henad_core::actions! {
30    const SEED_OUTBREAK = ActionDescriptor::new("seed_outbreak", "Seed outbreak");
31}
32
33/// Cell colours, indexed by state: susceptible, infected, then recovered.
34pub const PALETTE: [[u8; 4]; 3] = [
35    [0x00, 0x7A, 0xF5, 0xFF], // S - blue
36    [0xE4, 0x37, 0x48, 0xFF], // I - red
37    [0x80, 0x80, 0x80, 0xFF], // R - gray
38];
39
40/// The SIR epidemic as a [`GridModel`].
41#[derive(Debug)]
42pub struct SirGridModel;
43
44/// Rates of [`SirGridModel`], read once per tick.
45#[derive(Debug)]
46pub struct SirParams {
47    infection_rate: f32,
48    recovery_rate: f32,
49}
50
51impl GridModel for SirGridModel {
52    const NAME: &'static str = "SIR Epidemic";
53    const ID: &'static str = "sir";
54    const DESCRIPTION: &'static str = "Classic SIR compartmental model on a 2D grid with Moore neighborhood";
55    const PALETTE: &'static [[u8; 4]] = &PALETTE;
56    const NEIGHBORHOOD: NeighborhoodKind = NeighborhoodKind::Moore;
57    // --8<-- [start:stat_descriptors]
58    const STATS: &'static [StatDescriptor] = &[
59        StatDescriptor::new("Susceptible", PALETTE[0]),
60        StatDescriptor::new("Infected", PALETTE[1]),
61        StatDescriptor::new("Recovered", PALETTE[2]),
62    ];
63    // --8<-- [end:stat_descriptors]
64    const ACTIONS: &'static [ActionDescriptor] = ACTION_SPECS;
65    type Params = SirParams;
66
67    fn param_descriptors() -> Vec<ParamDescriptor> {
68        descriptors()
69    }
70
71    // --8<-- [start:from_params]
72    fn from_params(params: &[ParamValue]) -> SirParams {
73        SirParams {
74            infection_rate: extract_f32(params, INFECTION_RATE, 0.3),
75            recovery_rate: extract_f32(params, RECOVERY_RATE, 0.05),
76        }
77    }
78    // --8<-- [end:from_params]
79
80    fn init(grid: &mut Grid2D<u8>, params: &[ParamValue], rng: &mut u64) {
81        let initial_pct = extract_f32(params, INITIAL_INFECTED_PCT, 0.01);
82        let threshold = (initial_pct * u32::MAX as f32) as u32;
83        for cell in grid.current_mut().iter_mut() {
84            *cell = if below(next_bits(rng), threshold) { I } else { S };
85        }
86    }
87
88    fn step_cell(cell: u8, neighbors: &[u8], params: &SirParams, rng: &mut u64) -> u8 {
89        match cell {
90            S => {
91                let infected_count = neighbors.iter().filter(|&&n| n == I).count();
92                if infected_count > 0 {
93                    let prob_safe = (1.0 - params.infection_rate).powi(infected_count as i32);
94                    // A rate of 1 always infects under `>=`. The draw is half-open and can be zero,
95                    // and `>` would let a zero draw escape.
96                    if next_float(rng, 1.0) >= prob_safe { I } else { S }
97                } else {
98                    S
99                }
100            }
101            I => {
102                if next_float(rng, 1.0) < params.recovery_rate {
103                    R
104                } else {
105                    I
106                }
107            }
108            _ => cell,
109        }
110    }
111
112    #[expect(clippy::single_match, reason = "for future multi-action extendability")]
113    fn act(action: usize, grid: &mut Grid2D<u8>, params: &[ParamValue], rng: &mut u64) {
114        match action {
115            SEED_OUTBREAK => seed_outbreak(grid, params, rng),
116            _ => {}
117        }
118    }
119
120    // --8<-- [start:stats]
121    fn stats(grid: &Grid2D<u8>) -> Vec<StatValue> {
122        let (s, i, r) = count_sir(grid.current());
123        vec![
124            StatValue::Scalar(s as f64),
125            StatValue::Scalar(i as f64),
126            StatValue::Scalar(r as f64),
127        ]
128    }
129    // --8<-- [end:stats]
130}
131
132/// Counts S, I and R in one pass over `cells`.
133fn count_sir_seq(cells: &[u8]) -> (u64, u64, u64) {
134    let (mut s, mut i, mut r) = (0u64, 0u64, 0u64);
135    for &cell in cells {
136        match cell {
137            S => s += 1,
138            I => i += 1,
139            _ => r += 1,
140        }
141    }
142    (s, i, r)
143}
144
145/// Infects each susceptible cell with probability `initial_infected_pct`, so a run that has burnt out can be
146/// restarted without losing its recovered cells.
147fn seed_outbreak(grid: &mut Grid2D<u8>, params: &[ParamValue], rng: &mut u64) {
148    let initial_pct = extract_f32(params, INITIAL_INFECTED_PCT, 0.01);
149    let threshold = (initial_pct * u32::MAX as f32) as u32;
150    for cell in grid.current_mut().iter_mut() {
151        if *cell == S && below(next_bits(rng), threshold) {
152            *cell = I;
153        }
154    }
155}
156
157/// Counts S, I and R over the grid in chunks, folded in index order.
158fn count_sir(cells: &[u8]) -> (u64, u64, u64) {
159    reduce_chunks(
160        cells.len(),
161        STATS_CHUNK,
162        |r| count_sir_seq(&cells[r]),
163        |a, b| (a.0 + b.0, a.1 + b.1, a.2 + b.2),
164        (0, 0, 0),
165    )
166}
167
168#[cfg(test)]
169mod tests {
170
171    use super::*;
172    use henad_compute::cpu::grid_engine::GridModelState;
173    use henad_core::model::SimState as _;
174
175    #[test]
176    fn sir_population_conservation() {
177        let params = vec![
178            ParamValue::U32(100),
179            ParamValue::U32(100),
180            ParamValue::F32(0.3),
181            ParamValue::F32(0.05),
182            ParamValue::F32(0.01),
183        ];
184        let mut state = GridModelState::<SirGridModel>::from_params(&params);
185        let pop = state.population();
186        assert_eq!(pop, 10_000, "population should be 100x100");
187
188        for _ in 0..100 {
189            state.step();
190        }
191        assert_eq!(state.population(), pop, "S+I+R must remain constant over 100 steps");
192    }
193
194    #[test]
195    fn sir_step_cell_transitions() {
196        let params = SirParams {
197            infection_rate: 1.0,
198            recovery_rate: 0.0,
199        };
200        let mut rng = 42u64;
201        // At a rate of 1, a susceptible cell next to an infected one is always infected.
202        assert_eq!(
203            SirGridModel::step_cell(S, &[I, S, S, S, S, S, S, S], &params, &mut rng),
204            I
205        );
206        // At a recovery rate of 0, an infected cell stays infected.
207        assert_eq!(
208            SirGridModel::step_cell(I, &[S, S, S, S, S, S, S, S], &params, &mut rng),
209            I
210        );
211        // A recovered cell stays recovered.
212        assert_eq!(
213            SirGridModel::step_cell(R, &[I, I, I, I, I, I, I, I], &params, &mut rng),
214            R
215        );
216    }
217}