henad-models 0.3.0

Example models for Henad, a parallel agent-based modelling engine.
Documentation
//! The SIR (susceptible, infected, recovered) epidemic as a [`GridModel`] on a torus.
//!
//! Each infected cell among a susceptible cell's eight neighbours infects it with probability `infection_rate`,
//! independently of the others. An infected cell recovers with probability `recovery_rate` each tick.

use henad_compute::cpu::primitives::chunked::{STATS_CHUNK, reduce_chunks};
use henad_core::action::ActionDescriptor;
use henad_core::authoring::model::grid_model::GridModel;
use henad_core::authoring::primitives::rng::{below, next_bits, next_float};
use henad_core::grid::Grid2D;
use henad_core::helpers::{extract_f32, f32_param};
use henad_core::params::{ParamDescriptor, ParamValue};
use henad_core::topology::NeighborhoodKind;
use henad_core::view::{StatDescriptor, StatValue};

const S: u8 = 0;
const I: u8 = 1;
const R: u8 = 2;

henad_core::params! {
    const INFECTION_RATE = f32_param("infection_rate", "Infection Rate", 0.3, 0.0, 1.0, Some(0.01));
    const RECOVERY_RATE = f32_param("recovery_rate", "Recovery Rate", 0.05, 0.0, 1.0, Some(0.01));
    const INITIAL_INFECTED_PCT =
        f32_param("initial_infected_pct", "Initial Infected", 0.01, 0.0, 1.0, Some(0.001))
            .on_reload()
            .percent();
}

henad_core::actions! {
    const SEED_OUTBREAK = ActionDescriptor::new("seed_outbreak", "Seed outbreak");
}

/// Cell colours, indexed by state: susceptible, infected, then recovered.
pub const PALETTE: [[u8; 4]; 3] = [
    [0x00, 0x7A, 0xF5, 0xFF], // S - blue
    [0xE4, 0x37, 0x48, 0xFF], // I - red
    [0x80, 0x80, 0x80, 0xFF], // R - gray
];

/// The SIR epidemic as a [`GridModel`].
#[derive(Debug)]
pub struct SirGridModel;

/// Rates of [`SirGridModel`], read once per tick.
#[derive(Debug)]
pub struct SirParams {
    infection_rate: f32,
    recovery_rate: f32,
}

impl GridModel for SirGridModel {
    const NAME: &'static str = "SIR Epidemic";
    const ID: &'static str = "sir";
    const DESCRIPTION: &'static str = "Classic SIR compartmental model on a 2D grid with Moore neighborhood";
    const PALETTE: &'static [[u8; 4]] = &PALETTE;
    const NEIGHBORHOOD: NeighborhoodKind = NeighborhoodKind::Moore;
    // --8<-- [start:stat_descriptors]
    const STATS: &'static [StatDescriptor] = &[
        StatDescriptor::new("Susceptible", PALETTE[0]),
        StatDescriptor::new("Infected", PALETTE[1]),
        StatDescriptor::new("Recovered", PALETTE[2]),
    ];
    // --8<-- [end:stat_descriptors]
    const ACTIONS: &'static [ActionDescriptor] = ACTION_SPECS;
    type Params = SirParams;

    fn param_descriptors() -> Vec<ParamDescriptor> {
        descriptors()
    }

    // --8<-- [start:from_params]
    fn from_params(params: &[ParamValue]) -> SirParams {
        SirParams {
            infection_rate: extract_f32(params, INFECTION_RATE, 0.3),
            recovery_rate: extract_f32(params, RECOVERY_RATE, 0.05),
        }
    }
    // --8<-- [end:from_params]

    fn init(grid: &mut Grid2D<u8>, params: &[ParamValue], rng: &mut u64) {
        let initial_pct = extract_f32(params, INITIAL_INFECTED_PCT, 0.01);
        let threshold = (initial_pct * u32::MAX as f32) as u32;
        for cell in grid.current_mut().iter_mut() {
            *cell = if below(next_bits(rng), threshold) { I } else { S };
        }
    }

    fn step_cell(cell: u8, neighbors: &[u8], params: &SirParams, rng: &mut u64) -> u8 {
        match cell {
            S => {
                let infected_count = neighbors.iter().filter(|&&n| n == I).count();
                if infected_count > 0 {
                    let prob_safe = (1.0 - params.infection_rate).powi(infected_count as i32);
                    // A rate of 1 always infects under `>=`. The draw is half-open and can be zero,
                    // and `>` would let a zero draw escape.
                    if next_float(rng, 1.0) >= prob_safe { I } else { S }
                } else {
                    S
                }
            }
            I => {
                if next_float(rng, 1.0) < params.recovery_rate {
                    R
                } else {
                    I
                }
            }
            _ => cell,
        }
    }

    #[expect(clippy::single_match, reason = "for future multi-action extendability")]
    fn act(action: usize, grid: &mut Grid2D<u8>, params: &[ParamValue], rng: &mut u64) {
        match action {
            SEED_OUTBREAK => seed_outbreak(grid, params, rng),
            _ => {}
        }
    }

    // --8<-- [start:stats]
    fn stats(grid: &Grid2D<u8>) -> Vec<StatValue> {
        let (s, i, r) = count_sir(grid.current());
        vec![
            StatValue::Scalar(s as f64),
            StatValue::Scalar(i as f64),
            StatValue::Scalar(r as f64),
        ]
    }
    // --8<-- [end:stats]
}

