1use 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 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}