Skip to main content

gam_inference/
difference_smooth.rs

1//! Typed difference-smooth contrast and report orchestration.
2//!
3//! Front ends provide saved-model schema/term metadata and one design-builder
4//! callback. This module owns grid construction, group/pair resolution, row
5//! assembly, nuisance-column policy, contrast orientation, covariance-band
6//! configuration, and final report rows.
7
8use crate::effects::{
9    self, BandOptions, CovarianceSource, PointwiseBandOptions, SimultaneousBandOptions,
10};
11use gam_data::{ColumnKindTag, DataSchema};
12use gam_terms::smooth::TermCollectionSpec;
13use ndarray::{Array2, ArrayView1, ArrayView2, s};
14use serde::{Deserialize, Serialize};
15use std::collections::BTreeMap;
16
17#[derive(Clone, Debug, Deserialize, Serialize)]
18#[serde(deny_unknown_fields)]
19pub struct DifferenceSmoothRequest {
20    pub view: String,
21    pub group: Option<String>,
22    pub pairs: Option<Vec<(String, String)>>,
23    pub n: usize,
24    pub level: Option<f64>,
25    pub simultaneous: bool,
26    pub n_sim: Option<usize>,
27    pub seed: Option<u64>,
28    pub marginalise_random: bool,
29    pub group_means: bool,
30    pub template: Option<BTreeMap<String, String>>,
31}
32
33#[derive(Clone, Debug, Serialize)]
34pub struct DifferenceSmoothRow {
35    #[serde(flatten)]
36    pub view_value: BTreeMap<String, f64>,
37    pub group: String,
38    pub level_1: String,
39    pub level_2: String,
40    pub diff: f64,
41    pub se: f64,
42    pub lower: f64,
43    pub upper: f64,
44    pub level: f64,
45    pub simultaneous: bool,
46    pub critical: f64,
47    pub covariance_kind: String,
48    pub covariance_corrected: bool,
49}
50
51pub struct DifferenceSmoothInputs<'a> {
52    pub schema: &'a DataSchema,
53    pub training_feature_ranges: &'a [(f64, f64)],
54    pub termspec: &'a TermCollectionSpec,
55    pub beta: ArrayView1<'a, f64>,
56    pub covariance: ArrayView2<'a, f64>,
57    pub covariance_source: CovarianceSource,
58}
59
60pub fn difference_smooth_report(
61    inputs: DifferenceSmoothInputs<'_>,
62    request: DifferenceSmoothRequest,
63    mut build_design: impl FnMut(&[String], &[Vec<String>]) -> Result<Array2<f64>, String>,
64) -> Result<Vec<DifferenceSmoothRow>, String> {
65    let level = request.level.unwrap_or(effects::DEFAULT_BAND_LEVEL);
66    if !(0.0 < level && level < 1.0) {
67        return Err("difference_smooth level must be in (0, 1)".to_string());
68    }
69    if request.n < 2 {
70        return Err("difference_smooth n must be at least 2".to_string());
71    }
72    if inputs.schema.columns.len() != inputs.training_feature_ranges.len() {
73        return Err(format!(
74            "difference_smooth schema/range mismatch: {} columns but {} training ranges",
75            inputs.schema.columns.len(),
76            inputs.training_feature_ranges.len()
77        ));
78    }
79
80    let headers: Vec<String> = inputs
81        .schema
82        .columns
83        .iter()
84        .map(|column| column.name.clone())
85        .collect();
86    let view_idx = headers
87        .iter()
88        .position(|name| name == &request.view)
89        .ok_or_else(|| {
90            format!(
91                "view column {:?} not found in model schema: {:?}",
92                request.view, headers
93            )
94        })?;
95    let group = match request.group {
96        Some(group) => group,
97        None => inputs
98            .schema
99            .columns
100            .iter()
101            .find(|column| column.kind == ColumnKindTag::Categorical && column.name != request.view)
102            .map(|column| column.name.clone())
103            .ok_or_else(|| {
104                "difference_smooth could not infer a categorical group column; pass group="
105                    .to_string()
106            })?,
107    };
108    let group_column = inputs
109        .schema
110        .columns
111        .iter()
112        .find(|column| column.name == group)
113        .ok_or_else(|| format!("group column {group:?} not found in model schema: {headers:?}"))?;
114    if group_column.levels.len() < 2 {
115        return Err(format!(
116            "group column {group:?} must have at least two saved levels"
117        ));
118    }
119    let pairs = request.pairs.unwrap_or_else(|| {
120        let mut pairs = Vec::new();
121        for left in 0..group_column.levels.len() {
122            for right in (left + 1)..group_column.levels.len() {
123                pairs.push((
124                    group_column.levels[left].clone(),
125                    group_column.levels[right].clone(),
126                ));
127            }
128        }
129        pairs
130    });
131    if pairs.is_empty() {
132        return Err("difference_smooth requires at least one level pair".to_string());
133    }
134    for (level_1, level_2) in &pairs {
135        if level_1 == level_2 {
136            return Err(format!(
137                "difference_smooth pair levels must differ; got {level_1:?} twice"
138            ));
139        }
140        for level in [level_1, level_2] {
141            if !group_column.levels.contains(level) {
142                return Err(format!(
143                    "difference_smooth level {level:?} is not saved for group {group:?}"
144                ));
145            }
146        }
147    }
148
149    let (lo, hi) = inputs.training_feature_ranges[view_idx];
150    if !(lo.is_finite() && hi.is_finite() && lo < hi) {
151        return Err(format!(
152            "difference_smooth view range for {:?} must be finite and increasing; got ({lo:?}, {hi:?})",
153            request.view
154        ));
155    }
156    let step = (hi - lo) / (request.n - 1) as f64;
157    let grid: Vec<f64> = (0..request.n)
158        .map(|index| lo + step * index as f64)
159        .collect();
160    let template = complete_template(
161        request.template.unwrap_or_default(),
162        inputs.schema,
163        inputs.training_feature_ranges,
164    )?;
165    let (random_ranges, group_ranges) = random_effect_ranges(inputs.termspec, &group)?;
166    let band_options = if request.simultaneous {
167        BandOptions::Simultaneous(SimultaneousBandOptions {
168            level,
169            simulations: request.n_sim.unwrap_or(effects::DEFAULT_SIMULATIONS),
170            seed: request.seed.unwrap_or(effects::DEFAULT_SIMULATION_SEED),
171        })
172    } else {
173        BandOptions::Pointwise(PointwiseBandOptions { level })
174    };
175    let covariance_kind = inputs.covariance_source.to_string();
176    let covariance_corrected = inputs.covariance_source == CovarianceSource::SmoothingCorrected;
177    let mut output = Vec::with_capacity(pairs.len() * grid.len());
178
179    for (level_1, level_2) in pairs {
180        let rows_left = contrast_rows(&headers, &template, &request.view, &group, &level_1, &grid);
181        let rows_right = contrast_rows(&headers, &template, &request.view, &group, &level_2, &grid);
182        let left = build_design(&headers, &rows_left)?;
183        let right = build_design(&headers, &rows_right)?;
184        if left.raw_dim() != right.raw_dim() {
185            return Err(format!(
186                "difference_smooth candidate designs disagree in shape: {:?} vs {:?}",
187                left.raw_dim(),
188                right.raw_dim()
189            ));
190        }
191        // Pair orientation is level_1 - level_2, matching the row labels.
192        let mut contrast = &left - &right;
193        if request.marginalise_random {
194            let ranges = if request.group_means {
195                subtract_ranges(&random_ranges, &group_ranges)
196            } else {
197                random_ranges.clone()
198            };
199            zero_ranges(&mut contrast, &ranges)?;
200        }
201        if !request.group_means {
202            zero_ranges(&mut contrast, &group_ranges)?;
203        }
204        let report = effects::effect_report(
205            inputs.beta,
206            inputs.covariance,
207            contrast.view(),
208            band_options,
209        )
210        .map_err(|error| error.to_string())?;
211        for (index, &x) in grid.iter().enumerate() {
212            output.push(DifferenceSmoothRow {
213                view_value: BTreeMap::from([(request.view.clone(), x)]),
214                group: group.clone(),
215                level_1: level_1.clone(),
216                level_2: level_2.clone(),
217                diff: report.center[index],
218                se: report.se[index],
219                lower: report.lower[index],
220                upper: report.upper[index],
221                level,
222                simultaneous: request.simultaneous,
223                critical: report.critical,
224                covariance_kind: covariance_kind.clone(),
225                covariance_corrected,
226            });
227        }
228    }
229    Ok(output)
230}
231
232fn complete_template(
233    mut template: BTreeMap<String, String>,
234    schema: &DataSchema,
235    ranges: &[(f64, f64)],
236) -> Result<BTreeMap<String, String>, String> {
237    for (index, column) in schema.columns.iter().enumerate() {
238        if template.contains_key(&column.name) {
239            continue;
240        }
241        let value = match column.kind {
242            ColumnKindTag::Categorical => column.levels.first().cloned().ok_or_else(|| {
243                format!("categorical column {:?} has no saved levels", column.name)
244            })?,
245            ColumnKindTag::Binary => "0".to_string(),
246            ColumnKindTag::Continuous => {
247                let (lo, hi) = ranges[index];
248                if !(lo.is_finite() && hi.is_finite()) {
249                    return Err(format!(
250                        "training range for {:?} must be finite",
251                        column.name
252                    ));
253                }
254                (0.5 * (lo + hi)).to_string()
255            }
256        };
257        template.insert(column.name.clone(), value);
258    }
259    Ok(template)
260}
261
262fn contrast_rows(
263    headers: &[String],
264    template: &BTreeMap<String, String>,
265    view: &str,
266    group: &str,
267    level: &str,
268    grid: &[f64],
269) -> Vec<Vec<String>> {
270    grid.iter()
271        .map(|x| {
272            headers
273                .iter()
274                .map(|header| {
275                    if header == view {
276                        x.to_string()
277                    } else if header == group {
278                        level.to_string()
279                    } else {
280                        template
281                            .get(header)
282                            .expect("complete template must contain every header")
283                            .clone()
284                    }
285                })
286                .collect()
287        })
288        .collect()
289}
290
291fn random_effect_ranges(
292    termspec: &TermCollectionSpec,
293    group: &str,
294) -> Result<(Vec<(usize, usize)>, Vec<(usize, usize)>), String> {
295    let mut column = 1 + termspec.linear_terms.len();
296    let mut all = Vec::with_capacity(termspec.random_effect_terms.len());
297    let mut selected = Vec::new();
298    for term in &termspec.random_effect_terms {
299        let levels = term.frozen_levels.as_ref().ok_or_else(|| {
300            format!(
301                "difference_smooth random effect {:?} has no frozen levels",
302                term.name
303            )
304        })?;
305        let range = (column, column + levels.len());
306        all.push(range);
307        if term.name == group {
308            selected.push(range);
309        }
310        column += levels.len();
311    }
312    Ok((all, selected))
313}
314
315fn subtract_ranges(ranges: &[(usize, usize)], excluded: &[(usize, usize)]) -> Vec<(usize, usize)> {
316    let mut output = Vec::new();
317    for &(start, end) in ranges {
318        let mut segments = vec![(start, end)];
319        for &(excluded_start, excluded_end) in excluded {
320            let mut next = Vec::new();
321            for (segment_start, segment_end) in segments {
322                if excluded_end <= segment_start || excluded_start >= segment_end {
323                    next.push((segment_start, segment_end));
324                } else {
325                    if segment_start < excluded_start {
326                        next.push((segment_start, excluded_start));
327                    }
328                    if excluded_end < segment_end {
329                        next.push((excluded_end, segment_end));
330                    }
331                }
332            }
333            segments = next;
334        }
335        output.extend(segments.into_iter().filter(|(start, end)| start < end));
336    }
337    output
338}
339
340fn zero_ranges(design: &mut Array2<f64>, ranges: &[(usize, usize)]) -> Result<(), String> {
341    for &(start, end) in ranges {
342        if start > end || end > design.ncols() {
343            return Err(format!(
344                "difference_smooth design range {start}..{end} exceeds {} columns",
345                design.ncols()
346            ));
347        }
348        if start < end {
349            design.slice_mut(s![.., start..end]).fill(0.0);
350        }
351    }
352    Ok(())
353}
354
355#[cfg(test)]
356mod tests {
357    use super::*;
358    use gam_data::SchemaColumn;
359    use ndarray::{Array1, array};
360
361    #[test]
362    fn pair_orientation_and_report_are_owned_by_core() {
363        let schema = DataSchema {
364            columns: vec![
365                SchemaColumn {
366                    name: "x".to_string(),
367                    kind: ColumnKindTag::Continuous,
368                    levels: Vec::new(),
369                },
370                SchemaColumn {
371                    name: "g".to_string(),
372                    kind: ColumnKindTag::Categorical,
373                    levels: vec!["A".to_string(), "B".to_string()],
374                },
375            ],
376        };
377        let request = DifferenceSmoothRequest {
378            view: "x".to_string(),
379            group: Some("g".to_string()),
380            pairs: Some(vec![("B".to_string(), "A".to_string())]),
381            n: 2,
382            level: Some(0.95),
383            simultaneous: false,
384            n_sim: None,
385            seed: None,
386            marginalise_random: false,
387            group_means: true,
388            template: None,
389        };
390        let termspec = TermCollectionSpec {
391            linear_terms: Vec::new(),
392            smooth_terms: Vec::new(),
393            random_effect_terms: Vec::new(),
394        };
395        let beta = Array1::from_vec(vec![0.0, 1.5]);
396        let covariance = array![[0.1, 0.0], [0.0, 0.1]];
397        let rows = difference_smooth_report(
398            DifferenceSmoothInputs {
399                schema: &schema,
400                training_feature_ranges: &[(0.0, 1.0), (0.0, 1.0)],
401                termspec: &termspec,
402                beta: beta.view(),
403                covariance: covariance.view(),
404                covariance_source: CovarianceSource::Conditional,
405            },
406            request,
407            |_headers, rows| {
408                Ok(Array2::from_shape_fn((rows.len(), 2), |(row, column)| {
409                    if column == 0 {
410                        1.0
411                    } else if rows[row][1] == "B" {
412                        1.0
413                    } else {
414                        0.0
415                    }
416                }))
417            },
418        )
419        .expect("difference report");
420        assert_eq!(rows.len(), 2);
421        assert!(rows.iter().all(|row| (row.diff - 1.5).abs() < 1.0e-12));
422    }
423}