1use 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
46pub struct ChunkedKernelDesignOperator<K: SpatialKernelEvaluator> {
58 data: Arc<Array2<f64>>,
60 centers: Arc<Array2<f64>>,
62 kernel: K,
64 kernel_gauge: Option<Arc<Gauge>>,
66 poly_basis: Option<Arc<Array2<f64>>>,
68 n: usize,
69 total_cols: usize,
70 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 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, ¢ers[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 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 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 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 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 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 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 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 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 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}