Skip to main content

gam_terms/
chunked_kernel_design.rs

1//! Spatial-kernel design operators for basis construction.
2
3use faer::Accum;
4use faer::Par;
5use faer::linalg::matmul::matmul;
6use gam_linalg::faer_ndarray::{FaerArrayView, array2_to_matmut, fast_atv, fast_av};
7use gam_linalg::matrix::{DenseDesignOperator, FiniteSignedWeightsView, LinearOperator};
8use gam_problem::Gauge;
9use gam_runtime::resource::{MaterializationPolicy, MatrixMaterializationError};
10use ndarray::{Array1, Array2, ArrayViewMut2, s};
11use rayon::iter::{IndexedParallelIterator, ParallelIterator};
12use rayon::slice::ParallelSliceMut;
13use std::ops::Range;
14use std::sync::Arc;
15
16const KERNEL_OPERATOR_ROW_CHUNK_SIZE: usize = 2048;
17
18pub trait SpatialKernelEvaluator: Send + Sync + 'static {
19    fn eval(&self, x: &[f64], c: &[f64]) -> f64;
20}
21
22impl<F> SpatialKernelEvaluator for F
23where
24    F: Fn(&[f64], &[f64]) -> f64 + Send + Sync + 'static,
25{
26    fn eval(&self, x: &[f64], c: &[f64]) -> f64 {
27        self(x, c)
28    }
29}
30
31impl<F> SpatialKernelEvaluator for Arc<F>
32where
33    F: Fn(&[f64], &[f64]) -> f64 + Send + Sync + 'static + ?Sized,
34{
35    fn eval(&self, x: &[f64], c: &[f64]) -> f64 {
36        self.as_ref()(x, c)
37    }
38}
39
40impl SpatialKernelEvaluator for Arc<dyn SpatialKernelEvaluator> {
41    fn eval(&self, x: &[f64], c: &[f64]) -> f64 {
42        self.as_ref().eval(x, c)
43    }
44}
45
46/// Chunked kernel design operator for spatial smooths (TPS, Matérn, Duchon).
47///
48/// Instead of storing a dense n × k matrix, evaluates K(data[i], center[j])
49/// on-the-fly in row chunks. Memory usage is O(chunk_size × k) instead of O(n × k).
50///
51/// The optional `poly_basis` appends polynomial columns after the kernel columns
52/// (e.g., linear polynomial for TPS identifiability).
53///
54/// The optional `kernel_gauge` restricts the kernel coefficient block through
55/// a Gauge section, so the effective design is [K_reduced | poly] instead of
56/// [K | poly].
57pub struct ChunkedKernelDesignOperator<K: SpatialKernelEvaluator> {
58    /// Observation data points (n × d).
59    data: Arc<Array2<f64>>,
60    /// Radial basis centers (k × d).
61    centers: Arc<Array2<f64>>,
62    /// Kernel evaluator: (data_row, center_row) -> f64.
63    kernel: K,
64    /// Optional coefficient-space gauge applied to kernel columns.
65    kernel_gauge: Option<Arc<Gauge>>,
66    /// Optional polynomial basis columns (n × m) appended after kernel columns.
67    poly_basis: Option<Arc<Array2<f64>>>,
68    n: usize,
69    total_cols: usize,
70    /// The routing contract that selected this streamed representation. It is
71    /// propagated through coefficient/block wrappers so a later permissive
72    /// caller cannot silently materialize the full design.
73    materialization_policy: MaterializationPolicy,
74}
75
76impl<K: SpatialKernelEvaluator> ChunkedKernelDesignOperator<K> {
77    pub fn new(
78        data: Arc<Array2<f64>>,
79        centers: Arc<Array2<f64>>,
80        kernel: K,
81        kernel_gauge: Option<Arc<Gauge>>,
82        poly_basis: Option<Arc<Array2<f64>>>,
83        materialization_policy: MaterializationPolicy,
84    ) -> Result<Self, String> {
85        let n = data.nrows();
86        let k = centers.nrows();
87        if data.ncols() != centers.ncols() {
88            return Err(format!(
89                "ChunkedKernelDesignOperator: data dim {} != centers dim {}",
90                data.ncols(),
91                centers.ncols(),
92            ));
93        }
94        if let Some(gauge) = kernel_gauge.as_ref()
95            && gauge.raw_total() != k
96        {
97            return Err(format!(
98                "ChunkedKernelDesignOperator: kernel gauge raw width {} != centers rows {}",
99                gauge.raw_total(),
100                k,
101            ));
102        }
103        if let Some(poly) = poly_basis.as_ref()
104            && poly.nrows() != n
105        {
106            return Err(format!(
107                "ChunkedKernelDesignOperator: poly_basis rows {} != data rows {}",
108                poly.nrows(),
109                n,
110            ));
111        }
112        let k_eff = kernel_gauge.as_ref().map_or(k, |g| g.reduced_total());
113        let poly_cols = poly_basis.as_ref().map_or(0, |p| p.ncols());
114        Ok(Self {
115            data: Arc::new(data.as_standard_layout().to_owned()),
116            centers: Arc::new(centers.as_standard_layout().to_owned()),
117            kernel,
118            kernel_gauge,
119            poly_basis,
120            n,
121            total_cols: k_eff + poly_cols,
122            materialization_policy,
123        })
124    }
125
126    /// Evaluate kernel block for a range of rows, then restrict it through the
127    /// coefficient Gauge when present.
128    ///
129    /// This is not a matrix Kronecker product. The center rows are coordinate
130    /// arguments to `kernel.eval(data_row, center_row)`; each output entry is a
131    /// scalar kernel value before the optional column projection.
132    fn kernel_chunk(&self, rows: Range<usize>) -> Array2<f64> {
133        let chunk_n = rows.end - rows.start;
134        let k_raw = self.centers.nrows();
135        let dim = self.data.ncols();
136        let data = self
137            .data
138            .as_slice()
139            .expect("ChunkedKernelDesignOperator stores standard-layout data");
140        let centers = self
141            .centers
142            .as_slice()
143            .expect("ChunkedKernelDesignOperator stores standard-layout centers");
144        let kernel = &self.kernel;
145        let mut values = vec![0.0_f64; chunk_n * k_raw];
146        values
147            .par_chunks_mut(k_raw)
148            .enumerate()
149            .for_each(|(local, out_row)| {
150                let global = rows.start + local;
151                let x_start = global * dim;
152                let x = &data[x_start..x_start + dim];
153                for j in 0..k_raw {
154                    let c_start = j * dim;
155                    out_row[j] = kernel.eval(x, &centers[c_start..c_start + dim]);
156                }
157            });
158        let kernel_block = Array2::from_shape_vec((chunk_n, k_raw), values)
159            .expect("kernel chunk shape should match generated values");
160        if let Some(gauge) = self.kernel_gauge.as_ref() {
161            gauge.restrict_design(&kernel_block)
162        } else {
163            kernel_block
164        }
165    }
166}
167
168impl<K: SpatialKernelEvaluator> LinearOperator for ChunkedKernelDesignOperator<K> {
169    fn nrows(&self) -> usize {
170        self.n
171    }
172    fn ncols(&self) -> usize {
173        self.total_cols
174    }
175    fn apply(&self, vector: &Array1<f64>) -> Array1<f64> {
176        let k_eff = self
177            .kernel_gauge
178            .as_ref()
179            .map_or(self.centers.nrows(), |g| g.reduced_total());
180        let v_kernel = vector.slice(s![..k_eff]);
181        let mut result = Array1::<f64>::zeros(self.n);
182        // Process in chunks to limit memory.
183        for start in (0..self.n).step_by(KERNEL_OPERATOR_ROW_CHUNK_SIZE) {
184            let end = (start + KERNEL_OPERATOR_ROW_CHUNK_SIZE).min(self.n);
185            let chunk = self.kernel_chunk(start..end);
186            let partial = fast_av(&chunk, &v_kernel);
187            result.slice_mut(s![start..end]).assign(&partial);
188        }
189        if let Some(poly) = self.poly_basis.as_ref() {
190            let v_poly = vector.slice(s![k_eff..]);
191            let poly_part = fast_av(poly, &v_poly);
192            result += &poly_part;
193        }
194        result
195    }
196    fn apply_transpose(&self, vector: &Array1<f64>) -> Array1<f64> {
197        let k_eff = self
198            .kernel_gauge
199            .as_ref()
200            .map_or(self.centers.nrows(), |g| g.reduced_total());
201        let mut result = Array1::<f64>::zeros(self.total_cols);
202        // Kernel part: chunked accumulation of K^T v.
203        for start in (0..self.n).step_by(KERNEL_OPERATOR_ROW_CHUNK_SIZE) {
204            let end = (start + KERNEL_OPERATOR_ROW_CHUNK_SIZE).min(self.n);
205            let chunk = self.kernel_chunk(start..end);
206            let v_slice = vector.slice(s![start..end]);
207            let partial = fast_atv(&chunk, &v_slice);
208            result.slice_mut(s![..k_eff]).scaled_add(1.0, &partial);
209        }
210        // Poly part.
211        if let Some(poly) = self.poly_basis.as_ref() {
212            let poly_part = fast_atv(poly, vector);
213            result.slice_mut(s![k_eff..]).assign(&poly_part);
214        }
215        result
216    }
217    fn diag_xtw_x(&self, weights: &Array1<f64>) -> Result<Array2<f64>, String> {
218        if weights.len() != self.n {
219            return Err(format!(
220                "ChunkedSpatialKernelDesign::diag_xtw_x weight length mismatch: weights={}, nrows={}",
221                weights.len(),
222                self.n
223            ));
224        }
225        FiniteSignedWeightsView::try_from_array(weights)
226            .map_err(|reason| format!("ChunkedSpatialKernelDesign::diag_xtw_x: {reason}"))?;
227        let p = self.total_cols;
228        // The basis router chose this operator because the dense design was not
229        // admitted. Preserve that decision and accumulate XᵀWX from bounded
230        // row chunks instead of hiding a second, operator-local dense cache.
231        let n = self.n;
232        if n == 0 || p == 0 {
233            return Ok(Array2::<f64>::zeros((p, p)));
234        }
235        let chunk_starts: Vec<usize> = (0..n).step_by(KERNEL_OPERATOR_ROW_CHUNK_SIZE).collect();
236        // Deterministic parallel reduction over the row chunks: length-only
237        // pairwise tree, so the accumulated XᵀWX never depends on thread count
238        // or rayon's demand-driven fold/reduce grouping (#2228).
239        let xtwx = gam_linalg::pairwise_reduce::par_deterministic_block_fold(
240            chunk_starts.len(),
241            |idx_range: core::ops::Range<usize>| {
242                let mut acc = Array2::<f64>::zeros((p, p));
243                for &start in &chunk_starts[idx_range] {
244                    let end = (start + KERNEL_OPERATOR_ROW_CHUNK_SIZE).min(n);
245                    let chunk = self.row_chunk_combined(start..end);
246                    let mut wchunk = chunk.clone();
247                    for local in 0..(end - start) {
248                        let wi = weights[start + local];
249                        wchunk.row_mut(local).mapv_inplace(|v| v * wi);
250                    }
251                    let chunk_view = FaerArrayView::new(&chunk);
252                    let wchunk_view = FaerArrayView::new(&wchunk);
253                    let mut acc_view = array2_to_matmut(&mut acc);
254                    matmul(
255                        acc_view.as_mut(),
256                        Accum::Add,
257                        chunk_view.as_ref().transpose(),
258                        wchunk_view.as_ref(),
259                        1.0,
260                        Par::Seq,
261                    );
262                }
263                acc
264            },
265            |mut a, b| {
266                a += &b;
267                a
268            },
269        )
270        .unwrap_or_else(|| Array2::<f64>::zeros((p, p)));
271        Ok(xtwx)
272    }
273}
274
275impl<K: SpatialKernelEvaluator> ChunkedKernelDesignOperator<K> {
276    /// Build a combined row chunk `[kernel_chunk | poly_chunk]` without ever
277    /// materializing the full design.
278    pub(crate) fn row_chunk_combined(&self, rows: Range<usize>) -> Array2<f64> {
279        let chunk_n = rows.end - rows.start;
280        let k_eff = self
281            .kernel_gauge
282            .as_ref()
283            .map_or(self.centers.nrows(), |g| g.reduced_total());
284        let kernel = self.kernel_chunk(rows.clone());
285        let poly_cols = self.poly_basis.as_ref().map_or(0, |p| p.ncols());
286        let mut combined = Array2::<f64>::zeros((chunk_n, k_eff + poly_cols));
287        combined.slice_mut(s![.., ..k_eff]).assign(&kernel);
288        if let Some(poly) = self.poly_basis.as_ref() {
289            combined
290                .slice_mut(s![.., k_eff..])
291                .assign(&poly.slice(s![rows, ..]));
292        }
293        combined
294    }
295}
296
297impl<K: SpatialKernelEvaluator> DenseDesignOperator for ChunkedKernelDesignOperator<K> {
298    /// A chunked kernel design never exposes an implicit full-design cache.
299    /// Callers that explicitly accept densification must go through the
300    /// governed `DenseDesignMatrix` materialization APIs.
301    fn as_dense_ref(&self) -> Option<&Array2<f64>> {
302        None
303    }
304
305    fn materialization_policy(&self) -> Option<MaterializationPolicy> {
306        Some(self.materialization_policy.clone())
307    }
308
309    fn row_chunk_into(
310        &self,
311        rows: Range<usize>,
312        mut out: ArrayViewMut2<'_, f64>,
313    ) -> Result<(), MatrixMaterializationError> {
314        if out.nrows() != rows.end - rows.start || out.ncols() != self.total_cols {
315            return Err(MatrixMaterializationError::MissingRowChunk {
316                context: "ChunkedKernelDesignOperator::row_chunk_into shape mismatch",
317            });
318        }
319        out.assign(&self.row_chunk_combined(rows));
320        Ok(())
321    }
322
323    fn to_dense(&self) -> Array2<f64> {
324        self.row_chunk_combined(0..self.n)
325    }
326}
327
328#[cfg(test)]
329mod chunked_kernel_operator_tests {
330    use super::*;
331    use gam_linalg::matrix::DenseDesignMatrix;
332    use ndarray::{Array1, Array2, array};
333    use std::sync::Arc;
334
335    fn strict_materialization_policy() -> MaterializationPolicy {
336        gam_runtime::resource::ResourcePolicy::analytic_operator_required().material_policy()
337    }
338
339    #[test]
340    fn chunked_kernel_operator_uses_center_rows_for_column_count() {
341        let data = Arc::new(array![[0.0, 1.0], [1.0, 0.5]]);
342        let centers = Arc::new(array![[0.0, 0.0], [1.0, 1.0], [2.0, -1.0]]);
343        let kernel =
344            |x: &[f64], c: &[f64]| x.iter().zip(c.iter()).map(|(xi, ci)| xi * ci).sum::<f64>();
345        let operator = ChunkedKernelDesignOperator::new(
346            data,
347            centers,
348            kernel,
349            None,
350            None,
351            strict_materialization_policy(),
352        )
353        .expect("chunked kernel operator");
354
355        assert_eq!(operator.ncols(), 3);
356        let chunk = operator.row_chunk_combined(0..2);
357        assert_eq!(chunk.dim(), (2, 3));
358    }
359
360    #[test]
361    fn chunked_kernel_operator_rejects_incompatible_optional_shapes() {
362        let data = Arc::new(array![[0.0, 1.0], [1.0, 0.5]]);
363        let centers = Arc::new(array![[0.0, 0.0], [1.0, 1.0], [2.0, -1.0]]);
364        let kernel = |_: &[f64], _: &[f64]| 0.0;
365        let bad_gauge = Arc::new(gam_problem::Gauge::from_block_transforms(&[
366            Array2::<f64>::zeros((2, 1)),
367        ]));
368        let bad_poly = Arc::new(Array2::<f64>::zeros((3, 1)));
369
370        let gauge_err = match ChunkedKernelDesignOperator::new(
371            data.clone(),
372            centers.clone(),
373            kernel,
374            Some(bad_gauge),
375            None,
376            strict_materialization_policy(),
377        ) {
378            // SAFETY: test asserting validation rejects mismatched gauge raw width; Ok means the validator regressed.
379            Ok(_) => panic!("gauge raw width should match centers rows"),
380            Err(err) => err,
381        };
382        assert!(gauge_err.contains("kernel gauge raw width 2 != centers rows 3"));
383
384        let poly_err = match ChunkedKernelDesignOperator::new(
385            data,
386            centers,
387            kernel,
388            None,
389            Some(bad_poly),
390            strict_materialization_policy(),
391        ) {
392            // SAFETY: test asserting validation rejects mismatched poly rows; Ok means the validator regressed.
393            Ok(_) => panic!("poly rows should match data rows"),
394            Err(err) => err,
395        };
396        assert!(poly_err.contains("poly_basis rows 3 != data rows 2"));
397    }
398
399    #[test]
400    fn chunked_kernel_operator_canonicalizes_non_contiguous_inputs() {
401        let data = Arc::new(array![[0.0, 1.0], [1.0, 0.5]].reversed_axes());
402        let centers = Arc::new(array![[0.0, 1.0, 2.0], [0.0, 1.0, -1.0]].reversed_axes());
403        assert!(!data.is_standard_layout());
404        assert!(!centers.is_standard_layout());
405
406        let kernel =
407            |x: &[f64], c: &[f64]| x.iter().zip(c.iter()).map(|(xi, ci)| xi * ci).sum::<f64>();
408        let operator = ChunkedKernelDesignOperator::new(
409            data,
410            centers,
411            kernel,
412            None,
413            None,
414            strict_materialization_policy(),
415        )
416        .expect("chunked kernel operator");
417        let chunk = operator.row_chunk_combined(0..2);
418
419        assert_eq!(chunk.dim(), (2, 3));
420        assert_eq!(chunk[[0, 0]], 0.0);
421        assert_eq!(chunk[[1, 1]], 1.5);
422    }
423    #[test]
424    fn chunked_kernel_operator_never_exposes_an_implicit_dense_cache() {
425        let data = Arc::new(array![[0.0, 1.0], [1.0, 0.5], [2.0, -1.0]]);
426        let centers = Arc::new(array![[0.0, 0.0], [1.0, 1.0]]);
427        let kernel =
428            |x: &[f64], c: &[f64]| x.iter().zip(c.iter()).map(|(xi, ci)| xi * ci).sum::<f64>();
429        let op = ChunkedKernelDesignOperator::new(
430            data,
431            centers,
432            kernel,
433            None,
434            None,
435            strict_materialization_policy(),
436        )
437        .expect("chunked kernel operator");
438        let expected = op.to_dense();
439
440        let dense_design = DenseDesignMatrix::from(Arc::new(op));
441
442        let probe = Array1::from_elem(3, 1.0);
443        let applied = dense_design.apply_transpose(&probe);
444        let expected_applied = expected.t().dot(&probe);
445        for (got, want) in applied.iter().zip(expected_applied.iter()) {
446            assert!((got - want).abs() < 1e-12);
447        }
448        assert!(
449            dense_design.as_dense_ref().is_none(),
450            "chunked kernel operations must not warm a hidden full-design cache"
451        );
452    }
453
454    #[test]
455    fn chunked_kernel_gram_is_signed_and_rejects_nonfinite_rows() {
456        let data = Arc::new(array![[1.0], [2.0], [3.0]]);
457        let centers = Arc::new(array![[0.5], [-1.0]]);
458        let kernel = |x: &[f64], c: &[f64]| x[0] + c[0];
459        let op = ChunkedKernelDesignOperator::new(
460            data,
461            centers,
462            kernel,
463            None,
464            None,
465            strict_materialization_policy(),
466        )
467        .unwrap();
468        let dense = op.to_dense();
469        let weights = array![2.0, -3.0, 0.25];
470        let weighted_dense = dense.clone() * weights.view().insert_axis(ndarray::Axis(1));
471        let expected = dense.t().dot(&weighted_dense);
472        let got = op.diag_xtw_x(&weights).unwrap();
473        assert!((&got - &expected).iter().all(|value| value.abs() < 1e-12));
474
475        let err = op
476            .diag_xtw_x(&array![1.0, f64::NAN, f64::INFINITY])
477            .unwrap_err();
478        assert!(err.contains("row 1"), "unexpected diagnostic: {err}");
479    }
480}