use ndarray::{Array1, Array2, ArrayView1};
use gam_gpu::gpu_error::GpuError;
#[derive(Debug)]
pub enum SigmaCubatureGpuError {
#[cfg(target_os = "linux")]
Geometry(gam_problem::EstimationError),
Runtime(GpuError),
}
impl std::fmt::Display for SigmaCubatureGpuError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
#[cfg(target_os = "linux")]
Self::Geometry(error) => write!(f, "{error}"),
Self::Runtime(error) => write!(f, "{error}"),
}
}
}
impl std::error::Error for SigmaCubatureGpuError {}
impl From<GpuError> for SigmaCubatureGpuError {
fn from(error: GpuError) -> Self {
Self::Runtime(error)
}
}
pub(crate) fn certified_sigma_point_covariance(
hessian: &Array2<f64>,
label: &str,
) -> Result<Array2<f64>, gam_linalg::utils::CertifiedSymmetricSolveError> {
gam_linalg::utils::certified_spd_inverse(hessian, label)
.map(gam_linalg::utils::CertifiedSpdInverse::into_inverse)
}
pub struct SigmaPointGpuInput {
pub s_transformed: Array2<f64>,
pub qs: Array2<f64>,
pub linear_shift: Array1<f64>,
pub constant_shift: f64,
}
#[cfg(target_os = "linux")]
const STREAM_POOL_MAX: usize = 8;
#[cfg(target_os = "linux")]
const SIGMA_PIRLS_INITIAL_LM_LAMBDA: f64 = 1e-6;
#[cfg(target_os = "linux")]
#[inline]
fn pool_size(m: usize) -> usize {
m.min(STREAM_POOL_MAX).max(1)
}
pub fn try_gpu_sigma_stream_pool_eval(
x_original: ndarray::ArrayView2<'_, f64>,
y: ArrayView1<'_, f64>,
prior_w: ArrayView1<'_, f64>,
offset: ArrayView1<'_, f64>,
per_sigma: &[SigmaPointGpuInput],
admission: gam_gpu::policy::PirlsLoopAdmission,
likelihood_scale: crate::gpu::pirls_gpu::PirlsLoopLikelihoodScale,
convergence_tol: f64,
max_iter: usize,
) -> Result<Option<Vec<(ndarray::Array2<f64>, ndarray::Array1<f64>)>>, SigmaCubatureGpuError> {
if per_sigma.is_empty() {
return Ok(Some(Vec::new()));
}
validate_sigma_point_inputs(x_original.ncols(), per_sigma)?;
#[cfg(target_os = "linux")]
{
if gam_gpu::device_runtime::GpuRuntime::resolve(gam_gpu::global_policy())?.is_none() {
return Ok(None);
}
let Some(family_kind) = admission.family else {
return Ok(None);
};
let Some(family) = linux_impl::family_kind_to_row(family_kind) else {
return Err(
gam_gpu::gpu_err!("sigma stream pool: family not in JIT-cached set").into(),
);
};
let curvature = linux_impl::curvature_kind_to_row(admission.curvature);
return linux_impl::stream_pool_eval(
x_original,
y,
prior_w,
offset,
per_sigma,
family,
curvature,
likelihood_scale,
convergence_tol,
max_iter,
);
}
#[cfg(not(target_os = "linux"))]
{
log::trace!(
"[sigma stream pool] non-Linux target: skipping dispatch \
(x_original={}x{}, y_len={}, prior_w_len={}, offset_len={}, \
n_sigma={}, family={:?}, curvature={:?}, likelihood_scale={:?}, \
tol={}, max_iter={})",
x_original.nrows(),
x_original.ncols(),
y.len(),
prior_w.len(),
offset.len(),
per_sigma.len(),
admission.family,
admission.curvature,
likelihood_scale,
convergence_tol,
max_iter,
);
Ok(None)
}
}
fn validate_sigma_point_inputs(p: usize, per_sigma: &[SigmaPointGpuInput]) -> Result<(), GpuError> {
for (idx, pt) in per_sigma.iter().enumerate() {
if pt.s_transformed.shape() != [p, p] {
return Err(gam_gpu::gpu_err!(
"sigma stream pool: point[{idx}] S shape {:?} != [{p}, {p}]",
pt.s_transformed.shape()
));
}
if pt.qs.shape() != [p, p] {
return Err(gam_gpu::gpu_err!(
"sigma stream pool: point[{idx}] Qs shape {:?} != [{p}, {p}]",
pt.qs.shape()
));
}
if pt.linear_shift.len() != p {
return Err(gam_gpu::gpu_err!(
"sigma stream pool: point[{idx}] linear shift len {} != {p}",
pt.linear_shift.len()
));
}
if !pt.constant_shift.is_finite() {
return Err(gam_gpu::gpu_err!(
"sigma stream pool: point[{idx}] non-finite constant shift {}",
pt.constant_shift
));
}
}
Ok(())
}
#[cfg(target_os = "linux")]
mod linux_impl {
use crate::gpu_kernels::pirls_row::{CurvatureMode, PirlsRowFamily};
use crate::gpu_kernels::sigma_cubature::{SigmaCubatureGpuError, SigmaPointGpuInput};
use gam_gpu::gpu_error::GpuError;
use gam_gpu::policy::{PirlsLoopCurvatureKind, PirlsLoopFamilyKind};
use ndarray::{Array1, Array2, ArrayView1};
type SigmaPointResult = (Array2<f64>, Array1<f64>);
pub(super) fn family_kind_to_row(f: PirlsLoopFamilyKind) -> Option<PirlsRowFamily> {
match f {
PirlsLoopFamilyKind::BernoulliLogit => Some(PirlsRowFamily::BernoulliLogit),
PirlsLoopFamilyKind::BernoulliProbit => Some(PirlsRowFamily::BernoulliProbit),
PirlsLoopFamilyKind::BernoulliCLogLog => Some(PirlsRowFamily::BernoulliCLogLog),
PirlsLoopFamilyKind::PoissonLog => Some(PirlsRowFamily::PoissonLog),
PirlsLoopFamilyKind::GaussianIdentity => Some(PirlsRowFamily::GaussianIdentity),
PirlsLoopFamilyKind::GammaLog => Some(PirlsRowFamily::GammaLog),
}
}
pub(super) fn curvature_kind_to_row(c: PirlsLoopCurvatureKind) -> CurvatureMode {
match c {
PirlsLoopCurvatureKind::Fisher => CurvatureMode::Fisher,
PirlsLoopCurvatureKind::Observed => CurvatureMode::Observed,
}
}
fn hessian_to_original(
h_transformed: &ndarray::Array2<f64>,
qs: &ndarray::Array2<f64>,
) -> ndarray::Array2<f64> {
let tmp = qs.dot(h_transformed);
let mut h_orig = tmp.dot(&qs.t());
gam_linalg::matrix::symmetrize_in_place(&mut h_orig);
h_orig
}
pub(super) fn stream_pool_eval(
x_original: ndarray::ArrayView2<'_, f64>,
y: ArrayView1<'_, f64>,
prior_w: ArrayView1<'_, f64>,
offset: ArrayView1<'_, f64>,
per_sigma: &[SigmaPointGpuInput],
family: PirlsRowFamily,
curvature: CurvatureMode,
likelihood_scale: crate::gpu::pirls_gpu::PirlsLoopLikelihoodScale,
convergence_tol: f64,
max_iter: usize,
) -> Result<Option<Vec<SigmaPointResult>>, SigmaCubatureGpuError> {
use crate::gpu::pirls_gpu;
use crate::gpu_kernels::sigma_cubature::pool_size;
let m = per_sigma.len();
let p = x_original.ncols();
for (idx, pt) in per_sigma.iter().enumerate() {
if pt.s_transformed.shape() != [p, p] || pt.qs.shape() != [p, p] {
return Err(gam_gpu::gpu_err!(
"sigma stream pool: point[{idx}] shape mismatch against point[0]"
)
.into());
}
}
if family == PirlsRowFamily::GaussianIdentity {
return gaussian_sigma_pool_eval(x_original, y, prior_w, offset, per_sigma, p)
.map_err(SigmaCubatureGpuError::Runtime);
}
let bootstrap_shared =
pirls_gpu::upload_shared_pirls_gpu(x_original, y, prior_w, offset)
.map_err(|e| gam_gpu::gpu_err!("sigma stream pool bootstrap upload: {e}"))?;
let n_streams = pool_size(m);
let mut workspace_pairs: Vec<(
crate::gpu::pirls_gpu::SigmaPirlsGpuWorkspace,
crate::gpu::pirls_gpu::cuda::PirlsLoopWorkspace,
)> = Vec::with_capacity(n_streams);
for _ in 0..n_streams {
let ws = pirls_gpu::allocate_sigma_pirls_workspace(&bootstrap_shared)
.map_err(|e| gam_gpu::gpu_err!("sigma stream pool alloc workspace: {e}"))?;
let loop_ws = pirls_gpu::allocate_pirls_loop_workspace(&bootstrap_shared, &ws)
.map_err(|e| gam_gpu::gpu_err!("sigma stream pool alloc loop_ws: {e}"))?;
workspace_pairs.push((ws, loop_ws));
}
let beta0: Array1<f64> = Array1::zeros(p);
let mut outcomes: Vec<SigmaPointResult> = Vec::with_capacity(m);
for (idx, pt) in per_sigma.iter().enumerate() {
let stream_idx = idx % n_streams;
let (ws, loop_ws) = &mut workspace_pairs[stream_idx];
pirls_gpu::upload_qs_pirls(ws, pt.qs.view())
.map_err(|e| gam_gpu::gpu_err!("sigma stream pool upload Qs pt[{idx}]: {e}"))?;
let shared = &bootstrap_shared;
let outcome = pirls_gpu::pirls_loop_on_stream(
shared,
ws,
loop_ws,
family,
curvature,
likelihood_scale,
beta0.view(),
pt.s_transformed.view(),
pt.linear_shift.view(),
pt.constant_shift,
super::SIGMA_PIRLS_INITIAL_LM_LAMBDA,
0.0,
max_iter,
convergence_tol,
None,
);
let loop_out = match outcome {
Ok(loop_out) => loop_out,
Err(pirls_gpu::cuda::PirlsGpuLoopError::Geometry(error)) => {
return Err(SigmaCubatureGpuError::Geometry(error));
}
Err(pirls_gpu::cuda::PirlsGpuLoopError::Runtime(message)) => {
return Err(SigmaCubatureGpuError::Runtime(gam_gpu::gpu_err!(
"sigma point[{idx}] GPU PIRLS runtime failure: {message}"
)));
}
};
let h_orig = hessian_to_original(&loop_out.penalized_hessian, &pt.qs);
let cov = super::certified_sigma_point_covariance(&h_orig, "gpu sigma point")
.map_err(|error| {
gam_gpu::gpu_err!(
"gpu sigma point: exact SPD penalised-Hessian inverse failed: {error}"
)
})?;
let beta_orig = pt.qs.dot(&loop_out.beta);
let sigma_result = (cov, beta_orig);
outcomes.push(sigma_result);
}
Ok(Some(outcomes))
}
fn gaussian_sigma_pool_eval(
x_original: ndarray::ArrayView2<'_, f64>,
y: ArrayView1<'_, f64>,
prior_w: ArrayView1<'_, f64>,
offset: ArrayView1<'_, f64>,
per_sigma: &[SigmaPointGpuInput],
p: usize,
) -> Result<Option<Vec<SigmaPointResult>>, GpuError> {
use ndarray::Array1;
let xtwx = crate::gpu::pirls_gpu::weighted_crossprod_gpu(x_original, prior_w)
.map_err(|e| gam_gpu::gpu_err!("gaussian sigma: XᵀWX gpu failed: {e}"))?;
let mut yw = y.to_owned();
yw -= &offset;
yw *= &prior_w;
let xtwy: Array1<f64> = x_original.t().dot(&yw);
let prior_mean_zero: Array1<f64> = Array1::zeros(p);
let mut outcomes: Vec<SigmaPointResult> = Vec::with_capacity(per_sigma.len());
for (idx, pt) in per_sigma.iter().enumerate() {
let pls = crate::gpu::pirls_gpu::solve_gaussian_pls_gpu(
xtwx.view(),
xtwy.view(),
pt.s_transformed.view(),
pt.linear_shift.view(),
prior_mean_zero.view(),
0.0,
Some(pt.qs.view()),
)
.map_err(|e| gam_gpu::gpu_err!("gaussian sigma pool: point[{idx}] pls failed: {e}"))?;
let h_orig = hessian_to_original(&pls.penalized_hessian, &pt.qs);
let cov = super::certified_sigma_point_covariance(&h_orig, "gaussian sigma point")
.map_err(|error| {
gam_gpu::gpu_err!(
"gaussian sigma point: exact SPD penalised-Hessian inverse failed: {error}"
)
})?;
let beta_orig = pt.qs.dot(&pls.beta);
outcomes.push((cov, beta_orig));
}
Ok(Some(outcomes))
}
}
#[cfg(test)]
mod covariance_contract_tests {
use super::certified_sigma_point_covariance;
use gam_linalg::utils::CertifiedSymmetricSolveError;
use ndarray::array;
#[test]
fn cpu_and_gpu_sigma_postprocessing_share_one_unperturbed_inverse_contract() {
let hessian = array![[4.0, 1.0], [1.0, 3.0]];
let covariance = certified_sigma_point_covariance(&hessian, "shared sigma contract")
.expect("strict SPD inverse");
let identity = hessian.dot(&covariance);
assert!((identity[[0, 0]] - 1.0).abs() <= 4.0 * f64::EPSILON);
assert!((identity[[1, 1]] - 1.0).abs() <= 4.0 * f64::EPSILON);
assert!(identity[[0, 1]].abs() <= 4.0 * f64::EPSILON);
assert!(identity[[1, 0]].abs() <= 4.0 * f64::EPSILON);
}
#[test]
fn shared_sigma_contract_rejects_an_invertible_indefinite_hessian() {
let hessian = array![[1.0, 2.0], [2.0, 1.0]];
assert!(matches!(
certified_sigma_point_covariance(&hessian, "shared sigma contract"),
Err(CertifiedSymmetricSolveError::NotPositiveDefinite { .. })
));
}
}