1use 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
33pub const PALETTE: [[u8; 4]; 3] = [
35 [0x00, 0x7A, 0xF5, 0xFF], [0xE4, 0x37, 0x48, 0xFF], [0x80, 0x80, 0x80, 0xFF], ];
39
40#[derive(Debug)]
42pub struct SirGridModel;
43
44#[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 const STATS: &'static [StatDescriptor] = &[
59 StatDescriptor::new("Susceptible", PALETTE[0]),
60 StatDescriptor::new("Infected", PALETTE[1]),
61 StatDescriptor::new("Recovered", PALETTE[2]),
62 ];
63 const ACTIONS: &'static [ActionDescriptor] = ACTION_SPECS;
65 type Params = SirParams;
66
67 fn param_descriptors() -> Vec<ParamDescriptor> {
68 descriptors()
69 }
70
71 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 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 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 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 }
131
132fn 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
145fn 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
157fn 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(¶ms);
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 assert_eq!(
203 SirGridModel::step_cell(S, &[I, S, S, S, S, S, S, S], ¶ms, &mut rng),
204 I
205 );
206 assert_eq!(
208 SirGridModel::step_cell(I, &[S, S, S, S, S, S, S, S], ¶ms, &mut rng),
209 I
210 );
211 assert_eq!(
213 SirGridModel::step_cell(R, &[I, I, I, I, I, I, I, I], ¶ms, &mut rng),
214 R
215 );
216 }
217}