/// Counts S, I and R in one pass over `cells`.
fn count_sir_seq(cells: &[u8]) -> (u64, u64, u64) {
    let (mut s, mut i, mut r) = (0u64, 0u64, 0u64);
    for &cell in cells {
        match cell {
            S => s += 1,
            I => i += 1,
            _ => r += 1,
        }
    }
    (s, i, r)
}

/// Infects each susceptible cell with probability `initial_infected_pct`, so a run that has burnt out can be
/// restarted without losing its recovered cells.
fn seed_outbreak(grid: &mut Grid2D<u8>, params: &[ParamValue], rng: &mut u64) {
    let initial_pct = extract_f32(params, INITIAL_INFECTED_PCT, 0.01);
    let threshold = (initial_pct * u32::MAX as f32) as u32;
    for cell in grid.current_mut().iter_mut() {
        if *cell == S && below(next_bits(rng), threshold) {
            *cell = I;
        }
    }
}

/// Counts S, I and R over the grid in chunks, folded in index order.
fn count_sir(cells: &[u8]) -> (u64, u64, u64) {
    reduce_chunks(
        cells.len(),
        STATS_CHUNK,
        |r| count_sir_seq(&cells[r]),
        |a, b| (a.0 + b.0, a.1 + b.1, a.2 + b.2),
        (0, 0, 0),
    )
}

#[cfg(test)]
mod tests {

    use super::*;
    use henad_compute::cpu::grid_engine::GridModelState;
    use henad_core::model::SimState as _;

    #[test]
    fn sir_population_conservation() {
        let params = vec![
            ParamValue::U32(100),
            ParamValue::U32(100),
            ParamValue::F32(0.3),
            ParamValue::F32(0.05),
            ParamValue::F32(0.01),
        ];
        let mut state = GridModelState::<SirGridModel>::from_params(&params);
        let pop = state.population();
        assert_eq!(pop, 10_000, "population should be 100x100");

        for _ in 0..100 {
            state.step();
        }
        assert_eq!(state.population(), pop, "S+I+R must remain constant over 100 steps");
    }

    #[test]
    fn sir_step_cell_transitions() {
        let params = SirParams {
            infection_rate: 1.0,
            recovery_rate: 0.0,
        };
        let mut rng = 42u64;
        // At a rate of 1, a susceptible cell next to an infected one is always infected.
        assert_eq!(
            SirGridModel::step_cell(S, &[I, S, S, S, S, S, S, S], &params, &mut rng),
            I
        );
        // At a recovery rate of 0, an infected cell stays infected.
        assert_eq!(
            SirGridModel::step_cell(I, &[S, S, S, S, S, S, S, S], &params, &mut rng),
            I
        );
        // A recovered cell stays recovered.
        assert_eq!(
            SirGridModel::step_cell(R, &[I, I, I, I, I, I, I, I], &params, &mut rng),
            R
        );
    }
}