#[cfg(target_os = "linux")]
use std::sync::OnceLock;
use gam_gpu::gpu_error::GpuError;
#[cfg(target_os = "linux")]
use std::sync::Arc;
#[cfg(target_os = "linux")]
use cudarc::driver::{CudaModule, CudaSlice, CudaStream, LaunchConfig, PushKernelArg};
#[cfg(target_os = "linux")]
use super::super::flex_row_program::{
BmsFlexCalibrationOrder2Phase, BmsFlexRowOrder2FinalizerPhase, BmsFlexRowProgram,
};
#[cfg(target_os = "linux")]
pub(crate) const ROW_KERNEL_THREADS: u32 = 32;
pub(crate) const COEFF4: usize = 4;
pub(crate) const MOMENT_STRIDE: usize = 10;
pub(crate) enum CellMomentsSource<'a> {
Host(&'a [f64]),
#[cfg(target_os = "linux")]
Device(&'a CudaSlice<f64>),
}
impl<'a> CellMomentsSource<'a> {
pub(crate) fn len(&self) -> usize {
match self {
CellMomentsSource::Host(slice) => slice.len(),
#[cfg(target_os = "linux")]
CellMomentsSource::Device(d) => d.len(),
}
}
}
macro_rules! define_bms_flex_row_kernel_input_types {
(
f64_fields: [$($f64_field:ident),+ $(,)?],
u32_fields: [$($u32_field:ident),+ $(,)?],
moments_field: $moments_field:ident $(,)?
) => {
pub(crate) struct BmsFlexRowKernelInputs<'a> {
pub n_rows: usize,
pub r: usize,
pub p_h: usize,
pub p_w: usize,
pub s_f: f64,
$(pub $f64_field: &'a [f64],)+
$(pub $u32_field: &'a [u32],)+
pub $moments_field: CellMomentsSource<'a>,
}
pub(crate) struct BmsFlexRowKernelInputsOwned {
pub n_rows: usize,
pub r: usize,
pub p_h: usize,
pub p_w: usize,
pub s_f: f64,
$(pub $f64_field: Vec<f64>,)+
$(pub $u32_field: Vec<u32>,)+
pub $moments_field: Vec<f64>,
#[cfg(target_os = "linux")]
pub cell_moments_device: Option<CudaSlice<f64>>,
}
impl BmsFlexRowKernelInputsOwned {
pub(crate) fn as_borrowed(&self) -> BmsFlexRowKernelInputs<'_> {
#[cfg(target_os = "linux")]
let cell_moments = match self.cell_moments_device.as_ref() {
Some(d) => CellMomentsSource::Device(d),
None => CellMomentsSource::Host(&self.cell_moments),
};
#[cfg(not(target_os = "linux"))]
let cell_moments = CellMomentsSource::Host(&self.cell_moments);
BmsFlexRowKernelInputs {
n_rows: self.n_rows,
r: self.r,
p_h: self.p_h,
p_w: self.p_w,
s_f: self.s_f,
$($f64_field: &self.$f64_field,)+
$($u32_field: &self.$u32_field,)+
$moments_field: cell_moments,
}
}
}
};
}
define_bms_flex_row_kernel_input_types! {
f64_fields: [
q,
b,
mu_1,
mu_2,
z_obs,
y,
w,
e_obs,
cell_c0,
cell_c1,
cell_c2,
cell_c3,
cell_a,
cell_aa,
cell_r,
cell_ar,
cell_sbb,
cell_sbh,
cell_sbw,
chi_obs,
xi_obs,
rho_u,
tau_u,
r_uv,
],
u32_fields: [cell_offsets],
moments_field: cell_moments,
}
#[derive(Debug)]
pub(crate) struct BmsFlexRowKernelOutputs {
pub neglog: Vec<f64>,
pub grad: Vec<f64>,
pub hess: Vec<f64>,
}
fn checked_shape_len(context: &str, dimensions: &[usize]) -> Result<usize, GpuError> {
dimensions
.iter()
.copied()
.try_fold(1_usize, |product, dimension| {
product
.checked_mul(dimension)
.ok_or_else(|| GpuError::DriverCallFailed {
reason: format!(
"bms_flex_row {context}: shape product overflow for dimensions {dimensions:?}"
),
})
})
}
impl<'a> BmsFlexRowKernelInputs<'a> {
pub(crate) fn validate(&self) -> Result<(), GpuError> {
if self.n_rows == 0 {
return Err(GpuError::DriverCallFailed {
reason: "bms_flex_row inputs: n_rows must be > 0".to_string(),
});
}
if self.r == 0 {
return Err(GpuError::DriverCallFailed {
reason: "bms_flex_row inputs: r must be > 0".to_string(),
});
}
let decomposed_r = 2_usize
.checked_add(self.p_h)
.and_then(|value| value.checked_add(self.p_w))
.ok_or_else(|| GpuError::DriverCallFailed {
reason: format!(
"bms_flex_row inputs: primary decomposition overflow for p_h={} p_w={}",
self.p_h, self.p_w
),
})?;
if self.r != decomposed_r {
return Err(GpuError::DriverCallFailed {
reason: format!(
"bms_flex_row inputs: r={} must equal 2 + p_h({}) + p_w({}) = {}",
self.r, self.p_h, self.p_w, decomposed_r
),
});
}
let n = self.n_rows;
let check_len = |name: &str, have: usize, want: usize| -> Result<(), GpuError> {
if have != want {
return Err(GpuError::DriverCallFailed {
reason: format!("bms_flex_row inputs: {name}.len()={have} != {want}"),
});
}
Ok(())
};
check_len("q", self.q.len(), n)?;
check_len("b", self.b.len(), n)?;
check_len("mu_1", self.mu_1.len(), n)?;
check_len("mu_2", self.mu_2.len(), n)?;
check_len("z_obs", self.z_obs.len(), n)?;
check_len("y", self.y.len(), n)?;
check_len("w", self.w.len(), n)?;
check_len("e_obs", self.e_obs.len(), n)?;
check_len("chi_obs", self.chi_obs.len(), n)?;
check_len("xi_obs", self.xi_obs.len(), n)?;
let nr = checked_shape_len("validate [n,r]", &[n, self.r])?;
let nrr = checked_shape_len("validate [n,r,r]", &[n, self.r, self.r])?;
check_len("rho_u", self.rho_u.len(), nr)?;
check_len("tau_u", self.tau_u.len(), nr)?;
check_len("r_uv", self.r_uv.len(), nrr)?;
let offsets_len = n.checked_add(1).ok_or_else(|| GpuError::DriverCallFailed {
reason: format!("bms_flex_row inputs: n_rows={n} cannot form n+1 offsets"),
})?;
check_len("cell_offsets", self.cell_offsets.len(), offsets_len)?;
let total_cells_u32 = self.cell_offsets[n];
let total_cells = total_cells_u32 as usize;
check_len("cell_c0", self.cell_c0.len(), total_cells)?;
check_len("cell_c1", self.cell_c1.len(), total_cells)?;
check_len("cell_c2", self.cell_c2.len(), total_cells)?;
check_len("cell_c3", self.cell_c3.len(), total_cells)?;
let cells_coeff4 = checked_shape_len("validate cell coeff4", &[total_cells, COEFF4])?;
check_len("cell_a", self.cell_a.len(), cells_coeff4)?;
check_len("cell_aa", self.cell_aa.len(), cells_coeff4)?;
check_len(
"cell_r",
self.cell_r.len(),
checked_shape_len(
"validate cell_r",
&[total_cells, self.r.saturating_sub(1), COEFF4],
)?,
)?;
check_len(
"cell_ar",
self.cell_ar.len(),
checked_shape_len(
"validate cell_ar",
&[total_cells, self.r.saturating_sub(1), COEFF4],
)?,
)?;
check_len("cell_sbb", self.cell_sbb.len(), cells_coeff4)?;
check_len(
"cell_sbh",
self.cell_sbh.len(),
checked_shape_len("validate cell_sbh", &[total_cells, self.p_h, COEFF4])?,
)?;
check_len(
"cell_sbw",
self.cell_sbw.len(),
checked_shape_len("validate cell_sbw", &[total_cells, self.p_w, COEFF4])?,
)?;
check_len(
"cell_moments",
self.cell_moments.len(),
checked_shape_len("validate cell_moments", &[total_cells, MOMENT_STRIDE])?,
)?;
for i in 0..n {
if self.cell_offsets[i] > self.cell_offsets[i + 1] {
return Err(GpuError::DriverCallFailed {
reason: format!(
"bms_flex_row inputs: cell_offsets must be monotone (offset[{}]={} > offset[{}]={})",
i,
self.cell_offsets[i],
i + 1,
self.cell_offsets[i + 1]
),
});
}
}
Ok(())
}
}
#[cfg(target_os = "linux")]
const CUDA_ROW_KERNEL_TEMPLATE: &str = r#"
// One block per row. threadIdx.x parallelises per-cell sums.
// Semantic calibration/finalization visits are generated from BmsFlexRowProgram.
#define INV_TWO_PI 0.15915494309189535
#define BMS_FLEX_ROW_THREADS /*__BMS_FLEX_ROW_THREADS__*/
extern "C" __device__ __forceinline__ double atomic_add_f64(double *addr, double value) {
unsigned long long int *addr_as_ull = (unsigned long long int *)addr;
unsigned long long int old = *addr_as_ull;
unsigned long long int assumed;
do {
assumed = old;
double next = __longlong_as_double((long long int)assumed) + value;
old = atomicCAS(addr_as_ull, assumed, (unsigned long long int)__double_as_longlong(next));
} while (assumed != old);
return __longlong_as_double((long long int)old);
}
// `nan_fill_outputs`: thread-0-only path used when row inputs are degenerate
// (`F_a` non-finite or non-positive). The status channel makes the host reject
// the entire selected-GPU execution before any output can enter a cache.
extern "C" __device__ __forceinline__ void
nan_fill_outputs(int r,
int row,
double *out_neglog,
double *out_grad,
double *out_hess,
unsigned int *out_status) {
double nan_value = __longlong_as_double(0x7ff8000000000000ULL);
out_status[row] = 1U;
out_neglog[row] = nan_value;
size_t row_r = (size_t)row * (size_t)r;
for (int u = 0; u < r; ++u) {
out_grad[row_r + (size_t)u] = nan_value;
}
size_t rr = (size_t)r * (size_t)r;
size_t row_rr = (size_t)row * rr;
for (size_t idx = 0; idx < rr; ++idx) {
out_hess[row_rr + idx] = nan_value;
}
}
extern "C" __global__ void bms_flex_row_kernel(
int n_rows,
int r,
int p_h,
int p_w,
double s_f, // currently unused on device:
// host has already baked S_f
// into the cubic coefficients.
// Kept for diagnostic parity.
const double * __restrict__ row_q,
const double * __restrict__ row_b,
const double * __restrict__ row_mu1,
const double * __restrict__ row_mu2,
const double * __restrict__ row_zobs,
const double * __restrict__ row_y,
const double * __restrict__ row_w,
const unsigned int * __restrict__ cell_offsets,
const double * __restrict__ cell_c0,
const double * __restrict__ cell_c1,
const double * __restrict__ cell_c2,
const double * __restrict__ cell_c3,
const double * __restrict__ cell_a, // [n_cells, 4]
const double * __restrict__ cell_aa, // [n_cells, 4]
const double * __restrict__ cell_r, // [n_cells, r-1, 4]
const double * __restrict__ cell_ar, // [n_cells, r-1, 4]
const double * __restrict__ cell_sbb, // [n_cells, 4]
const double * __restrict__ cell_sbh, // [n_cells, p_h, 4]
const double * __restrict__ cell_sbw, // [n_cells, p_w, 4]
const double * __restrict__ cell_moments, // [n_cells, 10]
const double * __restrict__ row_chi,
const double * __restrict__ row_xi,
const double * __restrict__ row_rho, // [n_rows, r]
const double * __restrict__ row_tau, // [n_rows, r]
const double * __restrict__ row_ruv, // [n_rows, r*r]
const double * __restrict__ row_e_obs, // [n_rows] observed predictor VALUE
double * __restrict__ row_f_au, // [n_rows, r] general-width scratch
double * __restrict__ out_neglog,
double * __restrict__ out_grad,
double * __restrict__ out_hess,
unsigned int * __restrict__ out_status)
{
int row = blockIdx.x;
if (row >= n_rows) return;
int tid = threadIdx.x;
// Width-general row scratch. Reuse the final output allocations in-place:
// F_u → a_u → gradient and F_uv → a_uv → Hessian. Only F_au needs one
// additional checked [n,r] device allocation.
size_t row_r_base = (size_t)row * (size_t)r;
size_t rr = (size_t)r * (size_t)r;
size_t row_rr_base = (size_t)row * rr;
double *F_u = out_grad + row_r_base;
double *F_au = row_f_au + row_r_base;
double *F_uv = out_hess + row_rr_base;
__shared__ double reduce_a[BMS_FLEX_ROW_THREADS];
__shared__ double reduce_b[BMS_FLEX_ROW_THREADS];
__shared__ double F_a_shared;
__shared__ double F_aa_shared;
// Zero scratch.
if (tid == 0) { F_a_shared = 0.0; F_aa_shared = 0.0; }
for (int u = tid; u < r; u += blockDim.x) {
F_u[u] = 0.0;
F_au[u] = 0.0;
}
for (size_t uv = (size_t)tid; uv < rr; uv += (size_t)blockDim.x) {
F_uv[uv] = 0.0;
}
__syncthreads();
// ── per-cell sweep ───────────────────────────────────────────────────
unsigned int cell_lo = cell_offsets[row];
unsigned int cell_hi = cell_offsets[row + 1];
unsigned int n_cells = cell_hi - cell_lo;
double local_Fa = 0.0;
double local_Faa = 0.0;
for (unsigned int local_c = (unsigned int)tid;
local_c < n_cells;
local_c += (unsigned int)blockDim.x) {
unsigned int c = cell_lo + local_c;
// Load cubic predictor coeffs C0..C3.
double C[4];
C[0] = cell_c0[c]; C[1] = cell_c1[c];
C[2] = cell_c2[c]; C[3] = cell_c3[c];
// Load m_0..m_9.
const double *m = cell_moments + (size_t)c * 10;
// T_n = κ · Σ_e C_e · m_{e+n}, n = 0..6.
// CPU parity: equivalent to the `eta_rs ⊗ moments` contraction in
// `cell_second_derivative_from_moments` after folding the
// cubic predictor.
double T[7];
#pragma unroll
for (int n = 0; n < 7; ++n) {
double acc = 0.0;
#pragma unroll
for (int e = 0; e < 4; ++e) {
acc = fma(C[e], m[e + n], acc);
}
T[n] = acc * INV_TWO_PI;
}
// D(R) = κ · Σ_k R_k · m_k.
// CPU parity: `cell_first_derivative_from_moments`.
// The argument is parenthesized because callers pass pointer
// ARITHMETIC (`D_OF(base + offset)`): without it the expansion binds
// as `base + offset[0]`, which NVRTC rejects ("pointer-to-object
// type" on the integer term) — the calibration-phase emitter was the
// first caller to hit this.
#define D_OF(R) (INV_TWO_PI * ((R)[0]*m[0] + (R)[1]*m[1] + (R)[2]*m[2] + (R)[3]*m[3]))
// Q(R, S) = Σ_{p,q} R_p · S_q · T_{p+q}.
// CPU parity: the `eta_rs` folded dot in
// `cell_second_derivative_from_moments`.
#define Q_OF(R, S) \
(((R)[0]*(S)[0])*T[0] + ((R)[0]*(S)[1] + (R)[1]*(S)[0])*T[1] \
+ ((R)[0]*(S)[2] + (R)[1]*(S)[1] + (R)[2]*(S)[0])*T[2] \
+ ((R)[0]*(S)[3] + (R)[1]*(S)[2] + (R)[2]*(S)[1] + (R)[3]*(S)[0])*T[3] \
+ ((R)[1]*(S)[3] + (R)[2]*(S)[2] + (R)[3]*(S)[1])*T[4] \
+ ((R)[2]*(S)[3] + (R)[3]*(S)[2])*T[5] \
+ ((R)[3]*(S)[3])*T[6])
// The typed calibration schedule below consumes these primitive
// coefficient views through D_OF/Q_OF.
const double *A_c = cell_a + (size_t)c * 4;
const double *AA_c = cell_aa + (size_t)c * 4;
/*__BMS_FLEX_CALIBRATION_ORDER2__*/
#undef D_OF
#undef Q_OF
}
// Block reduction of local_Fa, local_Faa into shared.
reduce_a[tid] = local_Fa;
reduce_b[tid] = local_Faa;
__syncthreads();
for (int stride = blockDim.x / 2; stride > 0; stride >>= 1) {
if (tid < stride) {
reduce_a[tid] += reduce_a[tid + stride];
reduce_b[tid] += reduce_b[tid + stride];
}
__syncthreads();
}
if (tid == 0) {
F_a_shared = reduce_a[0];
F_aa_shared = reduce_b[0];
}
__syncthreads();
// ── thread-0 finalisation: IFT + observed-point + Mills + writes ──────
if (tid != 0) return;
double F_a = F_a_shared;
double F_aa = F_aa_shared;
double mu_1 = row_mu1[row];
double mu_2 = row_mu2[row];
// q-row overrides.
// F_q = -mu_1 ; F_qq = -mu_2 ; F_qv = 0 (v > 0) ; F_aq = 0.
F_u[0] = -mu_1;
F_au[0] = 0.0;
// Zero the q-cross row/column of F_uv (u == 0 or v == 0), then plant -mu_2 at (0,0).
for (int v = 0; v < r; ++v) {
F_uv[(size_t)v] = 0.0;
F_uv[(size_t)v * (size_t)r] = 0.0;
}
F_uv[0] = -mu_2;
// Guard: degenerate F_a ⇒ NaN-fill this row's outputs.
if (!isfinite(F_a) || F_a <= 0.0) {
nan_fill_outputs(r, row, out_neglog, out_grad, out_hess, out_status);
return;
}
double inv_Fa = 1.0 / F_a;
// Storage consumed by the generated dependency-ordered finalizer. Both
// aliases overwrite their no-longer-needed calibration predecessors.
double *a_u = F_u;
double *a_uv = F_uv;
double chi = row_chi[row];
double xi = row_xi[row];
const double *rho = row_rho + (size_t)row * r;
const double *tau = row_tau + (size_t)row * r;
const double *ruv = row_ruv + row_rr_base;
// Probit Mills.
double y = row_y[row];
double w = row_w[row];
double s = 2.0 * y - 1.0;
// The "observed predictor" e_obs is the VALUE (degree-0 term) of the
// observed jet η(a(θ), θ; z_obs) — NOT `bar_e_u[0]`, which is the u=0
// FIRST-derivative jet (`chi·a_0 + rho_0 = dη_obs/dq`). The host packs
// the observed value directly in `row_e_obs[row]` (see
// `pack_bms_flex_row_kernel_inputs`, `eta_val = eval_coeff4_at(obs.coeff,
// z_obs)`), matching the CPU family `lower_bms_flex_row_order2_from_parts`
// which forms `signed_margin = s_y · eta_val`. #415 parity lock.
double e_obs = row_e_obs[row];
double m_arg = s * e_obs;
double log_cdf, lambda, probit_curvature;
log_ndtr_mills_curvature(m_arg, &log_cdf, &lambda, &probit_curvature);
double A_i = -w * s * lambda;
double B_i = w * probit_curvature;
out_neglog[row] = -w * log_cdf;
/*__BMS_FLEX_ORDER2_FINALIZER__*/
if (!isfinite(out_neglog[row])) {
out_status[row] = 2U;
}
for (int u = 0; u < r; ++u) {
if (!isfinite(out_grad[row_r_base + (size_t)u])) {
out_status[row] = 2U;
}
}
for (size_t uv = 0; uv < rr; ++uv) {
if (!isfinite(out_hess[row_rr_base + uv])) {
out_status[row] = 2U;
}
}
}
"#;
#[cfg(target_os = "linux")]
fn build_generated_row_kernel_source() -> String {
const CALIBRATION_MARKER: &str = " /*__BMS_FLEX_CALIBRATION_ORDER2__*/";
const FINALIZER_MARKER: &str = " /*__BMS_FLEX_ORDER2_FINALIZER__*/";
let (prefix, remainder) = CUDA_ROW_KERNEL_TEMPLATE
.split_once(CALIBRATION_MARKER)
.expect("CUDA row template must contain the calibration marker");
let (between, suffix) = remainder
.split_once(FINALIZER_MARKER)
.expect("CUDA row template must contain the finalizer marker");
let mut source = String::with_capacity(CUDA_ROW_KERNEL_TEMPLATE.len() + 16_000);
source.push_str(prefix);
BmsFlexRowProgram::try_for_each_calibration_order2_phase(
true,
|phase| -> Result<(), std::convert::Infallible> {
source.push_str(&format!(
" // canonical calibration phase: {phase:?}\n"
));
match phase {
BmsFlexCalibrationOrder2Phase::InterceptFirst => {
source.push_str(" local_Fa += D_OF(A_c);\n");
}
BmsFlexCalibrationOrder2Phase::InterceptSecond => {
source.push_str(" local_Faa += D_OF(AA_c) - Q_OF(A_c, A_c);\n");
}
BmsFlexCalibrationOrder2Phase::PrimaryFirstAndInterceptSecond => {
source.push_str(
r#" for (int u = 1; u < r; ++u) {
const double *R_u = cell_r + ((size_t)c * (size_t)(r - 1) + (size_t)(u - 1)) * 4;
const double *AR_u = cell_ar + ((size_t)c * (size_t)(r - 1) + (size_t)(u - 1)) * 4;
atomic_add_f64(&F_u[u], D_OF(R_u));
atomic_add_f64(&F_au[u], D_OF(AR_u) - Q_OF(A_c, R_u));
}
"#,
);
}
BmsFlexCalibrationOrder2Phase::PrimaryPairSecond => {
source.push_str(
r#" for (int u = 1; u < r; ++u) {
const double *R_u = cell_r + ((size_t)c * (size_t)(r - 1) + (size_t)(u - 1)) * 4;
for (int v = u; v < r; ++v) {
const double *R_v = cell_r + ((size_t)c * (size_t)(r - 1) + (size_t)(v - 1)) * 4;
double explicit_second = 0.0;
if (u == 1 && v == 1) {
explicit_second = D_OF(cell_sbb + (size_t)c * 4);
} else if (u == 1 && v < 2 + p_h) {
int j = v - 2;
explicit_second = D_OF(cell_sbh + ((size_t)c * (size_t)p_h + (size_t)j) * 4);
} else if (u == 1) {
int l = v - (2 + p_h);
explicit_second = D_OF(cell_sbw + ((size_t)c * (size_t)p_w + (size_t)l) * 4);
}
atomic_add_f64(&F_uv[(size_t)u * (size_t)r + (size_t)v], explicit_second - Q_OF(R_u, R_v));
}
}
"#,
);
}
}
Ok(())
},
)
.expect("the infallible calibration phase emitter cannot fail");
source.push_str(between);
BmsFlexRowProgram::try_for_each_order2_finalizer_phase(
true,
|phase| -> Result<(), std::convert::Infallible> {
source.push_str(&format!(" // canonical finalizer phase: {phase:?}\n"));
match phase {
BmsFlexRowOrder2FinalizerPhase::ImplicitFirst => {
source.push_str(
r#" for (int u = 0; u < r; ++u) {
a_u[u] = -F_u[u] * inv_Fa;
}
"#,
);
}
BmsFlexRowOrder2FinalizerPhase::ImplicitFirstComplete => {
source.push_str(" // Canonical implicit-first stage complete.\n");
}
BmsFlexRowOrder2FinalizerPhase::ImplicitSecond => {
source.push_str(
r#" for (int u = 0; u < r; ++u) {
for (int v = u; v < r; ++v) {
size_t uv = (size_t)u * (size_t)r + (size_t)v;
size_t vu = (size_t)v * (size_t)r + (size_t)u;
double term = F_uv[uv]
+ F_au[v] * a_u[u]
+ F_au[u] * a_u[v]
+ F_aa * a_u[u] * a_u[v];
double value = -term * inv_Fa;
a_uv[uv] = value;
a_uv[vu] = value;
}
}
"#,
);
}
BmsFlexRowOrder2FinalizerPhase::ObservedFirst => {
source.push_str(" // Observed first derivatives are derived on demand.\n");
}
BmsFlexRowOrder2FinalizerPhase::ObservedScoreSensitivity => {
source.push_str(
" // Score sensitivity has no Stage-2 device output channel.\n",
);
}
BmsFlexRowOrder2FinalizerPhase::ObservedSecond => {
source.push_str(
r#" for (int u = 0; u < r; ++u) {
for (int v = u; v < r; ++v) {
size_t uv = (size_t)u * (size_t)r + (size_t)v;
size_t vu = (size_t)v * (size_t)r + (size_t)u;
double bar_e_u = chi * a_u[u] + rho[u];
double bar_e_v = chi * a_u[v] + rho[v];
double observed_second = chi * a_uv[uv]
+ xi * a_u[u] * a_u[v]
+ tau[u] * a_u[v]
+ a_u[u] * tau[v]
+ ruv[uv];
double hessian_value =
B_i * bar_e_u * bar_e_v + A_i * observed_second;
out_hess[row_rr_base + uv] = hessian_value;
out_hess[row_rr_base + vu] = hessian_value;
}
}
"#,
);
}
BmsFlexRowOrder2FinalizerPhase::NegLogFirst => {
source.push_str(
r#" for (int u = 0; u < r; ++u) {
double bar_e_u = chi * a_u[u] + rho[u];
out_grad[row_r_base + (size_t)u] = A_i * bar_e_u;
}
"#,
);
}
}
Ok(())
},
)
.expect("the infallible finalizer phase emitter cannot fail");
source.push_str(suffix);
source.replace(
"/*__BMS_FLEX_ROW_THREADS__*/",
&ROW_KERNEL_THREADS.to_string(),
)
}
#[cfg(target_os = "linux")]
pub(crate) fn generated_row_kernel_source() -> &'static str {
static SOURCE: OnceLock<String> = OnceLock::new();
SOURCE.get_or_init(build_generated_row_kernel_source)
}
#[inline]
pub(crate) fn s_f_diagnostic_finite(inputs: &BmsFlexRowKernelInputs<'_>) -> bool {
inputs.s_f.is_finite() && inputs.s_f > 0.0
}
#[cfg(target_os = "linux")]
pub(crate) struct RowKernelBackend {
pub(crate) stream: Arc<CudaStream>,
pub(crate) module: Arc<CudaModule>,
}
#[cfg(target_os = "linux")]
impl RowKernelBackend {
pub(crate) fn probe() -> Result<&'static Self, GpuError> {
static BACKEND: OnceLock<Result<RowKernelBackend, GpuError>> = OnceLock::new();
BACKEND
.get_or_init(|| {
gam_gpu::backend_probe::probe_backend_with_compile("bms_flex_row", |parts| {
let row_kernel_source = [
gam_gpu::numerics_device::PROBIT_NUMERICS_CU,
generated_row_kernel_source(),
]
.concat();
let ptx = gam_gpu::device_cache::compile_ptx_arch(&row_kernel_source).map_err(
|err| GpuError::DriverCallFailed {
reason: format!("bms_flex_row NVRTC compile failed: {err}"),
},
)?;
let module =
parts
.ctx
.load_module(ptx)
.map_err(|err| GpuError::DriverCallFailed {
reason: format!("bms_flex_row module load failed: {err}"),
})?;
Ok(RowKernelBackend {
stream: parts.stream.clone(),
module,
})
})
})
.as_ref()
.map_err(GpuError::clone)
}
}
pub(crate) fn launch_bms_flex_row_kernel(
inputs: BmsFlexRowKernelInputs<'_>,
) -> Result<BmsFlexRowKernelOutputs, GpuError> {
inputs.validate()?;
if !s_f_diagnostic_finite(&inputs) {
return Err(GpuError::DriverCallFailed {
reason: format!(
"bms_flex_row inputs: s_f must be positive and finite, got {}",
inputs.s_f
),
});
}
#[cfg(target_os = "linux")]
{
launch_linux(inputs)
}
#[cfg(not(target_os = "linux"))]
{
Err(GpuError::DriverLibraryUnavailable {
reason: "bms_flex_row GPU kernel is Linux-only".to_string(),
})
}
}
#[cfg(target_os = "linux")]
pub(crate) fn launch_linux(
inputs: BmsFlexRowKernelInputs<'_>,
) -> Result<BmsFlexRowKernelOutputs, GpuError> {
let backend = RowKernelBackend::probe()?;
let stream = &backend.stream;
let upload_f64 = |slice: &[f64], label: &str| {
stream
.clone_htod(slice)
.map_err(|err| GpuError::DriverCallFailed {
reason: format!("bms_flex_row upload {label}: {err}"),
})
};
let upload_u32 = |slice: &[u32], label: &str| {
stream
.clone_htod(slice)
.map_err(|err| GpuError::DriverCallFailed {
reason: format!("bms_flex_row upload {label}: {err}"),
})
};
let d_q = upload_f64(inputs.q, "q")?;
let d_b = upload_f64(inputs.b, "b")?;
let d_mu1 = upload_f64(inputs.mu_1, "mu_1")?;
let d_mu2 = upload_f64(inputs.mu_2, "mu_2")?;
let d_zobs = upload_f64(inputs.z_obs, "z_obs")?;
let d_y = upload_f64(inputs.y, "y")?;
let d_w = upload_f64(inputs.w, "w")?;
let d_offsets = upload_u32(inputs.cell_offsets, "cell_offsets")?;
let d_c0 = upload_f64(inputs.cell_c0, "cell_c0")?;
let d_c1 = upload_f64(inputs.cell_c1, "cell_c1")?;
let d_c2 = upload_f64(inputs.cell_c2, "cell_c2")?;
let d_c3 = upload_f64(inputs.cell_c3, "cell_c3")?;
let d_a = upload_f64(inputs.cell_a, "cell_a")?;
let d_aa = upload_f64(inputs.cell_aa, "cell_aa")?;
let d_r = upload_f64(inputs.cell_r, "cell_r")?;
let d_ar = upload_f64(inputs.cell_ar, "cell_ar")?;
let d_sbb = upload_f64(inputs.cell_sbb, "cell_sbb")?;
let d_sbh = upload_f64(inputs.cell_sbh, "cell_sbh")?;
let d_sbw = upload_f64(inputs.cell_sbw, "cell_sbw")?;
let owned_host_moments: CudaSlice<f64>;
let d_moments_ref: &CudaSlice<f64> = match &inputs.cell_moments {
CellMomentsSource::Host(slice) => {
owned_host_moments = upload_f64(slice, "cell_moments")?;
&owned_host_moments
}
CellMomentsSource::Device(d) => *d,
};
let d_chi = upload_f64(inputs.chi_obs, "chi_obs")?;
let d_xi = upload_f64(inputs.xi_obs, "xi_obs")?;
let d_rho = upload_f64(inputs.rho_u, "rho_u")?;
let d_tau = upload_f64(inputs.tau_u, "tau_u")?;
let d_ruv = upload_f64(inputs.r_uv, "r_uv")?;
let d_e_obs = upload_f64(inputs.e_obs, "e_obs")?;
let n = inputs.n_rows;
let r = inputs.r;
let nr = checked_shape_len("launch [n,r]", &[n, r])?;
let nrr = checked_shape_len("launch [n,r,r]", &[n, r, r])?;
let mut d_neglog = stream
.alloc_zeros::<f64>(n)
.map_err(|err| GpuError::DriverCallFailed {
reason: format!("bms_flex_row alloc neglog: {err}"),
})?;
let mut d_grad = stream
.alloc_zeros::<f64>(nr)
.map_err(|err| GpuError::DriverCallFailed {
reason: format!("bms_flex_row alloc grad: {err}"),
})?;
let mut d_hess = stream
.alloc_zeros::<f64>(nrr)
.map_err(|err| GpuError::DriverCallFailed {
reason: format!("bms_flex_row alloc hess: {err}"),
})?;
let mut d_f_au = stream
.alloc_zeros::<f64>(nr)
.map_err(|err| GpuError::DriverCallFailed {
reason: format!("bms_flex_row alloc F_au scratch: {err}"),
})?;
let mut d_status = stream
.alloc_zeros::<u32>(n)
.map_err(|err| GpuError::DriverCallFailed {
reason: format!("bms_flex_row alloc status: {err}"),
})?;
let func = backend
.module
.load_function("bms_flex_row_kernel")
.map_err(|err| GpuError::DriverCallFailed {
reason: format!("bms_flex_row load_function: {err}"),
})?;
let n_u32 = u32::try_from(n).map_err(|_| GpuError::DriverCallFailed {
reason: format!("bms_flex_row: n_rows={n} exceeds CUDA grid range"),
})?;
let cfg = LaunchConfig {
grid_dim: (n_u32, 1, 1),
block_dim: (ROW_KERNEL_THREADS, 1, 1),
shared_mem_bytes: 0,
};
let n_i32 = i32::try_from(n).map_err(|_| GpuError::DriverCallFailed {
reason: format!("bms_flex_row: n_rows={n} exceeds i32 range"),
})?;
let r_i32 = i32::try_from(r).map_err(|_| GpuError::DriverCallFailed {
reason: format!("bms_flex_row: r={r} exceeds i32 range"),
})?;
let p_h_i32 = i32::try_from(inputs.p_h).map_err(|_| GpuError::DriverCallFailed {
reason: format!("bms_flex_row: p_h={} exceeds i32 range", inputs.p_h),
})?;
let p_w_i32 = i32::try_from(inputs.p_w).map_err(|_| GpuError::DriverCallFailed {
reason: format!("bms_flex_row: p_w={} exceeds i32 range", inputs.p_w),
})?;
let s_f = inputs.s_f;
let mut builder = stream.launch_builder(&func);
builder
.arg(&n_i32)
.arg(&r_i32)
.arg(&p_h_i32)
.arg(&p_w_i32)
.arg(&s_f)
.arg(&d_q)
.arg(&d_b)
.arg(&d_mu1)
.arg(&d_mu2)
.arg(&d_zobs)
.arg(&d_y)
.arg(&d_w)
.arg(&d_offsets)
.arg(&d_c0)
.arg(&d_c1)
.arg(&d_c2)
.arg(&d_c3)
.arg(&d_a)
.arg(&d_aa)
.arg(&d_r)
.arg(&d_ar)
.arg(&d_sbb)
.arg(&d_sbh)
.arg(&d_sbw)
.arg(d_moments_ref)
.arg(&d_chi)
.arg(&d_xi)
.arg(&d_rho)
.arg(&d_tau)
.arg(&d_ruv)
.arg(&d_e_obs)
.arg(&mut d_f_au)
.arg(&mut d_neglog)
.arg(&mut d_grad)
.arg(&mut d_hess)
.arg(&mut d_status);
unsafe { builder.launch(cfg) }.map_err(|err| GpuError::DriverCallFailed {
reason: format!("bms_flex_row launch: {err}"),
})?;
stream
.synchronize()
.map_err(|err| GpuError::DriverCallFailed {
reason: format!("bms_flex_row synchronize: {err}"),
})?;
let status = stream
.clone_dtoh(&d_status)
.map_err(|err| GpuError::DriverCallFailed {
reason: format!("bms_flex_row download status: {err}"),
})?;
if let Some((row, code)) = status
.iter()
.copied()
.enumerate()
.find(|(_, code)| *code != 0)
{
return Err(GpuError::DriverCallFailed {
reason: format!("bms_flex_row rejected non-finite row {row} with status {code}"),
});
}
let neglog = stream
.clone_dtoh(&d_neglog)
.map_err(|err| GpuError::DriverCallFailed {
reason: format!("bms_flex_row download neglog: {err}"),
})?;
let grad = stream
.clone_dtoh(&d_grad)
.map_err(|err| GpuError::DriverCallFailed {
reason: format!("bms_flex_row download grad: {err}"),
})?;
let hess = stream
.clone_dtoh(&d_hess)
.map_err(|err| GpuError::DriverCallFailed {
reason: format!("bms_flex_row download hess: {err}"),
})?;
Ok(BmsFlexRowKernelOutputs { neglog, grad, hess })
}
#[cfg(target_os = "linux")]
#[derive(Clone, Debug)]
pub(crate) struct BmsFlexBlockLayout {
pub p_m: usize,
pub p_g: usize,
pub h: Option<std::ops::Range<usize>>,
pub w: Option<std::ops::Range<usize>>,
pub p_total: usize,
}
#[cfg(target_os = "linux")]
#[derive(Clone, Debug)]
pub(crate) struct BmsFlexPrimaryLayout {
pub h: Option<std::ops::Range<usize>>,
pub w: Option<std::ops::Range<usize>>,
pub r: usize,
}
#[cfg(target_os = "linux")]
pub(crate) const HVP_ROWS_PER_CTA: u32 = 256;
#[cfg(target_os = "linux")]
pub(crate) const HVP_THREADS: u32 = 128;
#[cfg(target_os = "linux")]
pub(crate) const REDUCTION_THREADS: u32 = 256;
#[cfg(target_os = "linux")]
pub(crate) const BMS_FLEX_ROW_HVP_MAX_RHS: usize = 8;
#[cfg(target_os = "linux")]
pub struct DeviceResidentRowHess {
pub(crate) neglog: CudaSlice<f64>,
pub(crate) grad: CudaSlice<f64>,
pub(crate) hess: CudaSlice<f64>,
pub(crate) marginal_design: CudaSlice<f64>,
pub(crate) logslope_design: CudaSlice<f64>,
pub(crate) n: usize,
pub(crate) r: usize,
pub(crate) block: BmsFlexBlockLayout,
pub(crate) primary: BmsFlexPrimaryLayout,
pub(crate) bytes: u64,
}
#[cfg(target_os = "linux")]
pub(crate) struct BmsFlexDeviceJointGradient {
pub(crate) log_likelihood: f64,
pub(crate) gradient: Vec<f64>,
}
#[cfg(target_os = "linux")]
impl std::fmt::Debug for DeviceResidentRowHess {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("DeviceResidentRowHess")
.field("n", &self.n)
.field("r", &self.r)
.field("p_total", &self.block.p_total)
.field("bytes", &self.bytes)
.finish()
}
}
#[cfg(target_os = "linux")]
pub(crate) fn num_hvp_chunks(n: usize) -> usize {
n.div_ceil(HVP_ROWS_PER_CTA as usize)
}
#[cfg(target_os = "linux")]
pub(crate) const HVP_KERNEL_SOURCE: &str = r#"
// CPU parity reference: cpu_oracle_bms_flex_row_hvp / cpu_oracle_bms_flex_row_diagonal
// in this module.
#define MAX_MULTI_RHS 8
__device__ __forceinline__ double bms_flex_primary_direction(
int primary_idx,
int h_block_start,
int h_block_len,
int w_block_start,
int w_block_len,
int h_primary_start,
int w_primary_start,
double direction_q,
double direction_g,
const double * __restrict__ v)
{
if (primary_idx == 0) return direction_q;
if (primary_idx == 1) return direction_g;
if (primary_idx >= h_primary_start && primary_idx < h_primary_start + h_block_len) {
return v[h_block_start + primary_idx - h_primary_start];
}
if (primary_idx >= w_primary_start && primary_idx < w_primary_start + w_block_len) {
return v[w_block_start + primary_idx - w_primary_start];
}
return 0.0;
}
extern "C" __global__ void bms_flex_row_hvp_partial(
int n_rows,
int r,
int p_m,
int p_g,
int p_total,
int h_block_start,
int h_block_len,
int w_block_start,
int w_block_len,
int h_primary_start,
int w_primary_start,
int rows_per_cta,
const double * __restrict__ row_hessians, // [n, r*r]
const double * __restrict__ marginal_design, // [n, p_m] row-major
const double * __restrict__ logslope_design, // [n, p_g] row-major
const double * __restrict__ v, // [p_total]
double * __restrict__ partial) // [num_chunks, p_total]
{
int chunk = blockIdx.x;
int tid = threadIdx.x;
int row_lo = chunk * rows_per_cta;
int remaining_rows = n_rows - row_lo;
int row_hi = row_lo + (remaining_rows < rows_per_cta ? remaining_rows : rows_per_cta);
// Zero this chunk's partial slice cooperatively.
double *out = partial + (size_t)chunk * (size_t)p_total;
for (int j = tid; j < p_total; j += blockDim.x) {
out[j] = 0.0;
}
__syncthreads();
// Width-general scratch: only the two design directions/actions are
// shared. Every h/w direction is read directly from v, and the thread
// owning primary coordinate u accumulates that coordinate's action.
__shared__ double direction_q;
__shared__ double direction_g;
__shared__ double action_q;
__shared__ double action_g;
__shared__ double dot_reduce[128];
for (int row = row_lo; row < row_hi; ++row) {
const double *mrow = marginal_design + (size_t)row * (size_t)p_m;
const double *grow = logslope_design + (size_t)row * (size_t)p_g;
const double *Hrow = row_hessians + (size_t)row * (size_t)r * (size_t)r;
// row_dir[0] = mrow · v[0..p_m]
double local = 0.0;
for (int j = tid; j < p_m; j += blockDim.x) {
local += mrow[j] * v[j];
}
dot_reduce[tid] = local;
__syncthreads();
for (int stride = blockDim.x / 2; stride > 0; stride >>= 1) {
if (tid < stride) dot_reduce[tid] += dot_reduce[tid + stride];
__syncthreads();
}
if (tid == 0) direction_q = dot_reduce[0];
// row_dir[1] = grow · v[p_m..p_m+p_g]
local = 0.0;
for (int j = tid; j < p_g; j += blockDim.x) {
local += grow[j] * v[p_m + j];
}
dot_reduce[tid] = local;
__syncthreads();
for (int stride = blockDim.x / 2; stride > 0; stride >>= 1) {
if (tid < stride) dot_reduce[tid] += dot_reduce[tid + stride];
__syncthreads();
}
if (tid == 0) direction_g = dot_reduce[0];
__syncthreads();
for (int u = tid; u < r; u += blockDim.x) {
double acc = 0.0;
for (int vv = 0; vv < r; ++vv) {
double row_direction = bms_flex_primary_direction(
vv,
h_block_start, h_block_len,
w_block_start, w_block_len,
h_primary_start, w_primary_start,
direction_q, direction_g, v);
acc += Hrow[(size_t)u * (size_t)r + (size_t)vv] * row_direction;
}
if (u == 0) {
action_q = acc;
} else if (u == 1) {
action_g = acc;
} else if (u >= h_primary_start && u < h_primary_start + h_block_len) {
out[h_block_start + u - h_primary_start] += acc;
} else if (u >= w_primary_start && u < w_primary_start + w_block_len) {
out[w_block_start + u - w_primary_start] += acc;
}
}
__syncthreads();
// Pull back into joint β slot.
double a0 = action_q;
for (int j = tid; j < p_m; j += blockDim.x) {
out[j] += a0 * mrow[j];
}
double a1 = action_g;
for (int j = tid; j < p_g; j += blockDim.x) {
out[p_m + j] += a1 * grow[j];
}
__syncthreads();
}
}
extern "C" __global__ void bms_flex_row_hvp_reduce(
int num_chunks,
int p_total,
const double * __restrict__ partial, // [num_chunks, p_total]
double * __restrict__ out) // [p_total]
{
int j = blockIdx.x * blockDim.x + threadIdx.x;
if (j >= p_total) return;
double acc = 0.0;
for (int c = 0; c < num_chunks; ++c) {
acc += partial[(size_t)c * (size_t)p_total + (size_t)j];
}
out[j] = acc;
}
extern "C" __global__ void bms_flex_row_joint_gradient_partial(
int n_rows,
int r,
int p_m,
int p_g,
int p_total,
int h_block_start,
int h_block_len,
int w_block_start,
int w_block_len,
int h_primary_start,
int w_primary_start,
int rows_per_cta,
const double * __restrict__ row_neglog, // [n]
const double * __restrict__ row_grad, // [n, r]
const double * __restrict__ marginal_design, // [n, p_m]
const double * __restrict__ logslope_design, // [n, p_g]
double * __restrict__ partial) // [num_chunks, 1+p_total]
{
int chunk = blockIdx.x;
int tid = threadIdx.x;
int row_lo = chunk * rows_per_cta;
int remaining_rows = n_rows - row_lo;
int row_hi = row_lo + (remaining_rows < rows_per_cta ? remaining_rows : rows_per_cta);
int output_width = p_total + 1;
double *out = partial + (size_t)chunk * (size_t)output_width;
// One thread owns each output coordinate for the whole row chunk. The
// inner row loop therefore has a fixed order and needs no atomics.
for (int output_idx = tid; output_idx < output_width; output_idx += blockDim.x) {
double acc = 0.0;
if (output_idx == 0) {
for (int row = row_lo; row < row_hi; ++row) {
acc -= row_neglog[row];
}
} else {
int beta_idx = output_idx - 1;
if (beta_idx < p_m) {
for (int row = row_lo; row < row_hi; ++row) {
acc -= row_grad[(size_t)row * (size_t)r]
* marginal_design[(size_t)row * (size_t)p_m + (size_t)beta_idx];
}
} else if (beta_idx < p_m + p_g) {
int j = beta_idx - p_m;
for (int row = row_lo; row < row_hi; ++row) {
acc -= row_grad[(size_t)row * (size_t)r + 1]
* logslope_design[(size_t)row * (size_t)p_g + (size_t)j];
}
} else if (beta_idx >= h_block_start && beta_idx < h_block_start + h_block_len) {
int primary_idx = h_primary_start + beta_idx - h_block_start;
for (int row = row_lo; row < row_hi; ++row) {
acc -= row_grad[(size_t)row * (size_t)r + (size_t)primary_idx];
}
} else if (beta_idx >= w_block_start && beta_idx < w_block_start + w_block_len) {
int primary_idx = w_primary_start + beta_idx - w_block_start;
for (int row = row_lo; row < row_hi; ++row) {
acc -= row_grad[(size_t)row * (size_t)r + (size_t)primary_idx];
}
}
}
out[output_idx] = acc;
}
}
extern "C" __global__ void bms_flex_row_joint_gradient_reduce(
int num_chunks,
int output_width,
const double * __restrict__ partial, // [num_chunks, output_width]
double * __restrict__ out) // [output_width]
{
int j = blockIdx.x * blockDim.x + threadIdx.x;
if (j >= output_width) return;
double acc = 0.0;
for (int c = 0; c < num_chunks; ++c) {
acc += partial[(size_t)c * (size_t)output_width + (size_t)j];
}
out[j] = acc;
}
extern "C" __global__ void bms_flex_row_hvp_multi_partial(
int n_rows,
int r,
int p_m,
int p_g,
int p_total,
int h_block_start,
int h_block_len,
int w_block_start,
int w_block_len,
int h_primary_start,
int w_primary_start,
int rows_per_cta,
int rhs_count,
const double * __restrict__ row_hessians, // [n, r*r]
const double * __restrict__ marginal_design, // [n, p_m]
const double * __restrict__ logslope_design, // [n, p_g]
const double * __restrict__ v_rhs, // [rhs_count, p_total]
double * __restrict__ partial) // [rhs_count, num_chunks, p_total]
{
int chunk = blockIdx.x;
int tid = threadIdx.x;
int row_lo = chunk * rows_per_cta;
int remaining_rows = n_rows - row_lo;
int row_hi = row_lo + (remaining_rows < rows_per_cta ? remaining_rows : rows_per_cta);
int num_chunks = 1 + (n_rows - 1) / rows_per_cta;
for (int idx = tid; idx < rhs_count * p_total; idx += blockDim.x) {
int rhs = idx / p_total;
int j = idx - rhs * p_total;
partial[((size_t)rhs * (size_t)num_chunks + (size_t)chunk) * (size_t)p_total + (size_t)j] = 0.0;
}
__syncthreads();
__shared__ double direction_q[MAX_MULTI_RHS];
__shared__ double direction_g[MAX_MULTI_RHS];
__shared__ double action_q[MAX_MULTI_RHS];
__shared__ double action_g[MAX_MULTI_RHS];
__shared__ double dot_reduce[128];
for (int row = row_lo; row < row_hi; ++row) {
const double *mrow = marginal_design + (size_t)row * (size_t)p_m;
const double *grow = logslope_design + (size_t)row * (size_t)p_g;
const double *Hrow = row_hessians + (size_t)row * (size_t)r * (size_t)r;
for (int rhs = 0; rhs < rhs_count; ++rhs) {
const double *v = v_rhs + (size_t)rhs * (size_t)p_total;
double local = 0.0;
for (int j = tid; j < p_m; j += blockDim.x) {
local += mrow[j] * v[j];
}
dot_reduce[tid] = local;
__syncthreads();
for (int stride = blockDim.x / 2; stride > 0; stride >>= 1) {
if (tid < stride) dot_reduce[tid] += dot_reduce[tid + stride];
__syncthreads();
}
if (tid == 0) direction_q[rhs] = dot_reduce[0];
local = 0.0;
for (int j = tid; j < p_g; j += blockDim.x) {
local += grow[j] * v[p_m + j];
}
dot_reduce[tid] = local;
__syncthreads();
for (int stride = blockDim.x / 2; stride > 0; stride >>= 1) {
if (tid < stride) dot_reduce[tid] += dot_reduce[tid + stride];
__syncthreads();
}
if (tid == 0) direction_g[rhs] = dot_reduce[0];
__syncthreads();
}
size_t total_actions = (size_t)rhs_count * (size_t)r;
for (size_t idx = (size_t)tid; idx < total_actions; idx += (size_t)blockDim.x) {
int rhs = (int)(idx / (size_t)r);
int u = (int)(idx - (size_t)rhs * (size_t)r);
const double *v = v_rhs + (size_t)rhs * (size_t)p_total;
double *out = partial + ((size_t)rhs * (size_t)num_chunks + (size_t)chunk) * (size_t)p_total;
double acc = 0.0;
for (int vv = 0; vv < r; ++vv) {
double row_direction = bms_flex_primary_direction(
vv,
h_block_start, h_block_len,
w_block_start, w_block_len,
h_primary_start, w_primary_start,
direction_q[rhs], direction_g[rhs], v);
acc += Hrow[(size_t)u * (size_t)r + (size_t)vv] * row_direction;
}
if (u == 0) {
action_q[rhs] = acc;
} else if (u == 1) {
action_g[rhs] = acc;
} else if (u >= h_primary_start && u < h_primary_start + h_block_len) {
out[h_block_start + u - h_primary_start] += acc;
} else if (u >= w_primary_start && u < w_primary_start + w_block_len) {
out[w_block_start + u - w_primary_start] += acc;
}
}
__syncthreads();
for (int rhs = 0; rhs < rhs_count; ++rhs) {
double *out = partial + ((size_t)rhs * (size_t)num_chunks + (size_t)chunk) * (size_t)p_total;
double a0 = action_q[rhs];
for (int j = tid; j < p_m; j += blockDim.x) {
out[j] += a0 * mrow[j];
}
double a1 = action_g[rhs];
for (int j = tid; j < p_g; j += blockDim.x) {
out[p_m + j] += a1 * grow[j];
}
__syncthreads();
}
}
}
extern "C" __global__ void bms_flex_row_hvp_multi_reduce(
int num_chunks,
int p_total,
int rhs_count,
const double * __restrict__ partial, // [rhs_count, num_chunks, p_total]
double * __restrict__ out) // [rhs_count, p_total]
{
int idx = blockIdx.x * blockDim.x + threadIdx.x;
int total = rhs_count * p_total;
if (idx >= total) return;
int rhs = idx / p_total;
int j = idx - rhs * p_total;
double acc = 0.0;
for (int c = 0; c < num_chunks; ++c) {
acc += partial[((size_t)rhs * (size_t)num_chunks + (size_t)c) * (size_t)p_total + (size_t)j];
}
out[(size_t)rhs * (size_t)p_total + (size_t)j] = acc;
}
extern "C" __global__ void bms_flex_row_diag_partial(
int n_rows,
int r,
int p_m,
int p_g,
int p_total,
int h_block_start,
int h_block_len,
int w_block_start,
int w_block_len,
int h_primary_start,
int w_primary_start,
int rows_per_cta,
const double * __restrict__ row_hessians,
const double * __restrict__ marginal_design,
const double * __restrict__ logslope_design,
double * __restrict__ partial)
{
int chunk = blockIdx.x;
int tid = threadIdx.x;
int row_lo = chunk * rows_per_cta;
int remaining_rows = n_rows - row_lo;
int row_hi = row_lo + (remaining_rows < rows_per_cta ? remaining_rows : rows_per_cta);
double *out = partial + (size_t)chunk * (size_t)p_total;
for (int j = tid; j < p_total; j += blockDim.x) {
out[j] = 0.0;
}
__syncthreads();
for (int row = row_lo; row < row_hi; ++row) {
const double *mrow = marginal_design + (size_t)row * (size_t)p_m;
const double *grow = logslope_design + (size_t)row * (size_t)p_g;
const double *Hrow = row_hessians + (size_t)row * (size_t)r * (size_t)r;
double h00 = Hrow[0];
double h11 = Hrow[(size_t)r + 1U];
for (int j = tid; j < p_m; j += blockDim.x) {
double v = mrow[j];
out[j] += h00 * v * v;
}
for (int j = tid; j < p_g; j += blockDim.x) {
double v = grow[j];
out[p_m + j] += h11 * v * v;
}
if (tid == 0) {
for (int k = 0; k < h_block_len; ++k) {
int ii = h_primary_start + k;
out[h_block_start + k] +=
Hrow[(size_t)ii * (size_t)r + (size_t)ii];
}
for (int k = 0; k < w_block_len; ++k) {
int ii = w_primary_start + k;
out[w_block_start + k] +=
Hrow[(size_t)ii * (size_t)r + (size_t)ii];
}
}
__syncthreads();
}
}
// ────────────────────────────────────────────────────────────────────────
// Phase 6 — dense joint-Hessian block kernel for the debug / exact-REML
// route. Materialises the full `[p_total, p_total]` row-major joint H
// from the per-row r×r Hessian via the P_i pullback. NOT the default
// Newton path: production Newton uses HVP (Phase 2/3); this kernel exists
// for exact-REML logdet / dense-H comparisons / diagnostic dumps where the
// caller genuinely needs the dense matrix on the device.
//
// Per-CTA partial: each CTA owns a contiguous chunk of rows
// `[chunk*rows_per_cta, (chunk+1)*rows_per_cta)`. Inside the CTA the
// per-row pullback computes `(P_i^T H_i P_i)[m, n]` and adds it to the
// CTA's shared-mem `[p_total, p_total]` partial. The reduce kernel sums
// chunk-major-fixed-order into a single `[p_total, p_total]` output.
//
// Math: for primary index u ∈ [0, r):
// * u = 0: phi_u = (X_i in slot 0..p_m, 0 elsewhere)
// * u = 1: phi_u = (0, G_i in slot p_m..p_m+p_g, 0 elsewhere)
// * u = 2+j: phi_u = e_{h_block_start + j} (j ∈ 0..h_block_len)
// * u = 2+h+l: phi_u = e_{w_block_start + l} (l ∈ 0..w_block_len)
// Then `H_full[m, n] += sum_{u,v} H_i[u,v] * phi_u[m] * phi_v[n]`.
//
// Shared-memory budget: at large-scale shape p_total = 44, a [44, 44] f64
// partial is 44*44*8 = 15.5 KiB — well below the V100 48 KiB/SM cap.
// At p_total ≤ 80 the kernel still fits (80*80*8 = 50 KiB → just over
// V100 cap; caller must enforce p_total ≤ DENSE_BLOCK_MAX_P). The
// launcher rejects oversize p_total cleanly.
extern "C" __global__ void bms_flex_row_dense_block_partial(
int n_rows,
int r,
int p_m,
int p_g,
int p_total,
int h_block_start,
int h_block_len,
int w_block_start,
int w_block_len,
int h_primary_start,
int w_primary_start,
int rows_per_cta,
const double * __restrict__ row_hessians, // [n, r*r]
const double * __restrict__ marginal_design, // [n, p_m]
const double * __restrict__ logslope_design, // [n, p_g]
double * __restrict__ partial) // [num_chunks, p_total, p_total]
{
extern __shared__ double shmem[];
int chunk = blockIdx.x;
int tid = threadIdx.x;
int row_lo = chunk * rows_per_cta;
int remaining_rows = n_rows - row_lo;
int row_hi = row_lo + (remaining_rows < rows_per_cta ? remaining_rows : rows_per_cta);
int pp = p_total * p_total;
double *acc = shmem; // CTA-private accumulator [p_total, p_total]
for (int j = tid; j < pp; j += blockDim.x) acc[j] = 0.0;
__syncthreads();
// Per-row work performed by thread 0 to avoid cross-thread RW
// contention on `acc[]`. Per-row complexity is O(r² + p_total²); the host
// selects this direct algorithm only for small p_total, while r remains a
// checked runtime width with no semantic ceiling.
// Tighter parallel implementations are possible (warp-stripe the
// 4-way nested u-v-m-n loop) but Phase 6 is a debug-only path and
// the simple version is easier to audit for correctness against
// the host-side P_i pullback oracle.
if (tid == 0) {
for (int row = row_lo; row < row_hi; ++row) {
const double *mrow = marginal_design + (size_t)row * (size_t)p_m;
const double *grow = logslope_design + (size_t)row * (size_t)p_g;
const double *Hrow = row_hessians + (size_t)row * (size_t)r * (size_t)r;
for (int u = 0; u < r; ++u) {
for (int v = 0; v < r; ++v) {
double huv = Hrow[(size_t)u * (size_t)r + (size_t)v];
if (huv == 0.0) continue;
// For each (u, v), iterate (m, n) over the non-zero
// outer-product support of phi_u and phi_v.
// Build a small (offset, len, src_ptr) descriptor for
// each operand block as we go.
int m_off, m_len; const double *m_src; bool m_indicator;
int n_off, n_len; const double *n_src; bool n_indicator;
if (u == 0) { m_off = 0; m_len = p_m; m_src = mrow; m_indicator = false; }
else if (u == 1) { m_off = p_m; m_len = p_g; m_src = grow; m_indicator = false; }
else if (u - 2 < h_block_len) {
m_off = h_block_start + (u - 2);
m_len = 1; m_src = NULL; m_indicator = true;
} else {
m_off = w_block_start + (u - 2 - h_block_len);
m_len = 1; m_src = NULL; m_indicator = true;
}
if (v == 0) { n_off = 0; n_len = p_m; n_src = mrow; n_indicator = false; }
else if (v == 1) { n_off = p_m; n_len = p_g; n_src = grow; n_indicator = false; }
else if (v - 2 < h_block_len) {
n_off = h_block_start + (v - 2);
n_len = 1; n_src = NULL; n_indicator = true;
} else {
n_off = w_block_start + (v - 2 - h_block_len);
n_len = 1; n_src = NULL; n_indicator = true;
}
// accumulate huv * phi_u[m] * phi_v[n] into acc[m, n]
for (int mi = 0; mi < m_len; ++mi) {
double pm = m_indicator ? 1.0 : m_src[mi];
if (pm == 0.0) continue;
double scaled = huv * pm;
int m_idx = m_off + mi;
for (int ni = 0; ni < n_len; ++ni) {
double pn = n_indicator ? 1.0 : n_src[ni];
int n_idx = n_off + ni;
acc[m_idx * p_total + n_idx] += scaled * pn;
}
}
}
}
}
}
__syncthreads();
// Write CTA accumulator out to global memory at its chunk slot.
double *out_chunk = partial + (size_t)chunk * (size_t)pp;
for (int j = tid; j < pp; j += blockDim.x) {
out_chunk[j] = acc[j];
}
}
extern "C" __global__ void bms_flex_row_dense_block_reduce(
int num_chunks,
int p_total,
const double * __restrict__ partial,
double * __restrict__ out)
{
int j = blockIdx.x * blockDim.x + threadIdx.x;
int pp = p_total * p_total;
if (j >= pp) return;
double acc = 0.0;
for (int c = 0; c < num_chunks; ++c) {
acc += partial[(size_t)c * (size_t)pp + (size_t)j];
}
out[j] = acc;
}
"#;
#[cfg(target_os = "linux")]
pub(crate) struct HvpKernelBackend {
pub(crate) stream: Arc<CudaStream>,
pub(crate) module: Arc<CudaModule>,
}
#[cfg(target_os = "linux")]
impl HvpKernelBackend {
pub(crate) fn probe() -> Result<&'static Self, GpuError> {
static BACKEND: OnceLock<Result<HvpKernelBackend, GpuError>> = OnceLock::new();
BACKEND
.get_or_init(|| {
gam_gpu::backend_probe::probe_backend_with_compile("bms_flex_row hvp", |parts| {
let ptx = gam_gpu::device_cache::compile_ptx_arch(HVP_KERNEL_SOURCE).map_err(
|err| GpuError::DriverCallFailed {
reason: format!("bms_flex_row hvp NVRTC compile failed: {err}"),
},
)?;
let module =
parts
.ctx
.load_module(ptx)
.map_err(|err| GpuError::DriverCallFailed {
reason: format!("bms_flex_row hvp module load failed: {err}"),
})?;
Ok(HvpKernelBackend {
stream: parts.stream.clone(),
module,
})
})
})
.as_ref()
.map_err(GpuError::clone)
}
}
#[cfg(target_os = "linux")]
pub(crate) fn launch_bms_flex_row_kernel_device_resident(
inputs: BmsFlexRowKernelInputs<'_>,
marginal_design_row_major: &[f64],
logslope_design_row_major: &[f64],
block: BmsFlexBlockLayout,
primary: BmsFlexPrimaryLayout,
) -> Result<DeviceResidentRowHess, GpuError> {
inputs.validate()?;
if !s_f_diagnostic_finite(&inputs) {
return Err(GpuError::DriverCallFailed {
reason: format!(
"bms_flex_row device-resident: s_f must be positive and finite, got {}",
inputs.s_f
),
});
}
let n = inputs.n_rows;
let r = inputs.r;
let nr = checked_shape_len("device-resident [n,r]", &[n, r])?;
let nrr = checked_shape_len("device-resident [n,r,r]", &[n, r, r])?;
let marginal_len = checked_shape_len("device-resident marginal design", &[n, block.p_m])?;
let logslope_len = checked_shape_len("device-resident logslope design", &[n, block.p_g])?;
if marginal_design_row_major.len() != marginal_len {
return Err(GpuError::DriverCallFailed {
reason: format!(
"bms_flex_row device-resident: marginal_design len={} != n*p_m={}",
marginal_design_row_major.len(),
marginal_len
),
});
}
if logslope_design_row_major.len() != logslope_len {
return Err(GpuError::DriverCallFailed {
reason: format!(
"bms_flex_row device-resident: logslope_design len={} != n*p_g={}",
logslope_design_row_major.len(),
logslope_len
),
});
}
if primary.r != r {
return Err(GpuError::DriverCallFailed {
reason: format!(
"bms_flex_row device-resident: primary.r={} != inputs.r={}",
primary.r, r
),
});
}
let backend = RowKernelBackend::probe()?;
HvpKernelBackend::probe()?;
let stream = backend.stream.clone();
let upload_f64 = |slice: &[f64], label: &str| {
stream
.clone_htod(slice)
.map_err(|err| GpuError::DriverCallFailed {
reason: format!("bms_flex_row device-resident upload {label}: {err}"),
})
};
let upload_u32 = |slice: &[u32], label: &str| {
stream
.clone_htod(slice)
.map_err(|err| GpuError::DriverCallFailed {
reason: format!("bms_flex_row device-resident upload {label}: {err}"),
})
};
let d_q = upload_f64(inputs.q, "q")?;
let d_b = upload_f64(inputs.b, "b")?;
let d_mu1 = upload_f64(inputs.mu_1, "mu_1")?;
let d_mu2 = upload_f64(inputs.mu_2, "mu_2")?;
let d_zobs = upload_f64(inputs.z_obs, "z_obs")?;
let d_y = upload_f64(inputs.y, "y")?;
let d_w = upload_f64(inputs.w, "w")?;
let d_offsets = upload_u32(inputs.cell_offsets, "cell_offsets")?;
let d_c0 = upload_f64(inputs.cell_c0, "cell_c0")?;
let d_c1 = upload_f64(inputs.cell_c1, "cell_c1")?;
let d_c2 = upload_f64(inputs.cell_c2, "cell_c2")?;
let d_c3 = upload_f64(inputs.cell_c3, "cell_c3")?;
let d_a = upload_f64(inputs.cell_a, "cell_a")?;
let d_aa = upload_f64(inputs.cell_aa, "cell_aa")?;
let d_r = upload_f64(inputs.cell_r, "cell_r")?;
let d_ar = upload_f64(inputs.cell_ar, "cell_ar")?;
let d_sbb = upload_f64(inputs.cell_sbb, "cell_sbb")?;
let d_sbh = upload_f64(inputs.cell_sbh, "cell_sbh")?;
let d_sbw = upload_f64(inputs.cell_sbw, "cell_sbw")?;
let owned_host_moments: CudaSlice<f64>;
let d_moments_ref: &CudaSlice<f64> = match &inputs.cell_moments {
CellMomentsSource::Host(slice) => {
owned_host_moments = upload_f64(slice, "cell_moments")?;
&owned_host_moments
}
CellMomentsSource::Device(d) => *d,
};
let d_chi = upload_f64(inputs.chi_obs, "chi_obs")?;
let d_xi = upload_f64(inputs.xi_obs, "xi_obs")?;
let d_rho = upload_f64(inputs.rho_u, "rho_u")?;
let d_tau = upload_f64(inputs.tau_u, "tau_u")?;
let d_ruv = upload_f64(inputs.r_uv, "r_uv")?;
let d_e_obs = upload_f64(inputs.e_obs, "e_obs")?;
let d_marginal = upload_f64(marginal_design_row_major, "marginal_design")?;
let d_logslope = upload_f64(logslope_design_row_major, "logslope_design")?;
let mut d_neglog = stream
.alloc_zeros::<f64>(n)
.map_err(|err| GpuError::DriverCallFailed {
reason: format!("bms_flex_row device-resident alloc neglog: {err}"),
})?;
let mut d_grad = stream
.alloc_zeros::<f64>(nr)
.map_err(|err| GpuError::DriverCallFailed {
reason: format!("bms_flex_row device-resident alloc grad: {err}"),
})?;
let mut d_hess = stream
.alloc_zeros::<f64>(nrr)
.map_err(|err| GpuError::DriverCallFailed {
reason: format!("bms_flex_row device-resident alloc hess: {err}"),
})?;
let mut d_f_au = stream
.alloc_zeros::<f64>(nr)
.map_err(|err| GpuError::DriverCallFailed {
reason: format!("bms_flex_row device-resident alloc F_au scratch: {err}"),
})?;
let mut d_status = stream
.alloc_zeros::<u32>(n)
.map_err(|err| GpuError::DriverCallFailed {
reason: format!("bms_flex_row device-resident alloc status: {err}"),
})?;
let func = backend
.module
.load_function("bms_flex_row_kernel")
.map_err(|err| GpuError::DriverCallFailed {
reason: format!("bms_flex_row device-resident load_function: {err}"),
})?;
let n_u32 = u32::try_from(n).map_err(|_| GpuError::DriverCallFailed {
reason: format!("bms_flex_row device-resident: n_rows={n} exceeds CUDA grid range"),
})?;
let cfg = LaunchConfig {
grid_dim: (n_u32, 1, 1),
block_dim: (ROW_KERNEL_THREADS, 1, 1),
shared_mem_bytes: 0,
};
let n_i32 = i32::try_from(n).map_err(|_| GpuError::DriverCallFailed {
reason: format!("bms_flex_row device-resident: n_rows={n} exceeds i32 range"),
})?;
let r_i32 = i32::try_from(r).map_err(|_| GpuError::DriverCallFailed {
reason: format!("bms_flex_row device-resident: r={r} exceeds i32 range"),
})?;
let p_h_i32 = i32::try_from(inputs.p_h).map_err(|_| GpuError::DriverCallFailed {
reason: format!(
"bms_flex_row device-resident: p_h={} exceeds i32 range",
inputs.p_h
),
})?;
let p_w_i32 = i32::try_from(inputs.p_w).map_err(|_| GpuError::DriverCallFailed {
reason: format!(
"bms_flex_row device-resident: p_w={} exceeds i32 range",
inputs.p_w
),
})?;
let s_f_val = inputs.s_f;
let mut builder = stream.launch_builder(&func);
builder
.arg(&n_i32)
.arg(&r_i32)
.arg(&p_h_i32)
.arg(&p_w_i32)
.arg(&s_f_val)
.arg(&d_q)
.arg(&d_b)
.arg(&d_mu1)
.arg(&d_mu2)
.arg(&d_zobs)
.arg(&d_y)
.arg(&d_w)
.arg(&d_offsets)
.arg(&d_c0)
.arg(&d_c1)
.arg(&d_c2)
.arg(&d_c3)
.arg(&d_a)
.arg(&d_aa)
.arg(&d_r)
.arg(&d_ar)
.arg(&d_sbb)
.arg(&d_sbh)
.arg(&d_sbw)
.arg(d_moments_ref)
.arg(&d_chi)
.arg(&d_xi)
.arg(&d_rho)
.arg(&d_tau)
.arg(&d_ruv)
.arg(&d_e_obs)
.arg(&mut d_f_au)
.arg(&mut d_neglog)
.arg(&mut d_grad)
.arg(&mut d_hess)
.arg(&mut d_status);
unsafe { builder.launch(cfg) }.map_err(|err| GpuError::DriverCallFailed {
reason: format!("bms_flex_row device-resident launch: {err}"),
})?;
stream
.synchronize()
.map_err(|err| GpuError::DriverCallFailed {
reason: format!("bms_flex_row device-resident synchronize: {err}"),
})?;
let status = stream
.clone_dtoh(&d_status)
.map_err(|err| GpuError::DriverCallFailed {
reason: format!("bms_flex_row device-resident download status: {err}"),
})?;
if let Some((row, code)) = status
.iter()
.copied()
.enumerate()
.find(|(_, code)| *code != 0)
{
return Err(GpuError::DriverCallFailed {
reason: format!(
"bms_flex_row device-resident rejected non-finite row {row} with status {code}"
),
});
}
drop(d_status);
drop(d_f_au);
drop(d_q);
drop(d_b);
drop(d_mu1);
drop(d_mu2);
drop(d_zobs);
drop(d_y);
drop(d_w);
drop(d_offsets);
drop(d_c0);
drop(d_c1);
drop(d_c2);
drop(d_c3);
drop(d_a);
drop(d_aa);
drop(d_r);
drop(d_ar);
drop(d_sbb);
drop(d_sbh);
drop(d_sbw);
drop(d_chi);
drop(d_xi);
drop(d_rho);
drop(d_tau);
drop(d_ruv);
let resident_elements = n
.checked_add(nr)
.and_then(|value| value.checked_add(nrr))
.and_then(|value| value.checked_add(marginal_len))
.and_then(|value| value.checked_add(logslope_len))
.ok_or_else(|| GpuError::DriverCallFailed {
reason: "bms_flex_row device-resident: resident element count overflow".to_string(),
})?;
let resident_bytes = resident_elements
.checked_mul(std::mem::size_of::<f64>())
.ok_or_else(|| GpuError::DriverCallFailed {
reason: "bms_flex_row device-resident: resident byte count overflow".to_string(),
})?;
let bytes = u64::try_from(resident_bytes).map_err(|_| GpuError::DriverCallFailed {
reason: format!(
"bms_flex_row device-resident: resident bytes={resident_bytes} exceed u64 range"
),
})?;
Ok(DeviceResidentRowHess {
neglog: d_neglog,
grad: d_grad,
hess: d_hess,
marginal_design: d_marginal,
logslope_design: d_logslope,
n,
r,
block,
primary,
bytes,
})
}
#[cfg(target_os = "linux")]
pub(crate) fn launch_bms_flex_row_joint_gradient(
storage: &DeviceResidentRowHess,
) -> Result<BmsFlexDeviceJointGradient, GpuError> {
let p_total = storage.block.p_total;
let output_width = p_total
.checked_add(1)
.ok_or_else(|| GpuError::DriverCallFailed {
reason: "bms_flex_row joint gradient: output width overflow".to_string(),
})?;
if storage.n == 0 {
return Ok(BmsFlexDeviceJointGradient {
log_likelihood: 0.0,
gradient: vec![0.0; p_total],
});
}
let backend = HvpKernelBackend::probe()?;
let stream = backend.stream.clone();
let args = PreparedBmsFlexRowLaunchArgs::from_storage(storage)?;
let partial_len = args
.num_chunks
.checked_mul(output_width)
.ok_or_else(|| GpuError::DriverCallFailed {
reason: format!(
"bms_flex_row joint gradient: partial length overflow for chunks={} width={output_width}",
args.num_chunks
),
})?;
let mut d_partial =
stream
.alloc_zeros::<f64>(partial_len)
.map_err(|err| GpuError::DriverCallFailed {
reason: format!("bms_flex_row joint gradient alloc partial: {err}"),
})?;
let mut d_out =
stream
.alloc_zeros::<f64>(output_width)
.map_err(|err| GpuError::DriverCallFailed {
reason: format!("bms_flex_row joint gradient alloc output: {err}"),
})?;
let partial_func = backend
.module
.load_function("bms_flex_row_joint_gradient_partial")
.map_err(|err| GpuError::DriverCallFailed {
reason: format!("bms_flex_row joint gradient load partial: {err}"),
})?;
let reduce_func = backend
.module
.load_function("bms_flex_row_joint_gradient_reduce")
.map_err(|err| GpuError::DriverCallFailed {
reason: format!("bms_flex_row joint gradient load reduce: {err}"),
})?;
let num_chunks_u32 =
u32::try_from(args.num_chunks).map_err(|_| GpuError::DriverCallFailed {
reason: format!(
"bms_flex_row joint gradient: num_chunks={} exceeds u32 range",
args.num_chunks
),
})?;
let cfg_partial = LaunchConfig {
grid_dim: (num_chunks_u32, 1, 1),
block_dim: (HVP_THREADS, 1, 1),
shared_mem_bytes: 0,
};
let mut builder = stream.launch_builder(&partial_func);
builder
.arg(&args.n_i32)
.arg(&args.r_i32)
.arg(&args.p_m_i32)
.arg(&args.p_g_i32)
.arg(&args.p_total_i32)
.arg(&args.h_block_start)
.arg(&args.h_block_len)
.arg(&args.w_block_start)
.arg(&args.w_block_len)
.arg(&args.h_primary_start)
.arg(&args.w_primary_start)
.arg(&args.rows_per_cta)
.arg(&storage.neglog)
.arg(&storage.grad)
.arg(&storage.marginal_design)
.arg(&storage.logslope_design)
.arg(&mut d_partial);
unsafe { builder.launch(cfg_partial) }.map_err(|err| GpuError::DriverCallFailed {
reason: format!("bms_flex_row joint gradient partial launch: {err}"),
})?;
let output_width_i32 = i32::try_from(output_width).map_err(|_| GpuError::DriverCallFailed {
reason: format!(
"bms_flex_row joint gradient: output_width={output_width} exceeds i32 range"
),
})?;
let num_chunks_i32 =
i32::try_from(args.num_chunks).map_err(|_| GpuError::DriverCallFailed {
reason: format!(
"bms_flex_row joint gradient: num_chunks={} exceeds i32 range",
args.num_chunks
),
})?;
let output_width_u32 = u32::try_from(output_width).map_err(|_| GpuError::DriverCallFailed {
reason: format!(
"bms_flex_row joint gradient: output_width={output_width} exceeds u32 range"
),
})?;
let reduce_blocks = output_width_u32.div_ceil(REDUCTION_THREADS);
let cfg_reduce = LaunchConfig {
grid_dim: (reduce_blocks, 1, 1),
block_dim: (REDUCTION_THREADS, 1, 1),
shared_mem_bytes: 0,
};
let mut builder = stream.launch_builder(&reduce_func);
builder
.arg(&num_chunks_i32)
.arg(&output_width_i32)
.arg(&d_partial)
.arg(&mut d_out);
unsafe { builder.launch(cfg_reduce) }.map_err(|err| GpuError::DriverCallFailed {
reason: format!("bms_flex_row joint gradient reduce launch: {err}"),
})?;
stream
.synchronize()
.map_err(|err| GpuError::DriverCallFailed {
reason: format!("bms_flex_row joint gradient synchronize: {err}"),
})?;
let host = stream
.clone_dtoh(&d_out)
.map_err(|err| GpuError::DriverCallFailed {
reason: format!("bms_flex_row joint gradient download: {err}"),
})?;
if let Some((index, value)) = host
.iter()
.copied()
.enumerate()
.find(|(_, value)| !value.is_finite())
{
return Err(GpuError::DriverCallFailed {
reason: format!(
"bms_flex_row joint gradient produced non-finite output[{index}]={value}"
),
});
}
Ok(BmsFlexDeviceJointGradient {
log_likelihood: host[0],
gradient: host[1..].to_vec(),
})
}
#[cfg(target_os = "linux")]
#[derive(Clone, Copy)]
pub(crate) enum BmsFlexRowLaunchMode {
HvpDeviceOut,
DiagonalHostOut,
}
#[cfg(target_os = "linux")]
impl BmsFlexRowLaunchMode {
pub(crate) fn partial_kernel_name(self) -> &'static str {
match self {
BmsFlexRowLaunchMode::HvpDeviceOut => "bms_flex_row_hvp_partial",
BmsFlexRowLaunchMode::DiagonalHostOut => "bms_flex_row_diag_partial",
}
}
}
#[cfg(target_os = "linux")]
pub(crate) struct PreparedBmsFlexRowLaunchArgs {
pub(crate) n_i32: i32,
pub(crate) r_i32: i32,
pub(crate) p_m_i32: i32,
pub(crate) p_g_i32: i32,
pub(crate) p_total_i32: i32,
pub(crate) h_block_start: i32,
pub(crate) h_block_len: i32,
pub(crate) w_block_start: i32,
pub(crate) w_block_len: i32,
pub(crate) h_primary_start: i32,
pub(crate) w_primary_start: i32,
pub(crate) rows_per_cta: i32,
pub(crate) num_chunks: usize,
pub(crate) num_chunks_i32: i32,
pub(crate) num_chunks_u32: u32,
pub(crate) p_total_u32: u32,
}
#[cfg(target_os = "linux")]
impl PreparedBmsFlexRowLaunchArgs {
pub(crate) fn from_storage(storage: &DeviceResidentRowHess) -> Result<Self, GpuError> {
if storage.n == 0 {
return Err(GpuError::DriverCallFailed {
reason: "bms_flex_row launch: n_rows must be > 0".to_string(),
});
}
if storage.r < 2 {
return Err(GpuError::DriverCallFailed {
reason: format!("bms_flex_row launch: r={} must be >= 2", storage.r),
});
}
let p_total = storage.block.p_total;
if p_total == 0 {
return Err(GpuError::DriverCallFailed {
reason: "bms_flex_row launch: p_total must be > 0".to_string(),
});
}
if storage.primary.r != storage.r {
return Err(GpuError::DriverCallFailed {
reason: format!(
"bms_flex_row launch: primary.r={} != storage.r={}",
storage.primary.r, storage.r
),
});
}
let h_block_len = storage.block.h.as_ref().map_or(0, |range| range.len());
let w_block_len = storage.block.w.as_ref().map_or(0, |range| range.len());
let h_primary_len = storage.primary.h.as_ref().map_or(0, |range| range.len());
let w_primary_len = storage.primary.w.as_ref().map_or(0, |range| range.len());
if h_block_len != h_primary_len || w_block_len != w_primary_len {
return Err(GpuError::DriverCallFailed {
reason: format!(
"bms_flex_row launch: block/primary direct lengths disagree: h={h_block_len}/{h_primary_len}, w={w_block_len}/{w_primary_len}"
),
});
}
let h_block_start = storage
.block
.p_m
.checked_add(storage.block.p_g)
.ok_or_else(|| GpuError::DriverCallFailed {
reason: "bms_flex_row launch: p_m+p_g overflow".to_string(),
})?;
let w_block_start =
h_block_start
.checked_add(h_block_len)
.ok_or_else(|| GpuError::DriverCallFailed {
reason: "bms_flex_row launch: h block end overflow".to_string(),
})?;
let expected_p_total =
w_block_start
.checked_add(w_block_len)
.ok_or_else(|| GpuError::DriverCallFailed {
reason: "bms_flex_row launch: w block end overflow".to_string(),
})?;
let w_primary_start =
2_usize
.checked_add(h_primary_len)
.ok_or_else(|| GpuError::DriverCallFailed {
reason: "bms_flex_row launch: h primary end overflow".to_string(),
})?;
let expected_r = w_primary_start.checked_add(w_primary_len).ok_or_else(|| {
GpuError::DriverCallFailed {
reason: "bms_flex_row launch: w primary end overflow".to_string(),
}
})?;
let check_range = |name: &str,
range: Option<&std::ops::Range<usize>>,
expected_start: usize,
expected_len: usize|
-> Result<(), GpuError> {
match (range, expected_len) {
(None, 0) => Ok(()),
(Some(range), len)
if len > 0
&& range.start == expected_start
&& range.end == expected_start + len =>
{
Ok(())
}
_ => Err(GpuError::DriverCallFailed {
reason: format!(
"bms_flex_row launch: {name}={range:?} must be {expected_start}..{}",
expected_start + expected_len
),
}),
}
};
check_range(
"block.h",
storage.block.h.as_ref(),
h_block_start,
h_block_len,
)?;
check_range(
"block.w",
storage.block.w.as_ref(),
w_block_start,
w_block_len,
)?;
check_range("primary.h", storage.primary.h.as_ref(), 2, h_primary_len)?;
check_range(
"primary.w",
storage.primary.w.as_ref(),
w_primary_start,
w_primary_len,
)?;
if p_total != expected_p_total || storage.r != expected_r {
return Err(GpuError::DriverCallFailed {
reason: format!(
"bms_flex_row launch: inconsistent layout p_total={p_total}/{expected_p_total}, r={}/{}",
storage.r, expected_r
),
});
}
let expected_nr = checked_shape_len("launch storage [n,r]", &[storage.n, storage.r])?;
let expected_nrr =
checked_shape_len("launch storage [n,r,r]", &[storage.n, storage.r, storage.r])?;
let expected_marginal = checked_shape_len(
"launch storage marginal design",
&[storage.n, storage.block.p_m],
)?;
let expected_logslope = checked_shape_len(
"launch storage logslope design",
&[storage.n, storage.block.p_g],
)?;
for (name, have, want) in [
("neglog", storage.neglog.len(), storage.n),
("grad", storage.grad.len(), expected_nr),
("hess", storage.hess.len(), expected_nrr),
(
"marginal_design",
storage.marginal_design.len(),
expected_marginal,
),
(
"logslope_design",
storage.logslope_design.len(),
expected_logslope,
),
] {
if have != want {
return Err(GpuError::DriverCallFailed {
reason: format!("bms_flex_row launch: storage {name}.len()={have} != {want}"),
});
}
}
let num_chunks = num_hvp_chunks(storage.n);
let to_i32 = |name: &str, value: usize| {
i32::try_from(value).map_err(|_| GpuError::DriverCallFailed {
reason: format!("bms_flex_row launch: {name}={value} exceeds i32 range"),
})
};
let to_u32 = |name: &str, value: usize| {
u32::try_from(value).map_err(|_| GpuError::DriverCallFailed {
reason: format!("bms_flex_row launch: {name}={value} exceeds u32 range"),
})
};
Ok(PreparedBmsFlexRowLaunchArgs {
n_i32: to_i32("n_rows", storage.n)?,
r_i32: to_i32("r", storage.r)?,
p_m_i32: to_i32("p_m", storage.block.p_m)?,
p_g_i32: to_i32("p_g", storage.block.p_g)?,
p_total_i32: to_i32("p_total", p_total)?,
h_block_start: storage
.block
.h
.as_ref()
.map(|range| to_i32("h_block_start", range.start))
.transpose()?
.unwrap_or(0),
h_block_len: storage
.block
.h
.as_ref()
.map(|range| to_i32("h_block_len", range.len()))
.transpose()?
.unwrap_or(0),
w_block_start: storage
.block
.w
.as_ref()
.map(|range| to_i32("w_block_start", range.start))
.transpose()?
.unwrap_or(0),
w_block_len: storage
.block
.w
.as_ref()
.map(|range| to_i32("w_block_len", range.len()))
.transpose()?
.unwrap_or(0),
h_primary_start: storage
.primary
.h
.as_ref()
.map(|range| to_i32("h_primary_start", range.start))
.transpose()?
.unwrap_or(0),
w_primary_start: storage
.primary
.w
.as_ref()
.map(|range| to_i32("w_primary_start", range.start))
.transpose()?
.unwrap_or(0),
rows_per_cta: i32::try_from(HVP_ROWS_PER_CTA).map_err(|_| {
GpuError::DriverCallFailed {
reason: format!(
"bms_flex_row launch: rows_per_cta={HVP_ROWS_PER_CTA} exceeds i32 range"
),
}
})?,
num_chunks,
num_chunks_i32: to_i32("num_chunks", num_chunks)?,
num_chunks_u32: to_u32("num_chunks", num_chunks)?,
p_total_u32: to_u32("p_total", p_total)?,
})
}
}
#[cfg(target_os = "linux")]
pub(crate) fn run_bms_flex_row_partial_reduce(
storage: &DeviceResidentRowHess,
mode: BmsFlexRowLaunchMode,
d_v: Option<&CudaSlice<f64>>,
d_out: &mut CudaSlice<f64>,
ctx: &str,
) -> Result<(), GpuError> {
let backend = HvpKernelBackend::probe()?;
let stream = backend.stream.clone();
let args = PreparedBmsFlexRowLaunchArgs::from_storage(storage)?;
let p_total = storage.block.p_total;
let partial_len = checked_shape_len(
&format!("{ctx} partial [num_chunks,p_total]"),
&[args.num_chunks, p_total],
)?;
let mut d_partial =
stream
.alloc_zeros::<f64>(partial_len)
.map_err(|err| GpuError::DriverCallFailed {
reason: format!("bms_flex_row {ctx} alloc partial: {err}"),
})?;
let partial_kernel_name = mode.partial_kernel_name();
let part_func = backend
.module
.load_function(partial_kernel_name)
.map_err(|err| GpuError::DriverCallFailed {
reason: format!("bms_flex_row {ctx} load {partial_kernel_name}: {err}"),
})?;
let red_func = backend
.module
.load_function("bms_flex_row_hvp_reduce")
.map_err(|err| GpuError::DriverCallFailed {
reason: format!("bms_flex_row {ctx} load reduce: {err}"),
})?;
let cfg_part = LaunchConfig {
grid_dim: (args.num_chunks_u32, 1, 1),
block_dim: (HVP_THREADS, 1, 1),
shared_mem_bytes: 0,
};
let mut builder = stream.launch_builder(&part_func);
builder
.arg(&args.n_i32)
.arg(&args.r_i32)
.arg(&args.p_m_i32)
.arg(&args.p_g_i32)
.arg(&args.p_total_i32)
.arg(&args.h_block_start)
.arg(&args.h_block_len)
.arg(&args.w_block_start)
.arg(&args.w_block_len)
.arg(&args.h_primary_start)
.arg(&args.w_primary_start)
.arg(&args.rows_per_cta)
.arg(&storage.hess)
.arg(&storage.marginal_design)
.arg(&storage.logslope_design);
if let Some(d_v) = d_v {
builder.arg(d_v);
}
builder.arg(&mut d_partial);
unsafe { builder.launch(cfg_part) }.map_err(|err| GpuError::DriverCallFailed {
reason: format!("bms_flex_row {ctx} partial launch: {err}"),
})?;
let red_threads: u32 = REDUCTION_THREADS;
let red_blocks = args.p_total_u32.div_ceil(red_threads);
let cfg_red = LaunchConfig {
grid_dim: (red_blocks, 1, 1),
block_dim: (red_threads, 1, 1),
shared_mem_bytes: 0,
};
let mut builder = stream.launch_builder(&red_func);
builder
.arg(&args.num_chunks_i32)
.arg(&args.p_total_i32)
.arg(&d_partial)
.arg(d_out);
unsafe { builder.launch(cfg_red) }.map_err(|err| GpuError::DriverCallFailed {
reason: format!("bms_flex_row {ctx} reduce launch: {err}"),
})?;
drop(d_partial);
Ok(())
}
#[cfg(target_os = "linux")]
pub(crate) fn launch_bms_flex_row_diagonal_host(
storage: &DeviceResidentRowHess,
) -> Result<Vec<f64>, GpuError> {
let p_total = storage.block.p_total;
let backend = HvpKernelBackend::probe()?;
let stream = backend.stream.clone();
let mut d_out =
stream
.alloc_zeros::<f64>(p_total)
.map_err(|err| GpuError::DriverCallFailed {
reason: format!("bms_flex_row diag alloc out: {err}"),
})?;
run_bms_flex_row_partial_reduce(
storage,
BmsFlexRowLaunchMode::DiagonalHostOut,
None,
&mut d_out,
"diag",
)?;
stream
.synchronize()
.map_err(|err| GpuError::DriverCallFailed {
reason: format!("bms_flex_row diag synchronize: {err}"),
})?;
stream
.clone_dtoh(&d_out)
.map_err(|err| GpuError::DriverCallFailed {
reason: format!("bms_flex_row diag download out: {err}"),
})
}
#[cfg(target_os = "linux")]
pub(crate) fn validate_bms_flex_row_hvp_multi_shape(
storage: &DeviceResidentRowHess,
rhs_count: usize,
v_rhs_len: usize,
out_len: Option<usize>,
ctx: &str,
) -> Result<usize, GpuError> {
if rhs_count == 0 || rhs_count > BMS_FLEX_ROW_HVP_MAX_RHS {
return Err(GpuError::DriverCallFailed {
reason: format!(
"bms_flex_row {ctx}: rhs_count={rhs_count} outside 1..={BMS_FLEX_ROW_HVP_MAX_RHS}"
),
});
}
let p_total = storage.block.p_total;
let rhs_elems = rhs_count
.checked_mul(p_total)
.ok_or_else(|| GpuError::DriverCallFailed {
reason: format!(
"bms_flex_row {ctx}: rhs_count({rhs_count})*p_total({p_total}) overflow"
),
})?;
i32::try_from(rhs_elems).map_err(|_| GpuError::DriverCallFailed {
reason: format!(
"bms_flex_row {ctx}: rhs_count({rhs_count})*p_total({p_total})={rhs_elems} exceeds CUDA int indexing range"
),
})?;
if v_rhs_len != rhs_elems {
return Err(GpuError::DriverCallFailed {
reason: format!(
"bms_flex_row {ctx}: v_rhs.len()={v_rhs_len} != rhs_count({rhs_count})*p_total({p_total})={rhs_elems}"
),
});
}
if let Some(out_len) = out_len
&& out_len != rhs_elems
{
return Err(GpuError::DriverCallFailed {
reason: format!(
"bms_flex_row {ctx}: out.len()={out_len} != rhs_count({rhs_count})*p_total({p_total})={rhs_elems}"
),
});
}
Ok(rhs_elems)
}
#[cfg(target_os = "linux")]
pub fn bms_flex_row_hvp_multi_scratch_bytes_for_shape(
n: usize,
p_total: usize,
rhs_count: usize,
) -> Result<u64, GpuError> {
if rhs_count == 0 || rhs_count > BMS_FLEX_ROW_HVP_MAX_RHS {
return Err(GpuError::DriverCallFailed {
reason: format!(
"bms_flex_row hvp_multi_scratch_bytes: rhs_count={rhs_count} outside 1..={BMS_FLEX_ROW_HVP_MAX_RHS}"
),
});
}
let num_chunks = num_hvp_chunks(n);
let partial = rhs_count
.checked_mul(num_chunks)
.and_then(|v| v.checked_mul(p_total))
.ok_or_else(|| GpuError::DriverCallFailed {
reason: format!(
"bms_flex_row hvp_multi_scratch_bytes: rhs_count({rhs_count})*num_chunks({num_chunks})*p_total({p_total}) overflow"
),
})?;
let rhs_vectors = rhs_count
.checked_mul(p_total)
.and_then(|v| v.checked_mul(2))
.ok_or_else(|| GpuError::DriverCallFailed {
reason: format!(
"bms_flex_row hvp_multi_scratch_bytes: 2*rhs_count({rhs_count})*p_total({p_total}) overflow"
),
})?;
let elems = partial
.checked_add(rhs_vectors)
.ok_or_else(|| GpuError::DriverCallFailed {
reason: "bms_flex_row hvp_multi_scratch_bytes: element count overflow".to_string(),
})?;
let bytes = elems
.checked_mul(std::mem::size_of::<f64>())
.ok_or_else(|| GpuError::DriverCallFailed {
reason: "bms_flex_row hvp_multi_scratch_bytes: byte count overflow".to_string(),
})?;
u64::try_from(bytes).map_err(|_| GpuError::DriverCallFailed {
reason: format!(
"bms_flex_row hvp_multi_scratch_bytes: byte count={bytes} exceeds u64 range"
),
})
}
#[cfg(target_os = "linux")]
pub(crate) fn run_bms_flex_row_multi_partial_reduce(
storage: &DeviceResidentRowHess,
rhs_count: usize,
d_v_rhs: &CudaSlice<f64>,
d_out: &mut CudaSlice<f64>,
ctx: &str,
) -> Result<(), GpuError> {
let rhs_elems = validate_bms_flex_row_hvp_multi_shape(
storage,
rhs_count,
d_v_rhs.len(),
Some(d_out.len()),
ctx,
)?;
let backend = HvpKernelBackend::probe()?;
let stream = backend.stream.clone();
let args = PreparedBmsFlexRowLaunchArgs::from_storage(storage)?;
let p_total = storage.block.p_total;
let partial_len = rhs_count
.checked_mul(args.num_chunks)
.and_then(|v| v.checked_mul(p_total))
.ok_or_else(|| GpuError::DriverCallFailed {
reason: format!(
"bms_flex_row {ctx}: partial length overflow for rhs_count={rhs_count}, num_chunks={}, p_total={p_total}",
args.num_chunks
),
})?;
let mut d_partial =
stream
.alloc_zeros::<f64>(partial_len)
.map_err(|err| GpuError::DriverCallFailed {
reason: format!("bms_flex_row {ctx} alloc multi partial: {err}"),
})?;
let part_func = backend
.module
.load_function("bms_flex_row_hvp_multi_partial")
.map_err(|err| GpuError::DriverCallFailed {
reason: format!("bms_flex_row {ctx} load multi partial: {err}"),
})?;
let red_func = backend
.module
.load_function("bms_flex_row_hvp_multi_reduce")
.map_err(|err| GpuError::DriverCallFailed {
reason: format!("bms_flex_row {ctx} load multi reduce: {err}"),
})?;
let rhs_count_i32 = i32::try_from(rhs_count).map_err(|_| GpuError::DriverCallFailed {
reason: format!("bms_flex_row {ctx}: rhs_count={rhs_count} exceeds i32 range"),
})?;
let cfg_part = LaunchConfig {
grid_dim: (args.num_chunks_u32, 1, 1),
block_dim: (HVP_THREADS, 1, 1),
shared_mem_bytes: 0,
};
let mut builder = stream.launch_builder(&part_func);
builder
.arg(&args.n_i32)
.arg(&args.r_i32)
.arg(&args.p_m_i32)
.arg(&args.p_g_i32)
.arg(&args.p_total_i32)
.arg(&args.h_block_start)
.arg(&args.h_block_len)
.arg(&args.w_block_start)
.arg(&args.w_block_len)
.arg(&args.h_primary_start)
.arg(&args.w_primary_start)
.arg(&args.rows_per_cta)
.arg(&rhs_count_i32)
.arg(&storage.hess)
.arg(&storage.marginal_design)
.arg(&storage.logslope_design)
.arg(d_v_rhs)
.arg(&mut d_partial);
unsafe { builder.launch(cfg_part) }.map_err(|err| GpuError::DriverCallFailed {
reason: format!("bms_flex_row {ctx} multi partial launch: {err}"),
})?;
let red_threads: u32 = REDUCTION_THREADS;
let rhs_elems_u32 = u32::try_from(rhs_elems).map_err(|_| GpuError::DriverCallFailed {
reason: format!("bms_flex_row {ctx}: rhs elements={rhs_elems} exceed u32 range"),
})?;
let red_blocks = rhs_elems_u32.div_ceil(red_threads);
let cfg_red = LaunchConfig {
grid_dim: (red_blocks, 1, 1),
block_dim: (red_threads, 1, 1),
shared_mem_bytes: 0,
};
let mut builder = stream.launch_builder(&red_func);
builder
.arg(&args.num_chunks_i32)
.arg(&args.p_total_i32)
.arg(&rhs_count_i32)
.arg(&d_partial)
.arg(d_out);
unsafe { builder.launch(cfg_red) }.map_err(|err| GpuError::DriverCallFailed {
reason: format!("bms_flex_row {ctx} multi reduce launch: {err}"),
})?;
drop(d_partial);
Ok(())
}
#[cfg(target_os = "linux")]
pub(crate) fn launch_bms_flex_row_hvp_multi(
storage: &DeviceResidentRowHess,
v_rhs: &[f64],
rhs_count: usize,
) -> Result<Vec<f64>, GpuError> {
let rhs_elems =
validate_bms_flex_row_hvp_multi_shape(storage, rhs_count, v_rhs.len(), None, "hvp_multi")?;
let backend = HvpKernelBackend::probe()?;
let stream = backend.stream.clone();
let d_v_rhs = stream
.clone_htod(v_rhs)
.map_err(|err| GpuError::DriverCallFailed {
reason: format!("bms_flex_row hvp_multi upload v_rhs: {err}"),
})?;
let mut d_out =
stream
.alloc_zeros::<f64>(rhs_elems)
.map_err(|err| GpuError::DriverCallFailed {
reason: format!("bms_flex_row hvp_multi alloc out: {err}"),
})?;
run_bms_flex_row_multi_partial_reduce(storage, rhs_count, &d_v_rhs, &mut d_out, "hvp_multi")?;
stream
.synchronize()
.map_err(|err| GpuError::DriverCallFailed {
reason: format!("bms_flex_row hvp_multi synchronize: {err}"),
})?;
stream
.clone_dtoh(&d_out)
.map_err(|err| GpuError::DriverCallFailed {
reason: format!("bms_flex_row hvp_multi download out: {err}"),
})
}
#[cfg(target_os = "linux")]
fn materialize_dense_from_hvp_batches(
p_total: usize,
mut launch: impl FnMut(&[f64], usize) -> Result<Vec<f64>, GpuError>,
) -> Result<Vec<f64>, GpuError> {
if p_total == 0 {
return Err(GpuError::DriverCallFailed {
reason: "bms_flex_row dense HVP materialization: p_total must be > 0".to_string(),
});
}
let dense_len = p_total
.checked_mul(p_total)
.ok_or_else(|| GpuError::DriverCallFailed {
reason: format!(
"bms_flex_row dense HVP materialization: p_total={p_total} square overflow"
),
})?;
let mut dense = vec![0.0_f64; dense_len];
for column_start in (0..p_total).step_by(BMS_FLEX_ROW_HVP_MAX_RHS) {
let rhs_count = (p_total - column_start).min(BMS_FLEX_ROW_HVP_MAX_RHS);
let batch_len = checked_shape_len(
"dense HVP materialization [rhs_count,p_total]",
&[rhs_count, p_total],
)?;
let mut basis = vec![0.0_f64; batch_len];
for local_column in 0..rhs_count {
basis[local_column * p_total + column_start + local_column] = 1.0;
}
let images = launch(&basis, rhs_count)?;
if images.len() != basis.len() {
return Err(GpuError::DriverCallFailed {
reason: format!(
"bms_flex_row dense HVP materialization: batch at column {column_start} returned {} values, expected {}",
images.len(),
basis.len()
),
});
}
for local_column in 0..rhs_count {
let column = column_start + local_column;
let image = &images[local_column * p_total..(local_column + 1) * p_total];
for (row, &value) in image.iter().enumerate() {
dense[row * p_total + column] = value;
}
}
}
Ok(dense)
}
#[cfg(target_os = "linux")]
pub(crate) fn launch_bms_flex_row_hvp_into_device(
storage: &DeviceResidentRowHess,
d_v: &CudaSlice<f64>,
d_out: &mut CudaSlice<f64>,
) -> Result<(), GpuError> {
let p_total = storage.block.p_total;
if d_v.len() != p_total {
return Err(GpuError::DriverCallFailed {
reason: format!(
"bms_flex_row hvp_into_device: d_v.len()={} != p_total={}",
d_v.len(),
p_total
),
});
}
if d_out.len() != p_total {
return Err(GpuError::DriverCallFailed {
reason: format!(
"bms_flex_row hvp_into_device: d_out.len()={} != p_total={}",
d_out.len(),
p_total
),
});
}
run_bms_flex_row_partial_reduce(
storage,
BmsFlexRowLaunchMode::HvpDeviceOut,
Some(d_v),
d_out,
"hvp_into_device",
)
}
#[cfg(target_os = "linux")]
pub(crate) fn launch_bms_flex_row_hvp(
storage: &DeviceResidentRowHess,
v: &[f64],
) -> Result<Vec<f64>, GpuError> {
launch_bms_flex_row_hvp_multi(storage, v, 1)
}
#[cfg(target_os = "linux")]
pub(crate) fn launch_bms_flex_row_diagonal(
storage: &DeviceResidentRowHess,
) -> Result<Vec<f64>, GpuError> {
launch_bms_flex_row_diagonal_host(storage)
}
#[cfg(target_os = "linux")]
pub(crate) const DENSE_BLOCK_MAX_P: usize = 72;
#[cfg(target_os = "linux")]
pub(crate) const DENSE_BLOCK_ROWS_PER_CTA: u32 = 32;
#[cfg(target_os = "linux")]
pub(crate) fn launch_bms_flex_row_dense(
storage: &DeviceResidentRowHess,
) -> Result<Vec<f64>, GpuError> {
let p_total = storage.block.p_total;
if p_total <= DENSE_BLOCK_MAX_P {
return launch_bms_flex_row_dense_block(storage);
}
materialize_dense_from_hvp_batches(p_total, |basis, rhs_count| {
launch_bms_flex_row_hvp_multi(storage, basis, rhs_count)
})
}
#[cfg(target_os = "linux")]
pub fn launch_bms_flex_row_dense_block(
storage: &DeviceResidentRowHess,
) -> Result<Vec<f64>, GpuError> {
let p_total = storage.block.p_total;
if p_total == 0 {
return Err(GpuError::DriverCallFailed {
reason: "bms_flex_row dense_block: p_total must be > 0".to_string(),
});
}
if p_total > DENSE_BLOCK_MAX_P {
return Err(GpuError::DriverCallFailed {
reason: format!(
"bms_flex_row dense_block: p_total={p_total} exceeds DENSE_BLOCK_MAX_P={DENSE_BLOCK_MAX_P} \
(per-CTA shmem accumulator p²*8 bytes would exceed V100's 48 KiB/block)"
),
});
}
let backend = HvpKernelBackend::probe()?;
let stream = backend.stream.clone();
let args = PreparedBmsFlexRowLaunchArgs::from_storage(storage)?;
let n = storage.n;
let rows_per_cta = DENSE_BLOCK_ROWS_PER_CTA as usize;
let num_chunks = n.div_ceil(rows_per_cta);
let pp = checked_shape_len("dense_block [p_total,p_total]", &[p_total, p_total])?;
let partial_len = checked_shape_len("dense_block partial", &[num_chunks, pp])?;
let mut d_partial =
stream
.alloc_zeros::<f64>(partial_len)
.map_err(|err| GpuError::DriverCallFailed {
reason: format!("bms_flex_row dense_block alloc partial: {err}"),
})?;
let mut d_out = stream
.alloc_zeros::<f64>(pp)
.map_err(|err| GpuError::DriverCallFailed {
reason: format!("bms_flex_row dense_block alloc out: {err}"),
})?;
let part_func = backend
.module
.load_function("bms_flex_row_dense_block_partial")
.map_err(|err| GpuError::DriverCallFailed {
reason: format!("bms_flex_row dense_block load partial: {err}"),
})?;
let red_func = backend
.module
.load_function("bms_flex_row_dense_block_reduce")
.map_err(|err| GpuError::DriverCallFailed {
reason: format!("bms_flex_row dense_block load reduce: {err}"),
})?;
let rows_per_cta_i32 = i32::try_from(DENSE_BLOCK_ROWS_PER_CTA).map_err(|_| {
GpuError::DriverCallFailed {
reason: format!(
"bms_flex_row dense_block: rows_per_cta={DENSE_BLOCK_ROWS_PER_CTA} exceeds i32 range"
),
}
})?;
let num_chunks_u32 = u32::try_from(num_chunks).map_err(|_| GpuError::DriverCallFailed {
reason: format!("bms_flex_row dense_block: num_chunks={num_chunks} exceeds u32 range"),
})?;
let num_chunks_i32 = i32::try_from(num_chunks).map_err(|_| GpuError::DriverCallFailed {
reason: format!("bms_flex_row dense_block: num_chunks={num_chunks} exceeds i32 range"),
})?;
let pp_u32 = u32::try_from(pp).map_err(|_| GpuError::DriverCallFailed {
reason: format!("bms_flex_row dense_block: p_total²={pp} exceeds u32 range"),
})?;
let shmem_bytes_usize =
pp.checked_mul(std::mem::size_of::<f64>())
.ok_or_else(|| GpuError::DriverCallFailed {
reason: format!("dense_block shmem bytes overflow for p_total={p_total}"),
})?;
let shmem_bytes: u32 =
u32::try_from(shmem_bytes_usize).map_err(|_| GpuError::DriverCallFailed {
reason: format!("dense_block shmem bytes overflow u32 for p_total={p_total}"),
})?;
let cfg_part = LaunchConfig {
grid_dim: (num_chunks_u32, 1, 1),
block_dim: (HVP_THREADS, 1, 1),
shared_mem_bytes: shmem_bytes,
};
let mut builder = stream.launch_builder(&part_func);
builder
.arg(&args.n_i32)
.arg(&args.r_i32)
.arg(&args.p_m_i32)
.arg(&args.p_g_i32)
.arg(&args.p_total_i32)
.arg(&args.h_block_start)
.arg(&args.h_block_len)
.arg(&args.w_block_start)
.arg(&args.w_block_len)
.arg(&args.h_primary_start)
.arg(&args.w_primary_start)
.arg(&rows_per_cta_i32)
.arg(&storage.hess)
.arg(&storage.marginal_design)
.arg(&storage.logslope_design)
.arg(&mut d_partial);
unsafe { builder.launch(cfg_part) }.map_err(|err| GpuError::DriverCallFailed {
reason: format!("bms_flex_row dense_block partial launch: {err}"),
})?;
let red_threads: u32 = REDUCTION_THREADS;
let red_blocks = pp_u32.div_ceil(red_threads);
let cfg_red = LaunchConfig {
grid_dim: (red_blocks, 1, 1),
block_dim: (red_threads, 1, 1),
shared_mem_bytes: 0,
};
let mut builder = stream.launch_builder(&red_func);
builder
.arg(&num_chunks_i32)
.arg(&args.p_total_i32)
.arg(&d_partial)
.arg(&mut d_out);
unsafe { builder.launch(cfg_red) }.map_err(|err| GpuError::DriverCallFailed {
reason: format!("bms_flex_row dense_block reduce launch: {err}"),
})?;
stream
.synchronize()
.map_err(|err| GpuError::DriverCallFailed {
reason: format!("bms_flex_row dense_block sync: {err}"),
})?;
stream
.clone_dtoh(&d_out)
.map_err(|err| GpuError::DriverCallFailed {
reason: format!("bms_flex_row dense_block download: {err}"),
})
}
#[cfg(test)]
mod row_kernel_tests {
pub(crate) fn host_log_ndtr_and_mills(x: f64) -> (f64, f64) {
gam_gpu::numerics_host::log_ndtr_and_mills(x)
}
#[cfg(target_os = "linux")]
pub(crate) fn host_log_ndtr_mills_curvature(x: f64) -> (f64, f64, f64) {
gam_gpu::numerics_host::log_ndtr_mills_curvature(x)
}
pub(crate) mod parity_415 {
use crate::bms::family::*;
use crate::bms::hessian_paths::*;
use crate::bms::{DeviationBlockConfig, LatentMeasureKind, exact_kernel};
use gam_linalg::matrix::{DenseDesignMatrix, DesignMatrix};
use gam_problem::{InverseLink, ParameterBlockState, StandardLink};
use ndarray::{Array1, Array2};
use std::sync::{Arc, Mutex};
pub(crate) fn make_flex_parity_family(
n: usize,
score_internal_knots: usize,
link_internal_knots: usize,
) -> (BernoulliMarginalSlopeFamily, Vec<ParameterBlockState>) {
let score_seed = Array1::linspace(-2.0, 2.0, n.max(6));
let link_seed = Array1::linspace(-1.8, 1.8, n.max(6));
let score_cfg = DeviationBlockConfig {
num_internal_knots: score_internal_knots,
..DeviationBlockConfig::default()
};
let link_cfg = DeviationBlockConfig {
num_internal_knots: link_internal_knots,
..DeviationBlockConfig::default()
};
let score_prepared =
build_score_warp_deviation_block_from_seed(&score_seed, &score_cfg)
.expect("build score warp block");
let link_prepared = build_link_deviation_block_from_knots_design_seed_and_weights(
&link_seed, &link_seed, &link_cfg,
)
.expect("build link deviation block");
let y: Array1<f64> =
Array1::from_iter((0..n).map(|i| if (i * 17 + 3) % 7 >= 4 { 1.0 } else { 0.0 }));
let weights: Array1<f64> =
Array1::from_iter((0..n).map(|i| 0.75 + ((i * 11 + 5) % 5) as f64 * 0.05));
let z: Array1<f64> =
Array1::from_iter((0..n).map(|i| -1.7 + 3.4 * (i as f64 + 0.5) / n as f64));
let marginal_x = Array2::from_shape_fn((n, 2), |(i, j)| {
if j == 0 {
1.0
} else {
-0.4 + 0.8 * ((i * 19 + 7) % n) as f64 / n as f64
}
});
let logslope_x = Array2::from_shape_fn((n, 2), |(i, j)| {
if j == 0 {
1.0
} else {
0.3 - 0.6 * ((i * 23 + 11) % n) as f64 / n as f64
}
});
let family = BernoulliMarginalSlopeFamily {
y: Arc::new(y),
weights: Arc::new(weights),
z: Arc::new(z.clone()),
latent_measure: LatentMeasureKind::StandardNormal,
gaussian_frailty_sd: Some(0.15),
base_link: InverseLink::Standard(StandardLink::Probit),
marginal_design: DesignMatrix::Dense(DenseDesignMatrix::from(marginal_x.clone())),
logslope_design: DesignMatrix::Dense(DenseDesignMatrix::from(logslope_x.clone())),
score_warp: Some(score_prepared.runtime.clone()),
link_dev: Some(link_prepared.runtime.clone()),
policy: gam_runtime::resource::ResourcePolicy::default_library(),
cell_moment_lru: Arc::new(exact_kernel::CellMomentLruCache::new(1024)),
cell_moment_cache_stats: Arc::new(exact_kernel::CellMomentCacheStats::default()),
intercept_warm_starts: None,
auto_subsample_phase_counter: Arc::new(std::sync::atomic::AtomicUsize::new(0)),
auto_subsample_last_rho: Arc::new(Mutex::new(None)),
};
let beta_m = Array1::from_vec(vec![0.12, -0.04]);
let beta_g = Array1::from_vec(vec![0.35, 0.03]);
let beta_h = Array1::from_iter(
(0..score_prepared.runtime.basis_dim()).map(|idx| 0.0015 * (idx as f64 + 1.0)),
);
let beta_w = Array1::from_iter(
(0..link_prepared.runtime.basis_dim()).map(|idx| -0.001 * (idx as f64 + 1.0)),
);
let states = vec![
ParameterBlockState {
eta: marginal_x.dot(&beta_m),
beta: beta_m,
},
ParameterBlockState {
eta: logslope_x.dot(&beta_g),
beta: beta_g,
},
ParameterBlockState {
beta: beta_h,
eta: Array1::zeros(z.len()),
},
ParameterBlockState {
beta: beta_w,
eta: Array1::zeros(z.len()),
},
];
(family, states)
}
fn assert_generated_cuda_row_kernel_matches_canonical_cpu_lowering(
n: usize,
score_internal_knots: usize,
link_internal_knots: usize,
expected_r: Option<usize>,
) {
let (family, states) =
make_flex_parity_family(n, score_internal_knots, link_internal_knots);
let cache = family
.build_exact_eval_cache(&states)
.expect("flex exact eval cache");
assert!(
cache.row_cell_moments.is_some(),
"#415 fixture must materialise production row-cell moments"
);
let primary = &cache.primary;
let r = primary.total;
let p_h = primary.h.as_ref().map(|range| range.len()).unwrap_or(0);
let p_w = primary.w.as_ref().map(|range| range.len()).unwrap_or(0);
assert!(
p_h > 0 && p_w > 0,
"fixture must activate both deviation blocks"
);
assert_eq!(r, 2 + p_h + p_w);
if let Some(expected_r) = expected_r {
assert_eq!(
r, expected_r,
"fixture knot counts must exercise the requested primary width"
);
}
let owned = family
.pack_bms_flex_row_kernel_inputs(&states, &cache)
.expect("packing production CUDA inputs must not error")
.expect("StandardNormal full-FLEX fixture must admit the CUDA row kernel");
let inputs = owned.as_borrowed();
let mut canonical_neglog = vec![0.0; n];
let mut canonical_grad = vec![0.0; n * r];
let mut canonical_hess = vec![0.0; n * r * r];
let mut scratch = BernoulliMarginalSlopeFlexRowScratch::new(r);
let mut checked_labels = [false, false];
for row in 0..n {
let row_ctx = BernoulliMarginalSlopeFamily::row_ctx(&cache, row);
let row_moments = cache
.row_cell_moments
.as_ref()
.and_then(|bundle| bundle.row(row, 9));
assert!(
row_moments.is_some(),
"row {row} must carry degree-9 moments"
);
canonical_neglog[row] = family
.lower_bms_flex_row_order2_with_moments(
row,
&states,
primary,
row_ctx,
row_moments,
cache.cell_family_forest.as_ref(),
true,
&mut scratch,
)
.expect("canonical production CPU row lowering");
for u in 0..r {
canonical_grad[row * r + u] = scratch.grad[u];
for v in 0..r {
let value = scratch.hess[[u, v]];
assert!(value.is_finite(), "row {row}: H[{u},{v}] is non-finite");
assert_eq!(
value.to_bits(),
scratch.hess[[v, u]].to_bits(),
"row {row}: canonical Hessian lost exact symmetry"
);
canonical_hess[row * r * r + u * r + v] = value;
}
}
checked_labels[family.y[row] as usize] = true;
}
assert!(checked_labels[0] && checked_labels[1]);
let mut separates_value_from_q_derivative = false;
for row in 0..n {
let sign = 2.0 * inputs.y[row] - 1.0;
let (_, lambda) = super::host_log_ndtr_and_mills(sign * inputs.e_obs[row]);
let scale = -inputs.w[row] * sign * lambda;
if scale.abs() > 1e-12 {
let observed_q_derivative = canonical_grad[row * r] / scale;
if (observed_q_derivative - inputs.e_obs[row]).abs() > 1e-8 {
separates_value_from_q_derivative = true;
break;
}
}
}
assert!(
separates_value_from_q_derivative,
"fixture must distinguish the observed value from its q derivative"
);
#[cfg(not(target_os = "linux"))]
{
eprintln!("[bms_flex_row parity] generated CUDA check requires Linux");
return;
}
#[cfg(target_os = "linux")]
{
match gam_gpu::device_runtime::GpuRuntime::resolve(gam_gpu::GpuPolicy::Auto) {
Ok(Some(_)) => {}
Ok(None) => {
eprintln!("[bms_flex_row parity] no CUDA device");
return;
}
Err(error) => panic!("[bms_flex_row parity] CUDA probe failed: {error}"),
}
let gpu = super::super::launch_bms_flex_row_kernel(owned.as_borrowed())
.expect("CUDA-selected canonical parity launch must succeed");
let check = |channel: &str, index: usize, cpu: f64, device: f64| {
let difference = (cpu - device).abs();
let tolerance = 1e-8 + 1e-8 * cpu.abs();
assert!(
difference <= tolerance,
"{channel}[{index}] CPU={cpu:.17e} CUDA={device:.17e} \
difference={difference:.3e} tolerance={tolerance:.3e}"
);
};
for (index, (&cpu, &device)) in
canonical_neglog.iter().zip(gpu.neglog.iter()).enumerate()
{
check("neglog", index, cpu, device);
}
for (index, (&cpu, &device)) in
canonical_grad.iter().zip(gpu.grad.iter()).enumerate()
{
check("gradient", index, cpu, device);
}
for (index, (&cpu, &device)) in
canonical_hess.iter().zip(gpu.hess.iter()).enumerate()
{
check("hessian", index, cpu, device);
}
}
}
#[test]
fn generated_cuda_row_kernel_matches_canonical_cpu_lowering_415() {
assert_generated_cuda_row_kernel_matches_canonical_cpu_lowering(12, 3, 3, None);
}
#[test]
fn full_flex_canonical_exact_cache_admits_material_finite_cell_curvature_2321() {
let (family, states) = make_flex_parity_family(256, 8, 6);
let cache = family
.build_exact_eval_cache(&states)
.expect("the full-FLEX host cache must preserve non-affine finite cells");
let score_width = cache
.primary
.h
.as_ref()
.expect("the full-FLEX fixture must retain its score-warp block")
.len();
let deviation_width = cache
.primary
.w
.as_ref()
.expect("the full-FLEX fixture must retain its link-deviation block")
.len();
assert!(score_width > 0 && deviation_width > 0);
assert_eq!(
cache.primary.total,
2 + score_width + deviation_width,
"the canonical primary layout must contain exactly q, logslope, score-warp, and link-deviation coordinates"
);
assert!(
cache.row_cell_moments.is_some(),
"the production full-FLEX fixture must materialize its exact row-cell cache"
);
}
#[test]
fn generated_cuda_row_kernel_r33_matches_canonical_cpu_lowering_932() {
gam_gpu::configure_global_policy(gam_gpu::GpuPolicy::Required);
assert_eq!(
gam_gpu::global_policy(),
gam_gpu::GpuPolicy::Required,
"fresh-process r=33 parity must claim Required before runtime discovery"
);
gam_gpu::device_runtime::GpuRuntime::require()
.expect("#932 mandatory r=33 CUDA runtime");
assert_generated_cuda_row_kernel_matches_canonical_cpu_lowering(40, 15, 14, Some(33));
}
}
}
#[cfg(all(test, target_os = "linux"))]
mod tests {
use super::row_kernel_tests::*;
use super::*;
use crate::bms::exact_eval_cache::RowPrimaryEvalCache;
use crate::bms::row_kernel::BernoulliMarginalSlopeExactNewtonJointHessianWorkspace;
use crate::custom_family::{BlockwiseFitOptions, ExactNewtonJointHessianWorkspace};
use gam_gpu::{GpuPolicy, configure_global_policy};
use ndarray::{Array1, Array2};
use std::hint::black_box;
use std::sync::atomic::AtomicUsize;
use std::time::{Duration, Instant};
#[cfg(target_os = "linux")]
fn assert_row_batch_dispatch_worthy(
label: &str,
policy: &gam_gpu::policy::GpuDispatchPolicy,
n: usize,
) {
assert!(
policy.row_batch_target_is_gpu(n),
"{label}: n={n} rows is below this device's calibrated row-kernel \
crossover ({}), so the fixture no longer exercises a shape the \
dispatch policy would send to the device — grow the fixture rather \
than lowering the crossover",
policy.row_kernel_min_n
);
assert!(
!policy.row_batch_target_is_gpu(0),
"{label}: the dispatch predicate admitted an empty batch, so the \
assertion above proves nothing about n={n}"
);
}
fn cuda_runtime_for_test(
test_name: &str,
) -> Option<&'static gam_gpu::device_runtime::GpuRuntime> {
match gam_gpu::device_runtime::GpuRuntime::resolve(GpuPolicy::Auto) {
Ok(Some(runtime)) => Some(runtime),
Ok(None) => {
eprintln!("[{test_name}] no CUDA device — skipping");
None
}
Err(error) => panic!("[{test_name}] CUDA probe failed: {error}"),
}
}
fn assert_array1_close_932(label: &str, expected: &Array1<f64>, actual: &Array1<f64>) {
assert_eq!(expected.len(), actual.len(), "{label}: length mismatch");
for (index, (&want, &got)) in expected.iter().zip(actual).enumerate() {
let tolerance = 2.0e-8 * (1.0 + want.abs());
assert!(
want.is_finite() && got.is_finite() && (want - got).abs() <= tolerance,
"{label}[{index}]: expected={want:.17e} actual={got:.17e} tolerance={tolerance:.3e}"
);
}
}
#[test]
fn mandatory_required_gpu_workspace_consumes_device_cache_end_to_end_932() {
configure_global_policy(GpuPolicy::Required);
assert_eq!(
gam_gpu::global_policy(),
GpuPolicy::Required,
"fresh-process acceptance test must claim Required before any competing policy"
);
gam_gpu::device_runtime::GpuRuntime::require().expect("#932 mandatory CUDA runtime");
let (family, states) = row_kernel_tests::parity_415::make_flex_parity_family(256, 8, 6);
let mut workspace = BernoulliMarginalSlopeExactNewtonJointHessianWorkspace::new(
family,
states,
BlockwiseFitOptions::default(),
)
.expect("#932 Required workspace must build its device row cache");
assert!(
matches!(
&workspace.cache.row_primary_hessians,
RowPrimaryEvalCache::Device(_)
),
"Required full-FLEX workspace must retain RowPrimaryEvalCache::Device"
);
{
let device = workspace
.cache
.row_primary_hessians
.device()
.expect("device cache variant");
assert!(
device
.primary
.h
.as_ref()
.is_some_and(|range| !range.is_empty())
&& device
.primary
.w
.as_ref()
.is_some_and(|range| !range.is_empty()),
"mandatory fixture must carry active h and w primary blocks"
);
assert!(
device
.block
.h
.as_ref()
.is_some_and(|range| !range.is_empty())
&& device
.block
.w
.as_ref()
.is_some_and(|range| !range.is_empty()),
"mandatory fixture must carry active h and w coefficient blocks"
);
}
for operation in ["host HVP replay", "host diagonal replay"] {
let error = workspace
.cache
.row_primary_hessians
.reject_device_cpu_recompute(operation)
.expect_err("a selected device cache must reject host row recomputation");
assert!(
error.contains("device-resident row evaluation selected")
&& error.contains("CPU row recomputation is forbidden"),
"unexpected fail-closed diagnostic: {error}"
);
}
let total = workspace.cache.slices.total;
let direction = Array1::from_shape_fn(total, |index| {
let sign = if index % 2 == 0 { 1.0 } else { -1.0 };
sign * (0.025 + 0.0075 * index as f64)
});
let joint_ll = workspace
.joint_log_likelihood_evaluation()
.expect("device joint log-likelihood")
.expect("device joint log-likelihood must be present");
let joint = workspace
.joint_gradient_evaluation()
.expect("device joint gradient")
.expect("device joint gradient must be present");
assert!(joint_ll.is_finite());
assert_eq!(joint.log_likelihood.to_bits(), joint_ll.to_bits());
assert_eq!(joint.gradient.len(), total);
assert!(joint.gradient.iter().all(|value| value.is_finite()));
let hvp = workspace
.hessian_matvec(&direction)
.expect("device HVP")
.expect("device HVP must be present");
let mut hvp_into = Array1::from_elem(total, f64::NAN);
assert!(
workspace
.hessian_matvec_into(&direction, &mut hvp_into)
.expect("device HVP-into"),
"device HVP-into must report that it handled the direction"
);
assert_array1_close_932("HVP owned/into", &hvp, &hvp_into);
let rhs = Array2::from_shape_fn((total, 3), |(row, column)| {
(row as f64 + 1.0)
* (column as f64 + 0.5)
* 0.011
* if (row + column) % 3 == 0 { -1.0 } else { 1.0 }
});
let mut applied = Array2::<f64>::from_elem((total, rhs.ncols()), f64::NAN);
assert!(
workspace
.hessian_apply_mat(&rhs, &mut applied)
.expect("device multi-RHS apply"),
"device multi-RHS apply must report that it handled the matrix"
);
let diagonal = workspace
.hessian_diagonal()
.expect("device diagonal")
.expect("device diagonal must be present");
let dense = workspace
.hessian_dense_forced()
.expect("device forced dense Hessian")
.expect("device forced dense Hessian must be present");
assert_eq!(dense.dim(), (total, total));
assert_array1_close_932("dense * v / HVP", &dense.dot(&direction), &hvp);
assert_array1_close_932("dense diagonal", &dense.diag().to_owned(), &diagonal);
let dense_applied = dense.dot(&rhs);
for column in 0..rhs.ncols() {
assert_array1_close_932(
&format!("dense * V / apply_mat column {column}"),
&dense_applied.column(column).to_owned(),
&applied.column(column).to_owned(),
);
}
for state in &mut workspace.block_states {
state.beta.fill(f64::NAN);
state.eta.fill(f64::NAN);
}
let poisoned_hvp = workspace
.hessian_matvec(&direction)
.expect("device HVP after host-state poison")
.expect("device HVP after host-state poison must be present");
let poisoned_diagonal = workspace
.hessian_diagonal()
.expect("device diagonal after host-state poison")
.expect("device diagonal after host-state poison must be present");
assert_eq!(
hvp.as_slice(),
poisoned_hvp.as_slice(),
"fixed-order device HVP changed after poisoning host block state"
);
assert_eq!(
diagonal.as_slice(),
poisoned_diagonal.as_slice(),
"fixed-order device diagonal changed after poisoning host block state"
);
}
#[test]
fn release_measure_generated_bms_full_row_vs_strongest_cpu_932() {
const N: usize = 32_768;
const WARMUPS: usize = 3;
const SAMPLES: usize = 21;
configure_global_policy(GpuPolicy::Required);
assert_eq!(gam_gpu::global_policy(), GpuPolicy::Required);
gam_gpu::device_runtime::GpuRuntime::require()
.expect("#932 full-row release measurement requires CUDA");
let (family, states) = row_kernel_tests::parity_415::make_flex_parity_family(N, 9, 7);
let cache = family
.build_exact_eval_cache(&states)
.expect("full-row timing exact cache");
let r = cache.primary.total;
assert_eq!(r, 20, "9/7 knot fixture must expose primary width r=20");
let marginal = family
.marginal_design
.as_dense_ref()
.expect("timing fixture marginal design must be dense");
let logslope = family
.logslope_design
.as_dense_ref()
.expect("timing fixture logslope design must be dense");
assert!(marginal.is_standard_layout() && logslope.is_standard_layout());
let marginal_slice = marginal
.as_slice()
.expect("timing fixture marginal design is contiguous");
let logslope_slice = logslope
.as_slice()
.expect("timing fixture logslope design is contiguous");
let block = BmsFlexBlockLayout {
p_m: cache.slices.marginal.len(),
p_g: cache.slices.logslope.len(),
h: cache.slices.h.clone(),
w: cache.slices.w.clone(),
p_total: cache.slices.total,
};
let primary = BmsFlexPrimaryLayout {
h: cache.primary.h.clone(),
w: cache.primary.w.clone(),
r,
};
assert!(
primary.h.as_ref().is_some_and(|range| !range.is_empty())
&& primary.w.as_ref().is_some_and(|range| !range.is_empty()),
"full-row timing fixture must exercise both h and w"
);
let pin_bytes =
crate::bms::family::BernoulliMarginalSlopeFamily::row_primary_eval_tile_bytes(N, r);
let run_cpu = || {
let completed = AtomicUsize::new(0);
family
.build_row_primary_hessian_pin(
&states,
&cache,
0..N,
&completed,
N.saturating_add(1),
Instant::now(),
pin_bytes,
)
.expect("production Rayon row-primary batch")
};
let run_gpu = || {
let owned = family
.pack_bms_flex_row_kernel_inputs(&states, &cache)
.expect("production BMS GPU packing")
.expect("StandardNormal full-FLEX timing fixture must pack");
launch_bms_flex_row_kernel_device_resident(
owned.as_borrowed(),
marginal_slice,
logslope_slice,
block.clone(),
primary.clone(),
)
.expect("production device-resident row launch")
};
let measure_cpu = || {
let started = Instant::now();
let output = black_box(run_cpu());
(started.elapsed(), output)
};
let measure_gpu = || {
let started = Instant::now();
let output = black_box(run_gpu());
(started.elapsed(), output)
};
let cold_started = Instant::now();
let cold_gpu = black_box(run_gpu());
let cold_gpu_e2e_nvrtc = cold_started.elapsed();
drop(cold_gpu);
for _ in 0..WARMUPS {
black_box(run_cpu());
black_box(run_gpu());
}
let mut cpu_samples = Vec::<Duration>::with_capacity(SAMPLES);
let mut gpu_samples = Vec::<Duration>::with_capacity(SAMPLES);
let mut last_cpu = None;
let mut last_gpu = None;
for sample in 0..SAMPLES {
if sample % 2 == 0 {
let (cpu_elapsed, cpu) = measure_cpu();
cpu_samples.push(cpu_elapsed);
let (gpu_elapsed, gpu) = measure_gpu();
gpu_samples.push(gpu_elapsed);
if sample + 1 == SAMPLES {
last_cpu = Some(cpu);
last_gpu = Some(gpu);
}
} else {
let (gpu_elapsed, gpu) = measure_gpu();
gpu_samples.push(gpu_elapsed);
let (cpu_elapsed, cpu) = measure_cpu();
cpu_samples.push(cpu_elapsed);
drop(gpu);
drop(cpu);
}
}
let cpu = last_cpu.expect("final CPU sample retained for parity");
let gpu = last_gpu.expect("final GPU sample retained for parity");
let stream = HvpKernelBackend::probe()
.expect("HVP backend remains available")
.stream
.clone();
let gpu_neglog = stream
.clone_dtoh(&gpu.neglog)
.expect("download timed GPU neglog for parity");
let gpu_grad = stream
.clone_dtoh(&gpu.grad)
.expect("download timed GPU gradient for parity");
let gpu_hess = stream
.clone_dtoh(&gpu.hess)
.expect("download timed GPU Hessian for parity");
let cpu_channels = [
cpu.neglog().as_slice().expect("CPU neglog is contiguous"),
cpu.grad().as_slice().expect("CPU gradient is contiguous"),
cpu.hess().as_slice().expect("CPU Hessian is contiguous"),
];
let gpu_channels = [
gpu_neglog.as_slice(),
gpu_grad.as_slice(),
gpu_hess.as_slice(),
];
let mut nonfinite = 0_usize;
let mut max_abs = 0.0_f64;
let mut max_scaled = 0.0_f64;
let mut cpu_digest = 0.0_f64;
let mut gpu_digest = 0.0_f64;
let mut digest_index = 0_usize;
for (cpu_channel, gpu_channel) in cpu_channels.iter().zip(gpu_channels) {
assert_eq!(cpu_channel.len(), gpu_channel.len());
for (&host, &device) in cpu_channel.iter().zip(gpu_channel) {
if !host.is_finite() || !device.is_finite() {
nonfinite += 1;
}
let difference = (host - device).abs();
let tolerance = 1.0e-8 * (1.0 + host.abs());
max_abs = max_abs.max(difference);
max_scaled = max_scaled.max(difference / tolerance);
let weight = 1.0 + (digest_index % 251) as f64 / 251.0;
cpu_digest += weight * host;
gpu_digest += weight * device;
digest_index += 1;
}
}
assert_eq!(
nonfinite, 0,
"full-row CPU/GPU output contains non-finite values"
);
assert!(
max_scaled <= 1.0,
"full-row CPU/GPU parity exceeded tolerance: max_abs={max_abs:.3e} max_scaled={max_scaled:.3e}"
);
let mut cpu_ms = cpu_samples
.iter()
.map(|sample| sample.as_secs_f64() * 1.0e3)
.collect::<Vec<_>>();
let mut gpu_ms = gpu_samples
.iter()
.map(|sample| sample.as_secs_f64() * 1.0e3)
.collect::<Vec<_>>();
cpu_ms.sort_by(f64::total_cmp);
gpu_ms.sort_by(f64::total_cmp);
let p25 = SAMPLES / 4;
let p50 = SAMPLES / 2;
let p75 = 3 * SAMPLES / 4;
let conservative_speedup = cpu_ms[p25] / gpu_ms[p75];
let median_speedup = cpu_ms[p50] / gpu_ms[p50];
let cpu_distribution = cpu_ms
.iter()
.map(|value| format!("{value:.6}"))
.collect::<Vec<_>>()
.join(",");
let gpu_distribution = gpu_ms
.iter()
.map(|value| format!("{value:.6}"))
.collect::<Vec<_>>()
.join(",");
println!(
"G932_BMS_FULL_ROW n={N} r={r} warmups={WARMUPS} samples={SAMPLES} \
cold_gpu_e2e_nvrtc_ms={:.6} cpu_ms_p25={:.6} cpu_ms_p50={:.6} cpu_ms_p75={:.6} \
gpu_ms_p25={:.6} gpu_ms_p50={:.6} gpu_ms_p75={:.6} \
speedup_conservative_cpu_p25_over_gpu_p75={conservative_speedup:.6} \
speedup_median={median_speedup:.6} parity_max_abs={max_abs:.9e} \
parity_max_scaled={max_scaled:.9e} cpu_digest={cpu_digest:.17e} \
gpu_digest={gpu_digest:.17e} nonfinite={nonfinite} \
cpu_ms_sorted=[{cpu_distribution}] gpu_ms_sorted=[{gpu_distribution}]",
cold_gpu_e2e_nvrtc.as_secs_f64() * 1.0e3,
cpu_ms[p25],
cpu_ms[p50],
cpu_ms[p75],
gpu_ms[p25],
gpu_ms[p50],
gpu_ms[p75],
);
}
#[test]
fn dense_hvp_batches_transpose_column_images_in_bounded_groups_932() {
let p_total = 2 * BMS_FLEX_ROW_HVP_MAX_RHS + 3;
let matrix = (0..p_total * p_total)
.map(|index| {
let row = index / p_total;
let column = index % p_total;
1000.0 * row as f64 + column as f64 + 0.25
})
.collect::<Vec<_>>();
let mut observed_batch_sizes = Vec::new();
let dense = materialize_dense_from_hvp_batches(p_total, |basis, rhs_count| {
observed_batch_sizes.push(rhs_count);
let mut images = vec![0.0_f64; rhs_count * p_total];
for rhs in 0..rhs_count {
for row in 0..p_total {
images[rhs * p_total + row] = (0..p_total)
.map(|column| {
matrix[row * p_total + column] * basis[rhs * p_total + column]
})
.sum();
}
}
Ok(images)
})
.expect("synthetic H*I batches must materialize");
assert_eq!(dense, matrix);
assert_eq!(
observed_batch_sizes,
vec![BMS_FLEX_ROW_HVP_MAX_RHS, BMS_FLEX_ROW_HVP_MAX_RHS, 3]
);
}
pub(crate) fn minimal_inputs<'a>(buffers: &'a TestBuffers) -> BmsFlexRowKernelInputs<'a> {
BmsFlexRowKernelInputs {
n_rows: 1,
r: 4,
p_h: 1,
p_w: 1,
q: &buffers.q,
b: &buffers.b,
mu_1: &buffers.mu_1,
mu_2: &buffers.mu_2,
z_obs: &buffers.z_obs,
y: &buffers.y,
w: &buffers.w,
e_obs: &buffers.e_obs,
s_f: 1.0,
cell_offsets: &buffers.cell_offsets,
cell_c0: &buffers.cell_c0,
cell_c1: &buffers.cell_c1,
cell_c2: &buffers.cell_c2,
cell_c3: &buffers.cell_c3,
cell_a: &buffers.cell_a,
cell_aa: &buffers.cell_aa,
cell_r: &buffers.cell_r,
cell_ar: &buffers.cell_ar,
cell_sbb: &buffers.cell_sbb,
cell_sbh: &buffers.cell_sbh,
cell_sbw: &buffers.cell_sbw,
cell_moments: CellMomentsSource::Host(&buffers.cell_moments),
chi_obs: &buffers.chi_obs,
xi_obs: &buffers.xi_obs,
rho_u: &buffers.rho_u,
tau_u: &buffers.tau_u,
r_uv: &buffers.r_uv,
}
}
pub(crate) struct TestBuffers {
pub(crate) q: Vec<f64>,
pub(crate) b: Vec<f64>,
pub(crate) mu_1: Vec<f64>,
pub(crate) mu_2: Vec<f64>,
pub(crate) z_obs: Vec<f64>,
pub(crate) y: Vec<f64>,
pub(crate) w: Vec<f64>,
pub(crate) e_obs: Vec<f64>,
pub(crate) cell_offsets: Vec<u32>,
pub(crate) cell_c0: Vec<f64>,
pub(crate) cell_c1: Vec<f64>,
pub(crate) cell_c2: Vec<f64>,
pub(crate) cell_c3: Vec<f64>,
pub(crate) cell_a: Vec<f64>,
pub(crate) cell_aa: Vec<f64>,
pub(crate) cell_r: Vec<f64>,
pub(crate) cell_ar: Vec<f64>,
pub(crate) cell_sbb: Vec<f64>,
pub(crate) cell_sbh: Vec<f64>,
pub(crate) cell_sbw: Vec<f64>,
pub(crate) cell_moments: Vec<f64>,
pub(crate) chi_obs: Vec<f64>,
pub(crate) xi_obs: Vec<f64>,
pub(crate) rho_u: Vec<f64>,
pub(crate) tau_u: Vec<f64>,
pub(crate) r_uv: Vec<f64>,
}
pub(crate) fn make_buffers(n_cells: u32, r: usize, p_h: usize, p_w: usize) -> TestBuffers {
let cells = n_cells as usize;
TestBuffers {
q: vec![0.1; 1],
b: vec![0.5; 1],
mu_1: vec![0.3; 1],
mu_2: vec![0.07; 1],
z_obs: vec![0.0; 1],
y: vec![1.0; 1],
w: vec![1.0; 1],
e_obs: vec![0.15; 1],
cell_offsets: vec![0, n_cells],
cell_c0: vec![0.2; cells],
cell_c1: vec![-0.1; cells],
cell_c2: vec![0.05; cells],
cell_c3: vec![-0.02; cells],
cell_a: vec![0.1; cells * 4],
cell_aa: vec![0.0; cells * 4],
cell_r: vec![0.05; cells * (r - 1) * 4],
cell_ar: vec![0.0; cells * (r - 1) * 4],
cell_sbb: vec![0.0; cells * 4],
cell_sbh: vec![0.0; cells * p_h * 4],
cell_sbw: vec![0.0; cells * p_w * 4],
cell_moments: vec![1.0; cells * MOMENT_STRIDE],
chi_obs: vec![1.0; 1],
xi_obs: vec![0.0; 1],
rho_u: vec![0.0; r],
tau_u: vec![0.0; r],
r_uv: vec![0.0; r * r],
}
}
#[test]
pub(crate) fn validate_accepts_minimal_inputs() {
let buffers = make_buffers(2, 4, 1, 1);
let inputs = minimal_inputs(&buffers);
assert!(inputs.validate().is_ok());
}
#[test]
pub(crate) fn validate_accepts_r33_with_active_h_and_w_blocks() {
let r = 33;
let p_h = 16;
let p_w = 15;
let buffers = make_buffers(1, r, p_h, p_w);
let inputs = BmsFlexRowKernelInputs {
r,
p_h,
p_w,
rho_u: &buffers.rho_u,
tau_u: &buffers.tau_u,
r_uv: &buffers.r_uv,
cell_r: &buffers.cell_r,
cell_ar: &buffers.cell_ar,
cell_sbh: &buffers.cell_sbh,
cell_sbw: &buffers.cell_sbw,
..minimal_inputs(&buffers)
};
inputs
.validate()
.expect("r=33 is a valid checked shape, not a semantic width boundary");
}
#[test]
pub(crate) fn checked_shape_len_rejects_arithmetic_overflow() {
let err = checked_shape_len("overflow test", &[usize::MAX, 2])
.expect_err("shape multiplication must fail closed");
assert!(err.to_string().contains("shape product overflow"));
}
#[test]
pub(crate) fn validate_rejects_zero_rows_before_cuda_grid_construction() {
let buffers = make_buffers(1, 4, 1, 1);
let inputs = BmsFlexRowKernelInputs {
n_rows: 0,
..minimal_inputs(&buffers)
};
let err = inputs
.validate()
.expect_err("zero-row launch must fail closed");
assert!(err.to_string().contains("n_rows must be > 0"));
}
#[test]
pub(crate) fn validate_rejects_mismatched_r_decomposition() {
let buffers = make_buffers(1, 4, 1, 1);
let bad_inputs = BmsFlexRowKernelInputs {
r: 4,
p_h: 1,
p_w: 2, ..minimal_inputs(&buffers)
};
let err = bad_inputs
.validate()
.expect_err("inconsistent r vs p_h+p_w must fail");
let msg = err.to_string();
assert!(msg.contains("p_h"), "got: {msg}");
assert!(msg.contains("p_w"), "got: {msg}");
}
#[test]
pub(crate) fn validate_rejects_non_monotone_offsets() {
let mut buffers = make_buffers(2, 4, 1, 1);
buffers.cell_offsets = vec![5, 2];
let inputs = minimal_inputs(&buffers);
let err = inputs
.validate()
.expect_err("non-monotone offsets must fail");
let msg = err.to_string();
assert!(msg.contains("monotone"), "got: {msg}");
}
#[test]
pub(crate) fn validate_rejects_mismatched_cell_moments_length() {
let mut buffers = make_buffers(2, 4, 1, 1);
buffers.cell_moments.pop(); let inputs = minimal_inputs(&buffers);
let err = inputs.validate().expect_err("short cell_moments must fail");
let msg = err.to_string();
assert!(msg.contains("cell_moments"), "got: {msg}");
}
#[test]
pub(crate) fn launch_on_non_linux_reports_driver_library_unavailable() {
#[cfg(target_os = "linux")]
{
if cuda_runtime_for_test("bms_flex_row launch smoke test").is_none() {
return;
}
let buffers = make_buffers(1, 4, 1, 1);
let inputs = minimal_inputs(&buffers);
launch_bms_flex_row_kernel(inputs)
.expect("BMS FLEX row kernel must launch after CUDA admission");
}
#[cfg(not(target_os = "linux"))]
{
let buffers = make_buffers(1, 4, 1, 1);
let inputs = minimal_inputs(&buffers);
match launch_bms_flex_row_kernel(inputs) {
Err(GpuError::DriverLibraryUnavailable { reason }) => {
assert!(
reason.contains("Linux-only"),
"expected Linux-only hint, got: {reason}"
);
}
other => panic!("expected DriverLibraryUnavailable on non-Linux, got {other:?}"),
}
}
}
#[test]
pub(crate) fn s_f_must_be_positive_and_finite() {
let buffers = make_buffers(1, 4, 1, 1);
let mut inputs = minimal_inputs(&buffers);
inputs.s_f = 0.0;
match launch_bms_flex_row_kernel(inputs) {
Err(GpuError::DriverCallFailed { reason }) => {
assert!(reason.contains("s_f"), "got: {reason}");
}
other => panic!("expected DriverCallFailed for s_f=0, got {other:?}"),
}
}
#[test]
pub(crate) fn device_mills_layer_matches_finite_differences() {
let neglog_of = |e: f64, y: f64, w: f64| -> f64 {
let s = 2.0 * y - 1.0;
let (log_cdf, _) = host_log_ndtr_and_mills(s * e);
-w * log_cdf
};
let ab_of = |e: f64, y: f64, w: f64| -> (f64, f64) {
let s = 2.0 * y - 1.0;
let m_arg = s * e;
let (_, lambda, probit_curvature) = host_log_ndtr_mills_curvature(m_arg);
let a_i = -w * s * lambda;
let b_i = w * probit_curvature;
(a_i, b_i)
};
let cases: [(f64, f64, f64); 12] = [
(-1.6, 1.0, 1.0),
(-0.7, 1.0, 1.0),
(0.0, 1.0, 1.0),
(0.9, 1.0, 1.0),
(1.8, 1.0, 1.0),
(-1.4, 0.0, 1.0),
(-0.3, 0.0, 1.0),
(0.0, 0.0, 1.0),
(0.6, 0.0, 1.0),
(1.5, 0.0, 1.0),
(0.4, 1.0, 0.75),
(-0.8, 0.0, 1.3),
];
let h = 1e-3_f64;
for (e, y, w) in cases {
let (a_ana, b_ana) = ab_of(e, y, w);
let fp2 = neglog_of(e + 2.0 * h, y, w);
let fp1 = neglog_of(e + h, y, w);
let f0 = neglog_of(e, y, w);
let fm1 = neglog_of(e - h, y, w);
let fm2 = neglog_of(e - 2.0 * h, y, w);
let d1_fd = (-fp2 + 8.0 * fp1 - 8.0 * fm1 + fm2) / (12.0 * h);
let d2_fd = (-fp2 + 16.0 * fp1 - 30.0 * f0 + 16.0 * fm1 - fm2) / (12.0 * h * h);
let a_abs = (a_ana - d1_fd).abs();
let a_rel = a_abs / a_ana.abs().max(1.0);
assert!(
a_abs <= 5e-8 || a_rel <= 5e-8,
"Mills A (∂neglog/∂e) drift at e={e} y={y} w={w}: \
analytic={a_ana:.17e} fd={d1_fd:.17e} abs={a_abs:.3e} rel={a_rel:.3e}"
);
let b_abs = (b_ana - d2_fd).abs();
let b_rel = b_abs / b_ana.abs().max(1.0);
assert!(
b_abs <= 5e-6 || b_rel <= 5e-6,
"Mills B (∂²neglog/∂e²) drift at e={e} y={y} w={w}: \
analytic={b_ana:.17e} fd={d2_fd:.17e} abs={b_abs:.3e} rel={b_rel:.3e}"
);
}
}
#[test]
pub(crate) fn generated_source_interprets_compact_canonical_phase_streams() {
let source = generated_row_kernel_source();
assert!(!source.contains("__BMS_FLEX_CALIBRATION_ORDER2__"));
assert!(!source.contains("__BMS_FLEX_ORDER2_FINALIZER__"));
assert!(!source.contains("__BMS_FLEX_ROW_THREADS__"));
assert!(source.contains("for (int u = 1; u < r; ++u)"));
assert!(source.contains("for (int v = u; v < r; ++v)"));
assert!(source.contains("Canonical implicit-first stage complete"));
assert!(source.contains("double *F_u = out_grad + row_r_base"));
assert!(source.contains("double *F_au = row_f_au + row_r_base"));
assert!(source.contains("double *F_uv = out_hess + row_rr_base"));
for forbidden in [
"MAX_R",
"double F_u[",
"double F_au[",
"double F_uv[",
"double a_u[",
"double a_uv[",
"double bar_e_u[",
] {
assert!(
!source.contains(forbidden),
"generated row source restored width-bound scratch: {forbidden}"
);
}
for forbidden in [
"MAX_R",
"double row_dir[",
"double action[",
"bms_flex_row_hvp_partial_packed",
"bms_flex_row_diag_partial_packed",
"bms_flex_row_pack_upper",
] {
assert!(
!HVP_KERNEL_SOURCE.contains(forbidden),
"HVP source restored a dead or width-bound path: {forbidden}"
);
}
assert!(HVP_KERNEL_SOURCE.contains("bms_flex_primary_direction"));
assert!(HVP_KERNEL_SOURCE.contains("direction_q[MAX_MULTI_RHS]"));
assert!(HVP_KERNEL_SOURCE.contains("action_g[MAX_MULTI_RHS]"));
let mut cursor = 0usize;
for marker in [
"canonical calibration phase: InterceptFirst",
"canonical calibration phase: InterceptSecond",
"canonical calibration phase: PrimaryFirstAndInterceptSecond",
"canonical calibration phase: PrimaryPairSecond",
"canonical finalizer phase: ImplicitFirst",
"canonical finalizer phase: ImplicitFirstComplete",
"canonical finalizer phase: ImplicitSecond",
"canonical finalizer phase: ObservedFirst",
"canonical finalizer phase: ObservedScoreSensitivity",
"canonical finalizer phase: ObservedSecond",
"canonical finalizer phase: NegLogFirst",
] {
let relative = source[cursor..]
.find(marker)
.unwrap_or_else(|| panic!("generated CUDA source omitted phase {marker}"));
cursor += relative + marker.len();
}
assert!(
source.len() < 40_000,
"generated CUDA source unexpectedly bloated"
);
}
pub(crate) fn cpu_oracle_bms_flex_row_hvp(
row_hessians: &[f64],
marginal_design: &[f64],
logslope_design: &[f64],
block: &BmsFlexBlockLayout,
primary: &BmsFlexPrimaryLayout,
n: usize,
v: &[f64],
) -> Vec<f64> {
let r = primary.r;
let p_m = block.p_m;
let p_g = block.p_g;
assert_eq!(v.len(), block.p_total);
assert_eq!(row_hessians.len(), n * r * r);
assert_eq!(marginal_design.len(), n * p_m);
assert_eq!(logslope_design.len(), n * p_g);
let mut out = vec![0.0_f64; block.p_total];
let mut row_dir = vec![0.0_f64; r];
let mut action = vec![0.0_f64; r];
for row in 0..n {
let mrow = &marginal_design[row * p_m..(row + 1) * p_m];
let grow = &logslope_design[row * p_g..(row + 1) * p_g];
let mut acc_q = 0.0_f64;
for j in 0..p_m {
acc_q += mrow[j] * v[j];
}
let mut acc_g = 0.0_f64;
for j in 0..p_g {
acc_g += grow[j] * v[p_m + j];
}
row_dir[0] = acc_q;
row_dir[1] = acc_g;
if let (Some(prange), Some(brange)) = (primary.h.as_ref(), block.h.as_ref()) {
for (k, ii) in prange.clone().enumerate() {
row_dir[ii] = v[brange.start + k];
}
}
if let (Some(prange), Some(brange)) = (primary.w.as_ref(), block.w.as_ref()) {
for (k, ii) in prange.clone().enumerate() {
row_dir[ii] = v[brange.start + k];
}
}
let h_slice = &row_hessians[row * r * r..(row + 1) * r * r];
for u in 0..r {
let mut acc = 0.0_f64;
for v_idx in 0..r {
acc += h_slice[u * r + v_idx] * row_dir[v_idx];
}
action[u] = acc;
}
let a0 = action[0];
for j in 0..p_m {
out[j] += a0 * mrow[j];
}
let a1 = action[1];
for j in 0..p_g {
out[p_m + j] += a1 * grow[j];
}
if let (Some(prange), Some(brange)) = (primary.h.as_ref(), block.h.as_ref()) {
for (k, ii) in prange.clone().enumerate() {
out[brange.start + k] += action[ii];
}
}
if let (Some(prange), Some(brange)) = (primary.w.as_ref(), block.w.as_ref()) {
for (k, ii) in prange.clone().enumerate() {
out[brange.start + k] += action[ii];
}
}
}
out
}
pub(crate) fn cpu_oracle_bms_flex_row_diagonal(
row_hessians: &[f64],
marginal_design: &[f64],
logslope_design: &[f64],
block: &BmsFlexBlockLayout,
primary: &BmsFlexPrimaryLayout,
n: usize,
) -> Vec<f64> {
let r = primary.r;
let p_m = block.p_m;
let p_g = block.p_g;
let mut out = vec![0.0_f64; block.p_total];
for row in 0..n {
let h_slice = &row_hessians[row * r * r..(row + 1) * r * r];
let h00 = h_slice[0];
let h11 = h_slice[r + 1];
let mrow = &marginal_design[row * p_m..(row + 1) * p_m];
let grow = &logslope_design[row * p_g..(row + 1) * p_g];
for j in 0..p_m {
out[j] += h00 * mrow[j] * mrow[j];
}
for j in 0..p_g {
out[p_m + j] += h11 * grow[j] * grow[j];
}
if let (Some(prange), Some(brange)) = (primary.h.as_ref(), block.h.as_ref()) {
for (k, ii) in prange.clone().enumerate() {
out[brange.start + k] += h_slice[ii * r + ii];
}
}
if let (Some(prange), Some(brange)) = (primary.w.as_ref(), block.w.as_ref()) {
for (k, ii) in prange.clone().enumerate() {
out[brange.start + k] += h_slice[ii * r + ii];
}
}
}
out
}
pub(crate) fn cpu_oracle_bms_flex_row_joint_gradient(
row_neglog: &[f64],
row_grad: &[f64],
marginal_design: &[f64],
logslope_design: &[f64],
block: &BmsFlexBlockLayout,
primary: &BmsFlexPrimaryLayout,
n: usize,
) -> (f64, Vec<f64>) {
let r = primary.r;
assert_eq!(row_neglog.len(), n);
assert_eq!(row_grad.len(), n * r);
assert_eq!(marginal_design.len(), n * block.p_m);
assert_eq!(logslope_design.len(), n * block.p_g);
let mut log_likelihood = 0.0_f64;
let mut gradient = vec![0.0_f64; block.p_total];
for row in 0..n {
log_likelihood -= row_neglog[row];
let grow = &row_grad[row * r..(row + 1) * r];
for j in 0..block.p_m {
gradient[j] -= grow[0] * marginal_design[row * block.p_m + j];
}
for j in 0..block.p_g {
gradient[block.p_m + j] -= grow[1] * logslope_design[row * block.p_g + j];
}
if let (Some(primary_h), Some(block_h)) = (primary.h.as_ref(), block.h.as_ref()) {
for (offset, primary_idx) in primary_h.clone().enumerate() {
gradient[block_h.start + offset] -= grow[primary_idx];
}
}
if let (Some(primary_w), Some(block_w)) = (primary.w.as_ref(), block.w.as_ref()) {
for (offset, primary_idx) in primary_w.clone().enumerate() {
gradient[block_w.start + offset] -= grow[primary_idx];
}
}
}
(log_likelihood, gradient)
}
#[test]
fn cpu_joint_gradient_oracle_pins_score_sign_and_active_hw_pullback() {
let n = 2_usize;
let r = 5_usize;
let block = BmsFlexBlockLayout {
p_m: 2,
p_g: 1,
h: Some(3..5),
w: Some(5..6),
p_total: 6,
};
let primary = BmsFlexPrimaryLayout {
h: Some(2..4),
w: Some(4..5),
r,
};
let row_neglog = [1.25, 0.75];
let row_grad = [
2.0, -3.0, 5.0, -7.0, 11.0, -13.0, 17.0, -19.0, 23.0, -29.0, ];
let marginal = [1.0, 2.0, -0.5, 3.0];
let logslope = [4.0, -2.0];
let (log_likelihood, gradient) = cpu_oracle_bms_flex_row_joint_gradient(
&row_neglog,
&row_grad,
&marginal,
&logslope,
&block,
&primary,
n,
);
assert_eq!(log_likelihood, -2.0);
assert_eq!(
gradient,
vec![-8.5, 35.0, 46.0, 14.0, -16.0, 18.0],
"joint output must be the score/log-likelihood sign, with h/w direct slots"
);
}
#[test]
pub(crate) fn cpu_oracle_hvp_matches_hand_computation_no_hw() {
let n = 4_usize;
let r = 4_usize; let p_m = 2_usize;
let p_g = 2_usize;
let p_h_dim = 1_usize;
let p_w_dim = 1_usize;
let p_total = p_m + p_g + p_h_dim + p_w_dim;
let block = BmsFlexBlockLayout {
p_m,
p_g,
h: Some(p_m + p_g..p_m + p_g + p_h_dim),
w: Some(p_m + p_g + p_h_dim..p_m + p_g + p_h_dim + p_w_dim),
p_total,
};
let primary = BmsFlexPrimaryLayout {
h: Some(2..3),
w: Some(3..4),
r,
};
let mut row_hessians = vec![0.0_f64; n * r * r];
for row in 0..n {
for u in 0..r {
for v in u..r {
let val = ((row + 1) as f64) * (1.0 + (u as f64) + 2.0 * (v as f64));
row_hessians[row * r * r + u * r + v] = val;
row_hessians[row * r * r + v * r + u] = val;
}
}
}
let mut marginal = vec![0.0_f64; n * p_m];
for row in 0..n {
for j in 0..p_m {
marginal[row * p_m + j] = 0.5 + (row as f64) * 0.1 - (j as f64) * 0.2;
}
}
let mut logslope = vec![0.0_f64; n * p_g];
for row in 0..n {
for j in 0..p_g {
logslope[row * p_g + j] = -0.3 + (row as f64) * 0.05 + (j as f64) * 0.15;
}
}
let v: Vec<f64> = (0..p_total).map(|i| 0.1 + (i as f64) * 0.25).collect();
let out = cpu_oracle_bms_flex_row_hvp(
&row_hessians,
&marginal,
&logslope,
&block,
&primary,
n,
&v,
);
let mut expect_out_0 = 0.0_f64;
for row in 0..n {
let mrow = &marginal[row * p_m..(row + 1) * p_m];
let grow = &logslope[row * p_g..(row + 1) * p_g];
let mut row_dir = vec![0.0_f64; r];
row_dir[0] = mrow[0] * v[0] + mrow[1] * v[1];
row_dir[1] = grow[0] * v[p_m] + grow[1] * v[p_m + 1];
row_dir[2] = v[p_m + p_g];
row_dir[3] = v[p_m + p_g + p_h_dim];
let h_slice = &row_hessians[row * r * r..(row + 1) * r * r];
let mut action0 = 0.0_f64;
for vv in 0..r {
action0 += h_slice[vv] * row_dir[vv];
}
expect_out_0 += action0 * mrow[0];
}
assert!(
(out[0] - expect_out_0).abs() < 1e-12,
"cpu oracle HVP out[0] mismatch: {} vs hand-check {}",
out[0],
expect_out_0
);
assert!(out.iter().all(|x| x.is_finite()));
assert_eq!(out.len(), p_total);
}
#[test]
pub(crate) fn cpu_oracle_diagonal_matches_hand_computation() {
let n = 3_usize;
let r = 4_usize;
let p_m = 2_usize;
let p_g = 2_usize;
let p_h_dim = 1_usize;
let p_w_dim = 1_usize;
let p_total = p_m + p_g + p_h_dim + p_w_dim;
let block = BmsFlexBlockLayout {
p_m,
p_g,
h: Some(p_m + p_g..p_m + p_g + p_h_dim),
w: Some(p_m + p_g + p_h_dim..p_m + p_g + p_h_dim + p_w_dim),
p_total,
};
let primary = BmsFlexPrimaryLayout {
h: Some(2..3),
w: Some(3..4),
r,
};
let mut row_hessians = vec![0.0_f64; n * r * r];
for row in 0..n {
for u in 0..r {
row_hessians[row * r * r + u * r + u] = 1.0 + (row as f64) + (u as f64) * 0.5;
}
}
let mut marginal = vec![0.0_f64; n * p_m];
let mut logslope = vec![0.0_f64; n * p_g];
for row in 0..n {
for j in 0..p_m {
marginal[row * p_m + j] = 0.2 + (row as f64) * 0.3 + (j as f64) * 0.1;
}
for j in 0..p_g {
logslope[row * p_g + j] = -0.4 + (row as f64) * 0.1 + (j as f64) * 0.2;
}
}
let out = cpu_oracle_bms_flex_row_diagonal(
&row_hessians,
&marginal,
&logslope,
&block,
&primary,
n,
);
let mut expect = 0.0_f64;
for row in 0..n {
let h00 = row_hessians[row * r * r];
expect += h00 * marginal[row * p_m].powi(2);
}
assert!(
(out[0] - expect).abs() < 1e-12,
"out[0] {} vs {}",
out[0],
expect
);
let mut expect_h = 0.0_f64;
for row in 0..n {
expect_h += row_hessians[row * r * r + 2 * r + 2];
}
let h_slot = p_m + p_g;
assert!(
(out[h_slot] - expect_h).abs() < 1e-12,
"h slot {} vs {}",
out[h_slot],
expect_h
);
}
#[test]
pub(crate) fn bms_flex_row_r33_consumers_match_cpu_oracles_when_cuda_available() {
configure_global_policy(GpuPolicy::Required);
assert_eq!(
gam_gpu::global_policy(),
GpuPolicy::Required,
"fresh-process r=33 consumer parity must claim Required before runtime discovery"
);
gam_gpu::device_runtime::GpuRuntime::require()
.expect("#932 mandatory r=33 consumer CUDA runtime");
let n = 3_usize;
let p_h_dim = 16_usize;
let p_w_dim = 15_usize;
let r = 2 + p_h_dim + p_w_dim;
let p_m = 2_usize;
let p_g = 2_usize;
let p_total = p_m + p_g + p_h_dim + p_w_dim;
let block = BmsFlexBlockLayout {
p_m,
p_g,
h: Some(p_m + p_g..p_m + p_g + p_h_dim),
w: Some(p_m + p_g + p_h_dim..p_m + p_g + p_h_dim + p_w_dim),
p_total,
};
let primary = BmsFlexPrimaryLayout {
h: Some(2..2 + p_h_dim),
w: Some(2 + p_h_dim..2 + p_h_dim + p_w_dim),
r,
};
let mut row_hessians = vec![0.0_f64; n * r * r];
for row in 0..n {
for u in 0..r {
for v in u..r {
let val = 0.001 * ((row + 1) as f64) * (1.0 + (u as f64) + 2.0 * (v as f64));
row_hessians[row * r * r + u * r + v] = val;
row_hessians[row * r * r + v * r + u] = val;
}
}
}
let mut marginal = vec![0.0_f64; n * p_m];
for row in 0..n {
for j in 0..p_m {
marginal[row * p_m + j] = 0.5 + (row as f64) * 0.1 - (j as f64) * 0.2;
}
}
let mut logslope = vec![0.0_f64; n * p_g];
for row in 0..n {
for j in 0..p_g {
logslope[row * p_g + j] = -0.3 + (row as f64) * 0.05 + (j as f64) * 0.15;
}
}
let v: Vec<f64> = (0..p_total).map(|i| 0.1 + (i as f64) * 0.25).collect();
let cpu_hvp = cpu_oracle_bms_flex_row_hvp(
&row_hessians,
&marginal,
&logslope,
&block,
&primary,
n,
&v,
);
let cpu_diag = cpu_oracle_bms_flex_row_diagonal(
&row_hessians,
&marginal,
&logslope,
&block,
&primary,
n,
);
let row_neglog = (0..n)
.map(|row| 0.25 + 0.125 * row as f64)
.collect::<Vec<_>>();
let row_grad = (0..n * r)
.map(|index| {
let row = index / r;
let primary_idx = index % r;
(row as f64 + 0.75) * (primary_idx as f64 - 1.25)
})
.collect::<Vec<_>>();
let (cpu_log_likelihood, cpu_gradient) = cpu_oracle_bms_flex_row_joint_gradient(
&row_neglog,
&row_grad,
&marginal,
&logslope,
&block,
&primary,
n,
);
let mut cpu_dense = vec![0.0_f64; p_total * p_total];
for column in 0..p_total {
let mut basis = vec![0.0_f64; p_total];
basis[column] = 1.0;
let image = cpu_oracle_bms_flex_row_hvp(
&row_hessians,
&marginal,
&logslope,
&block,
&primary,
n,
&basis,
);
for (row, value) in image.into_iter().enumerate() {
cpu_dense[row * p_total + column] = value;
}
}
let backend = HvpKernelBackend::probe()
.expect("[bms_flex_row hvp parity] backend probe must succeed on CUDA host");
let stream = backend.stream.clone();
let d_h = stream
.clone_htod(&row_hessians)
.expect("[bms_flex_row hvp parity] upload h must succeed on CUDA host");
let d_m = stream
.clone_htod(&marginal)
.expect("[bms_flex_row hvp parity] upload marg must succeed on CUDA host");
let d_g = stream
.clone_htod(&logslope)
.expect("[bms_flex_row hvp parity] upload logslope must succeed on CUDA host");
let storage = DeviceResidentRowHess {
neglog: stream
.clone_htod(&row_neglog)
.expect("[bms_flex_row hvp parity] upload neglog"),
grad: stream
.clone_htod(&row_grad)
.expect("[bms_flex_row hvp parity] upload grad"),
hess: d_h,
marginal_design: d_m,
logslope_design: d_g,
n,
r,
block: block.clone(),
primary: primary.clone(),
bytes: ((n + n * r + n * r * r + n * p_m + n * p_g) * std::mem::size_of::<f64>())
as u64,
};
let gpu_hvp =
launch_bms_flex_row_hvp(&storage, &v).expect("HVP kernel must launch on CUDA host");
let gpu_diag = launch_bms_flex_row_diagonal(&storage)
.expect("diagonal kernel must launch on CUDA host");
let gpu_joint = launch_bms_flex_row_joint_gradient(&storage)
.expect("joint-gradient kernel must launch on CUDA host");
let gpu_dense = launch_bms_flex_row_dense(&storage)
.expect("dense kernel must launch at r=33 on CUDA host");
assert_eq!(gpu_hvp.len(), cpu_hvp.len());
assert_eq!(gpu_diag.len(), cpu_diag.len());
assert_eq!(gpu_joint.gradient.len(), cpu_gradient.len());
assert!(
(gpu_joint.log_likelihood - cpu_log_likelihood).abs() <= 1e-12,
"loglik: cpu={} gpu={}",
cpu_log_likelihood,
gpu_joint.log_likelihood
);
for i in 0..p_total {
let diff = (cpu_hvp[i] - gpu_hvp[i]).abs();
assert!(
diff <= 1e-10,
"HVP[{i}]: cpu={} gpu={} |Δ|={diff:.3e}",
cpu_hvp[i],
gpu_hvp[i]
);
let ddiff = (cpu_diag[i] - gpu_diag[i]).abs();
assert!(
ddiff <= 1e-10,
"diag[{i}]: cpu={} gpu={} |Δ|={ddiff:.3e}",
cpu_diag[i],
gpu_diag[i]
);
let gdiff = (cpu_gradient[i] - gpu_joint.gradient[i]).abs();
assert!(
gdiff <= 1e-10,
"joint gradient[{i}]: cpu={} gpu={} |Δ|={gdiff:.3e}",
cpu_gradient[i],
gpu_joint.gradient[i]
);
}
assert_eq!(gpu_dense.len(), cpu_dense.len());
for (index, (&cpu, &gpu)) in cpu_dense.iter().zip(&gpu_dense).enumerate() {
let tolerance = 1e-10 * (1.0 + cpu.abs());
assert!(
(cpu - gpu).abs() <= tolerance,
"dense[{index}] at r=33: cpu={cpu} gpu={gpu} tolerance={tolerance}"
);
}
}
#[test]
pub(crate) fn bms_flex_row_hvp_multi_scratch_is_bounded_at_large_scale_shape() {
let n = 195_000_usize;
let r = 20_usize;
let p_total = 44_usize;
let rhs_count = 4_usize;
let scratch = bms_flex_row_hvp_multi_scratch_bytes_for_shape(n, p_total, rhs_count)
.expect("large-scale multi-RHS scratch budget");
let per_rhs_full_row_cache =
(n * r * r * std::mem::size_of::<f64>()) as u64 * rhs_count as u64;
assert!(
scratch < per_rhs_full_row_cache / 100,
"multi-RHS scratch must tile by row chunks instead of materializing \
a row-Hessian copy per RHS: scratch={scratch} full_per_rhs={per_rhs_full_row_cache}"
);
assert!(
bms_flex_row_hvp_multi_scratch_bytes_for_shape(
n,
p_total,
BMS_FLEX_ROW_HVP_MAX_RHS + 1
)
.is_err(),
"multi-RHS launch must reject unbounded RHS counts"
);
}
#[test]
pub(crate) fn bms_flex_row_hvp_multi_kernel_matches_cpu_oracle_when_cuda_available() {
if cuda_runtime_for_test("bms_flex_row hvp_multi parity").is_none() {
return;
}
let n = 5_usize;
let r = 4_usize;
let p_m = 2_usize;
let p_g = 2_usize;
let p_h_dim = 1_usize;
let p_w_dim = 1_usize;
let p_total = p_m + p_g + p_h_dim + p_w_dim;
let rhs_count = 3_usize;
let block = BmsFlexBlockLayout {
p_m,
p_g,
h: Some(p_m + p_g..p_m + p_g + p_h_dim),
w: Some(p_m + p_g + p_h_dim..p_m + p_g + p_h_dim + p_w_dim),
p_total,
};
let primary = BmsFlexPrimaryLayout {
h: Some(2..3),
w: Some(3..4),
r,
};
let mut row_hessians = vec![0.0_f64; n * r * r];
for row in 0..n {
for u in 0..r {
for v in u..r {
let val = ((row + 1) as f64) * (1.0 + (u as f64) + 2.0 * (v as f64));
row_hessians[row * r * r + u * r + v] = val;
row_hessians[row * r * r + v * r + u] = val;
}
}
}
let mut marginal = vec![0.0_f64; n * p_m];
let mut logslope = vec![0.0_f64; n * p_g];
for row in 0..n {
for j in 0..p_m {
marginal[row * p_m + j] = 0.5 + (row as f64) * 0.1 - (j as f64) * 0.2;
}
for j in 0..p_g {
logslope[row * p_g + j] = -0.3 + (row as f64) * 0.05 + (j as f64) * 0.15;
}
}
let mut v_rhs = vec![0.0_f64; rhs_count * p_total];
for rhs in 0..rhs_count {
for j in 0..p_total {
let seed = (rhs as f64) * 0.37 + (j as f64) * 0.19 + 0.4;
v_rhs[rhs * p_total + j] = seed.sin() * 0.4 + seed.cos() * 0.2;
}
}
let backend = HvpKernelBackend::probe()
.expect("[bms_flex_row hvp_multi parity] backend probe must succeed on CUDA host");
let stream = backend.stream.clone();
let d_h = stream
.clone_htod(&row_hessians)
.expect("[bms_flex_row hvp_multi parity] upload h must succeed on CUDA host");
let d_m = stream
.clone_htod(&marginal)
.expect("[bms_flex_row hvp_multi parity] upload marg must succeed on CUDA host");
let d_g = stream
.clone_htod(&logslope)
.expect("[bms_flex_row hvp_multi parity] upload logslope must succeed on CUDA host");
let storage = DeviceResidentRowHess {
neglog: stream
.alloc_zeros::<f64>(n)
.expect("[bms_flex_row hvp_multi parity] alloc neglog"),
grad: stream
.alloc_zeros::<f64>(n * r)
.expect("[bms_flex_row hvp_multi parity] alloc grad"),
hess: d_h,
marginal_design: d_m,
logslope_design: d_g,
n,
r,
block: block.clone(),
primary: primary.clone(),
bytes: ((n + n * r + n * r * r + n * p_m + n * p_g) * std::mem::size_of::<f64>())
as u64,
};
let scratch = bms_flex_row_hvp_multi_scratch_bytes_for_shape(n, p_total, rhs_count)
.expect("storage scratch budget");
assert!(
scratch < storage.bytes,
"multi-RHS scratch should stay below resident cache bytes"
);
let gpu = launch_bms_flex_row_hvp_multi(&storage, &v_rhs, rhs_count)
.expect("multi-RHS HVP kernel must launch on CUDA host");
assert_eq!(gpu.len(), rhs_count * p_total);
for rhs in 0..rhs_count {
let v = &v_rhs[rhs * p_total..(rhs + 1) * p_total];
let cpu = cpu_oracle_bms_flex_row_hvp(
&row_hessians,
&marginal,
&logslope,
&block,
&primary,
n,
v,
);
let single = launch_bms_flex_row_hvp(&storage, v)
.expect("single-RHS HVP kernel must launch on CUDA host");
for j in 0..p_total {
let got = gpu[rhs * p_total + j];
let diff = (cpu[j] - got).abs();
assert!(
diff <= 1e-10,
"multi-RHS HVP rhs={rhs} j={j}: cpu={} gpu={} |diff|={diff:.3e}",
cpu[j],
got
);
assert_eq!(
got, single[j],
"multi-RHS and single-RHS host launch diverged at rhs={rhs} j={j}"
);
}
}
}
#[test]
pub(crate) fn bms_flex_row_hvp_into_device_matches_cpu_oracle_and_host_out() {
#[cfg(not(target_os = "linux"))]
{
eprintln!(
"[bms_flex_row hvp_into_device parity] non-Linux host — skipping \
CUDA parity (CPU oracle exercised by sibling tests)"
);
}
#[cfg(target_os = "linux")]
{
if cuda_runtime_for_test("bms_flex_row hvp_into_device parity").is_none() {
return;
}
let n = 4_usize;
let r = 4_usize;
let p_m = 2_usize;
let p_g = 2_usize;
let p_h_dim = 1_usize;
let p_w_dim = 1_usize;
let p_total = p_m + p_g + p_h_dim + p_w_dim;
let block = BmsFlexBlockLayout {
p_m,
p_g,
h: Some(p_m + p_g..p_m + p_g + p_h_dim),
w: Some(p_m + p_g + p_h_dim..p_m + p_g + p_h_dim + p_w_dim),
p_total,
};
let primary = BmsFlexPrimaryLayout {
h: Some(2..3),
w: Some(3..4),
r,
};
let mut row_hessians = vec![0.0_f64; n * r * r];
for row in 0..n {
for u in 0..r {
for v in u..r {
let val = ((row + 1) as f64) * (1.0 + (u as f64) + 2.0 * (v as f64));
row_hessians[row * r * r + u * r + v] = val;
row_hessians[row * r * r + v * r + u] = val;
}
}
}
let mut marginal = vec![0.0_f64; n * p_m];
for row in 0..n {
for j in 0..p_m {
marginal[row * p_m + j] = 0.5 + (row as f64) * 0.1 - (j as f64) * 0.2;
}
}
let mut logslope = vec![0.0_f64; n * p_g];
for row in 0..n {
for j in 0..p_g {
logslope[row * p_g + j] = -0.3 + (row as f64) * 0.05 + (j as f64) * 0.15;
}
}
let v: Vec<f64> = (0..p_total).map(|i| 0.1 + (i as f64) * 0.25).collect();
let cpu_hvp = cpu_oracle_bms_flex_row_hvp(
&row_hessians,
&marginal,
&logslope,
&block,
&primary,
n,
&v,
);
let backend = HvpKernelBackend::probe().expect(
"[bms_flex_row hvp_into_device parity] backend probe must succeed on CUDA host",
);
let stream = backend.stream.clone();
let d_h = stream
.clone_htod(&row_hessians)
.expect("[bms_flex_row hvp_into_device parity] upload h must succeed on CUDA host");
let d_m = stream.clone_htod(&marginal).expect(
"[bms_flex_row hvp_into_device parity] upload marg must succeed on CUDA host",
);
let d_g = stream.clone_htod(&logslope).expect(
"[bms_flex_row hvp_into_device parity] upload logslope must succeed on CUDA host",
);
let storage = DeviceResidentRowHess {
neglog: stream
.alloc_zeros::<f64>(n)
.expect("[bms_flex_row hvp_into_device parity] alloc neglog"),
grad: stream
.alloc_zeros::<f64>(n * r)
.expect("[bms_flex_row hvp_into_device parity] alloc grad"),
hess: d_h,
marginal_design: d_m,
logslope_design: d_g,
n,
r,
block: block.clone(),
primary: primary.clone(),
bytes: ((n + n * r + n * r * r + n * p_m + n * p_g) * std::mem::size_of::<f64>())
as u64,
};
let host_out_hvp = launch_bms_flex_row_hvp(&storage, &v)
.expect("host-out HVP kernel must launch on CUDA host");
let d_v = stream
.clone_htod(&v)
.expect("upload direction for device-out HVP");
let mut d_out = stream
.alloc_zeros::<f64>(p_total)
.expect("alloc device-out HVP output");
launch_bms_flex_row_hvp_into_device(&storage, &d_v, &mut d_out)
.expect("device-out HVP kernel must launch on CUDA host");
stream
.synchronize()
.expect("synchronize after device-out HVP");
let device_out_hvp = stream
.clone_dtoh(&d_out)
.expect("download device-out HVP output");
assert_eq!(device_out_hvp.len(), cpu_hvp.len());
assert_eq!(device_out_hvp.len(), host_out_hvp.len());
for i in 0..p_total {
let diff = (cpu_hvp[i] - device_out_hvp[i]).abs();
assert!(
diff <= 1e-10,
"device-out HVP[{i}] vs CPU: cpu={} gpu={} |Δ|={diff:.3e}",
cpu_hvp[i],
device_out_hvp[i]
);
let host_diff = (host_out_hvp[i] - device_out_hvp[i]).abs();
assert!(
host_diff == 0.0,
"device-out vs host-out HVP[{i}]: host={} device={} |Δ|={host_diff:.3e}",
host_out_hvp[i],
device_out_hvp[i]
);
}
}
}
#[test]
pub(crate) fn bms_flex_row_hvp_kernel_matches_cpu_oracle_at_n64_r20_p44() {
#[cfg(not(target_os = "linux"))]
{
eprintln!(
"[bms_flex_row hvp parity n64_r20_p44] non-Linux host — \
skipping CUDA parity"
);
}
#[cfg(target_os = "linux")]
{
if cuda_runtime_for_test("bms_flex_row hvp parity n64_r20_p44").is_none() {
return;
}
let n = 64_usize;
let p_m = 14_usize;
let p_g = 12_usize;
let p_h_dim = 10_usize;
let p_w_dim = 8_usize;
let r = 2 + p_h_dim + p_w_dim;
assert_eq!(r, 20);
let p_total = p_m + p_g + p_h_dim + p_w_dim;
assert_eq!(p_total, 44);
let block = BmsFlexBlockLayout {
p_m,
p_g,
h: Some(p_m + p_g..p_m + p_g + p_h_dim),
w: Some(p_m + p_g + p_h_dim..p_m + p_g + p_h_dim + p_w_dim),
p_total,
};
let primary = BmsFlexPrimaryLayout {
h: Some(2..2 + p_h_dim),
w: Some(2 + p_h_dim..2 + p_h_dim + p_w_dim),
r,
};
let mut row_hessians = vec![0.0_f64; n * r * r];
for row in 0..n {
let base = row * r * r;
for u in 0..r {
for v in 0..r {
let seed = (row as f64) * 0.137 + (u as f64) * 1.901 + (v as f64) * 0.317;
let a = (seed.sin() * 1.7 + (seed * 0.5).cos() * 0.9) * 0.5;
row_hessians[base + u * r + v] = a;
}
}
for u in 0..r {
for v in (u + 1)..r {
let upper = row_hessians[base + u * r + v];
let lower = row_hessians[base + v * r + u];
let sym = 0.5 * (upper + lower);
row_hessians[base + u * r + v] = sym;
row_hessians[base + v * r + u] = sym;
}
row_hessians[base + u * r + u] += r as f64;
}
}
let mut marginal = vec![0.0_f64; n * p_m];
for row in 0..n {
for j in 0..p_m {
let seed = (row as f64) * 0.073 + (j as f64) * 0.211 + 0.4;
marginal[row * p_m + j] = seed.sin() * 0.8 - (seed * 0.7).cos() * 0.3;
}
}
let mut logslope = vec![0.0_f64; n * p_g];
for row in 0..n {
for j in 0..p_g {
let seed = (row as f64) * 0.091 + (j as f64) * 0.179 - 0.2;
logslope[row * p_g + j] = seed.cos() * 0.7 + (seed * 0.3).sin() * 0.25;
}
}
let v: Vec<f64> = (0..p_total)
.map(|i| {
let seed = (i as f64) * 0.157 + 0.6;
seed.sin() * 0.55 + (seed * 0.4).cos() * 0.35
})
.collect();
let cpu_hvp = cpu_oracle_bms_flex_row_hvp(
&row_hessians,
&marginal,
&logslope,
&block,
&primary,
n,
&v,
);
let cpu_diag = cpu_oracle_bms_flex_row_diagonal(
&row_hessians,
&marginal,
&logslope,
&block,
&primary,
n,
);
let backend = match HvpKernelBackend::probe() {
Ok(b) => b,
Err(err) => {
eprintln!(
"[bms_flex_row hvp parity n64_r20_p44] backend probe \
failed: {err}"
);
return;
}
};
let stream = backend.stream.clone();
let d_h = match stream.clone_htod(&row_hessians) {
Ok(s) => s,
Err(err) => {
eprintln!(
"[bms_flex_row hvp parity n64_r20_p44] upload h \
failed: {err}"
);
return;
}
};
let d_m = match stream.clone_htod(&marginal) {
Ok(s) => s,
Err(err) => {
eprintln!(
"[bms_flex_row hvp parity n64_r20_p44] upload marg \
failed: {err}"
);
return;
}
};
let d_g = match stream.clone_htod(&logslope) {
Ok(s) => s,
Err(err) => {
eprintln!(
"[bms_flex_row hvp parity n64_r20_p44] upload logslope \
failed: {err}"
);
return;
}
};
let storage = DeviceResidentRowHess {
neglog: stream
.alloc_zeros::<f64>(n)
.expect("[bms_flex_row hvp parity n64_r20_p44] alloc neglog"),
grad: stream
.alloc_zeros::<f64>(n * r)
.expect("[bms_flex_row hvp parity n64_r20_p44] alloc grad"),
hess: d_h,
marginal_design: d_m,
logslope_design: d_g,
n,
r,
block: block.clone(),
primary: primary.clone(),
bytes: ((n + n * r + n * r * r + n * p_m + n * p_g) * std::mem::size_of::<f64>())
as u64,
};
let gpu_hvp = launch_bms_flex_row_hvp(&storage, &v)
.expect("HVP kernel must launch on CUDA host at n64/r20/p44");
let gpu_diag = launch_bms_flex_row_diagonal(&storage)
.expect("diagonal kernel must launch on CUDA host at n64/r20/p44");
assert_eq!(gpu_hvp.len(), cpu_hvp.len());
assert_eq!(gpu_diag.len(), cpu_diag.len());
for i in 0..p_total {
let diff = (cpu_hvp[i] - gpu_hvp[i]).abs();
assert!(
diff <= 1e-8,
"n64_r20_p44 HVP[{i}]: cpu={} gpu={} |Δ|={diff:.3e}",
cpu_hvp[i],
gpu_hvp[i]
);
let ddiff = (cpu_diag[i] - gpu_diag[i]).abs();
assert!(
ddiff <= 1e-8,
"n64_r20_p44 diag[{i}]: cpu={} gpu={} |Δ|={ddiff:.3e}",
cpu_diag[i],
gpu_diag[i]
);
}
}
}
#[test]
pub(crate) fn bms_flex_row_dense_block_kernel_matches_cpu_pullback() {
#[cfg(not(target_os = "linux"))]
{
eprintln!("[bms_flex_row dense_block parity] non-Linux host — skipping CUDA parity");
}
#[cfg(target_os = "linux")]
{
if cuda_runtime_for_test("bms_flex_row dense_block parity").is_none() {
return;
}
let n = 24_usize;
let p_m = 4_usize;
let p_g = 4_usize;
let p_h_dim = 3_usize;
let p_w_dim = 3_usize;
let r = 2 + p_h_dim + p_w_dim;
let p_total = p_m + p_g + p_h_dim + p_w_dim;
let block = BmsFlexBlockLayout {
p_m,
p_g,
h: Some(p_m + p_g..p_m + p_g + p_h_dim),
w: Some(p_m + p_g + p_h_dim..p_m + p_g + p_h_dim + p_w_dim),
p_total,
};
let primary = BmsFlexPrimaryLayout {
h: Some(2..2 + p_h_dim),
w: Some(2 + p_h_dim..2 + p_h_dim + p_w_dim),
r,
};
let mut row_hessians = vec![0.0_f64; n * r * r];
for row in 0..n {
let base = row * r * r;
for u in 0..r {
for v in 0..r {
let seed = (row as f64) * 0.21 + (u as f64) * 1.13 + (v as f64) * 0.47;
let a = (seed.sin() * 1.4 + (seed * 0.6).cos() * 0.7) * 0.5;
row_hessians[base + u * r + v] = a;
}
}
for u in 0..r {
for v in (u + 1)..r {
let upper = row_hessians[base + u * r + v];
let lower = row_hessians[base + v * r + u];
let sym = 0.5 * (upper + lower);
row_hessians[base + u * r + v] = sym;
row_hessians[base + v * r + u] = sym;
}
row_hessians[base + u * r + u] += r as f64;
}
}
let mut marginal = vec![0.0_f64; n * p_m];
for row in 0..n {
for j in 0..p_m {
let seed = (row as f64) * 0.083 + (j as f64) * 0.171 + 0.31;
marginal[row * p_m + j] = seed.sin() * 0.7 - (seed * 0.5).cos() * 0.25;
}
}
let mut logslope = vec![0.0_f64; n * p_g];
for row in 0..n {
for j in 0..p_g {
let seed = (row as f64) * 0.097 + (j as f64) * 0.143 - 0.15;
logslope[row * p_g + j] = seed.cos() * 0.65 + (seed * 0.4).sin() * 0.2;
}
}
let h_block_start = block.h.as_ref().map(|r| r.start).unwrap_or(0);
let h_block_len = block.h.as_ref().map(|r| r.len()).unwrap_or(0);
let w_block_start = block.w.as_ref().map(|r| r.start).unwrap_or(0);
let w_block_len = block.w.as_ref().map(|r| r.len()).unwrap_or(0);
let h_primary_start = primary.h.as_ref().map(|r| r.start).unwrap_or(0);
let w_primary_start = primary.w.as_ref().map(|r| r.start).unwrap_or(0);
let mut h_cpu = vec![0.0_f64; p_total * p_total];
for row in 0..n {
let mrow = &marginal[row * p_m..(row + 1) * p_m];
let grow = &logslope[row * p_g..(row + 1) * p_g];
let hrow = &row_hessians[row * r * r..(row + 1) * r * r];
let mut phi = vec![vec![0.0_f64; p_total]; r];
for k in 0..p_m {
phi[0][k] = mrow[k];
}
for k in 0..p_g {
phi[1][p_m + k] = grow[k];
}
for k in 0..h_block_len {
phi[h_primary_start + k][h_block_start + k] = 1.0;
}
for k in 0..w_block_len {
phi[w_primary_start + k][w_block_start + k] = 1.0;
}
for u in 0..r {
for v in 0..r {
let huv = hrow[u * r + v];
if huv == 0.0 {
continue;
}
for m in 0..p_total {
let pm = phi[u][m];
if pm == 0.0 {
continue;
}
let scaled = huv * pm;
for nn in 0..p_total {
h_cpu[m * p_total + nn] += scaled * phi[v][nn];
}
}
}
}
}
let backend = HvpKernelBackend::probe().expect(
"[bms_flex_row dense_block parity] backend probe must succeed on CUDA host",
);
let stream = backend.stream.clone();
let d_h = stream
.clone_htod(&row_hessians)
.expect("[bms_flex_row dense_block parity] upload h must succeed on CUDA host");
let d_m = stream
.clone_htod(&marginal)
.expect("[bms_flex_row dense_block parity] upload marg must succeed on CUDA host");
let d_g = stream.clone_htod(&logslope).expect(
"[bms_flex_row dense_block parity] upload logslope must succeed on CUDA host",
);
let storage = DeviceResidentRowHess {
neglog: stream
.alloc_zeros::<f64>(n)
.expect("[bms_flex_row dense_block parity] alloc neglog"),
grad: stream
.alloc_zeros::<f64>(n * r)
.expect("[bms_flex_row dense_block parity] alloc grad"),
hess: d_h,
marginal_design: d_m,
logslope_design: d_g,
n,
r,
block: block.clone(),
primary: primary.clone(),
bytes: ((n + n * r + n * r * r + n * p_m + n * p_g) * std::mem::size_of::<f64>())
as u64,
};
let h_gpu = launch_bms_flex_row_dense_block(&storage)
.expect("dense_block kernel must launch on CUDA host");
assert_eq!(h_gpu.len(), p_total * p_total);
let mut max_abs = 0.0_f64;
for i in 0..p_total {
for j in 0..p_total {
let a = h_cpu[i * p_total + j];
let b = h_gpu[i * p_total + j];
let diff = (a - b).abs();
if diff > max_abs {
max_abs = diff;
}
assert!(
diff <= 1e-9 * a.abs().max(b.abs()).max(1.0),
"dense_block[{i},{j}]: cpu={a} gpu={b} |Δ|={diff:.3e}"
);
}
}
eprintln!(
"[bms_flex_row dense_block parity] n={n} r={r} p={p_total}: max|Δ|={max_abs:.3e}"
);
}
}
#[test]
pub(crate) fn bms_flex_row_dense_hvp_materialization_matches_cpu_above_block_cap_932() {
if cuda_runtime_for_test("bms_flex_row dense HVP parity").is_none() {
return;
}
let n = 1_usize;
let r = 2_usize;
let p_m = 37_usize;
let p_g = 36_usize;
let p_total = p_m + p_g;
assert_eq!(p_total, DENSE_BLOCK_MAX_P + 1);
let block = BmsFlexBlockLayout {
p_m,
p_g,
h: None,
w: None,
p_total,
};
let primary = BmsFlexPrimaryLayout {
h: None,
w: None,
r,
};
let row_hessians = vec![2.5_f64, -0.75, -0.75, 1.25];
let marginal = (0..p_m)
.map(|column| 0.15 + (column as f64 * 0.17).sin())
.collect::<Vec<_>>();
let logslope = (0..p_g)
.map(|column| -0.2 + (column as f64 * 0.11).cos())
.collect::<Vec<_>>();
let mut expected = vec![0.0_f64; p_total * p_total];
for row in 0..p_total {
for column in 0..p_total {
expected[row * p_total + column] = match (row < p_m, column < p_m) {
(true, true) => row_hessians[0] * marginal[row] * marginal[column],
(true, false) => row_hessians[1] * marginal[row] * logslope[column - p_m],
(false, true) => row_hessians[2] * logslope[row - p_m] * marginal[column],
(false, false) => {
row_hessians[3] * logslope[row - p_m] * logslope[column - p_m]
}
};
}
}
let backend = HvpKernelBackend::probe()
.expect("[bms_flex_row dense HVP parity] backend probe must succeed");
let stream = backend.stream.clone();
let storage = DeviceResidentRowHess {
neglog: stream
.alloc_zeros::<f64>(n)
.expect("dense HVP parity neglog allocation"),
grad: stream
.alloc_zeros::<f64>(n * r)
.expect("dense HVP parity grad allocation"),
hess: stream
.clone_htod(&row_hessians)
.expect("dense HVP parity hessian upload"),
marginal_design: stream
.clone_htod(&marginal)
.expect("dense HVP parity marginal upload"),
logslope_design: stream
.clone_htod(&logslope)
.expect("dense HVP parity logslope upload"),
n,
r,
block,
primary,
bytes: ((n + n * r + n * r * r + n * p_m + n * p_g) * std::mem::size_of::<f64>())
as u64,
};
let actual = launch_bms_flex_row_dense(&storage)
.expect("wide dense HVP materialization must stay on CUDA");
assert_eq!(actual.len(), expected.len());
for (index, (&actual, &expected)) in actual.iter().zip(&expected).enumerate() {
let tolerance = 1.0e-10 * actual.abs().max(expected.abs()).max(1.0);
assert!(
(actual - expected).abs() <= tolerance,
"wide dense entry {index}: CUDA={actual:.17e} CPU={expected:.17e} tolerance={tolerance:.3e}"
);
}
}
#[test]
pub(crate) fn bms_flex_row_hvp_dispatch_worthiness_at_large_scale() {
#[cfg(not(target_os = "linux"))]
{
eprintln!("[bms_flex_row hvp hill-climb] non-Linux host — skipping V100 perf gate");
}
#[cfg(target_os = "linux")]
{
let Some(runtime) = cuda_runtime_for_test("bms_flex_row hvp hill-climb") else {
return;
};
let n = 195_000_usize;
let p_m = 14_usize;
let p_g = 12_usize;
let p_h_dim = 10_usize;
let p_w_dim = 8_usize;
let r = 2 + p_h_dim + p_w_dim;
let p_total = p_m + p_g + p_h_dim + p_w_dim;
let block = BmsFlexBlockLayout {
p_m,
p_g,
h: Some(p_m + p_g..p_m + p_g + p_h_dim),
w: Some(p_m + p_g + p_h_dim..p_m + p_g + p_h_dim + p_w_dim),
p_total,
};
let primary = BmsFlexPrimaryLayout {
h: Some(2..2 + p_h_dim),
w: Some(2 + p_h_dim..2 + p_h_dim + p_w_dim),
r,
};
let mut row_hessians = vec![0.0_f64; n * r * r];
for row in 0..n {
let base = row * r * r;
for u in 0..r {
for vv in 0..r {
let seed = (row as f64) * 0.137 + (u as f64) * 1.901 + (vv as f64) * 0.317;
let a = (seed.sin() * 1.7 + (seed * 0.5).cos() * 0.9) * 0.5;
row_hessians[base + u * r + vv] = a;
}
}
for u in 0..r {
for vv in (u + 1)..r {
let upper = row_hessians[base + u * r + vv];
let lower = row_hessians[base + vv * r + u];
let sym = 0.5 * (upper + lower);
row_hessians[base + u * r + vv] = sym;
row_hessians[base + vv * r + u] = sym;
}
row_hessians[base + u * r + u] += r as f64;
}
}
let mut marginal = vec![0.0_f64; n * p_m];
for row in 0..n {
for j in 0..p_m {
let seed = (row as f64) * 0.073 + (j as f64) * 0.211 + 0.4;
marginal[row * p_m + j] = seed.sin() * 0.8 - (seed * 0.7).cos() * 0.3;
}
}
let mut logslope = vec![0.0_f64; n * p_g];
for row in 0..n {
for j in 0..p_g {
let seed = (row as f64) * 0.091 + (j as f64) * 0.179 - 0.2;
logslope[row * p_g + j] = seed.cos() * 0.7 + (seed * 0.3).sin() * 0.25;
}
}
let v: Vec<f64> = (0..p_total)
.map(|i| {
let seed = (i as f64) * 0.157 + 0.6;
seed.sin() * 0.55 + (seed * 0.4).cos() * 0.35
})
.collect();
let backend = match HvpKernelBackend::probe() {
Ok(b) => b,
Err(err) => {
eprintln!("[bms_flex_row hvp hill-climb] backend probe failed: {err}");
return;
}
};
let stream = backend.stream.clone();
let d_h = match stream.clone_htod(&row_hessians) {
Ok(s) => s,
Err(err) => {
eprintln!("[bms_flex_row hvp hill-climb] upload h failed (likely OOM): {err}");
return;
}
};
let d_m = match stream.clone_htod(&marginal) {
Ok(s) => s,
Err(err) => {
eprintln!("[bms_flex_row hvp hill-climb] upload marg failed: {err}");
return;
}
};
let d_g = match stream.clone_htod(&logslope) {
Ok(s) => s,
Err(err) => {
eprintln!("[bms_flex_row hvp hill-climb] upload logslope failed: {err}");
return;
}
};
let storage = DeviceResidentRowHess {
neglog: stream
.alloc_zeros::<f64>(n)
.expect("[bms_flex_row hvp hill-climb] alloc neglog"),
grad: stream
.alloc_zeros::<f64>(n * r)
.expect("[bms_flex_row hvp hill-climb] alloc grad"),
hess: d_h,
marginal_design: d_m,
logslope_design: d_g,
n,
r,
block: block.clone(),
primary: primary.clone(),
bytes: ((n + n * r + n * r * r + n * p_m + n * p_g) * std::mem::size_of::<f64>())
as u64,
};
let warmup: usize = 3;
let iters: usize = 15;
for _ in 0..warmup {
let out =
launch_bms_flex_row_hvp(&storage, &v).expect("warmup GPU HVP must launch");
assert_eq!(out.len(), p_total);
}
let mut gpu_us: Vec<u128> = Vec::with_capacity(iters);
for _ in 0..iters {
let t0 = std::time::Instant::now();
let out = launch_bms_flex_row_hvp(&storage, &v).expect("GPU HVP must launch");
gpu_us.push(t0.elapsed().as_micros());
assert_eq!(out.len(), p_total);
}
gpu_us.sort_unstable();
let gpu_median = gpu_us[iters / 2];
const CHUNK_ROWS: usize = 4096;
let cpu_hvp_parallel = || -> Vec<f64> {
let nchunks = n.div_ceil(CHUNK_ROWS);
gam_linalg::pairwise_reduce::par_deterministic_block_fold(
nchunks,
|ci_range| {
let mut acc = vec![0.0_f64; p_total];
for ci in ci_range {
let lo = ci * CHUNK_ROWS;
let hi = (lo + CHUNK_ROWS).min(n);
let m = hi - lo;
let partial = cpu_oracle_bms_flex_row_hvp(
&row_hessians[lo * r * r..hi * r * r],
&marginal[lo * p_m..hi * p_m],
&logslope[lo * p_g..hi * p_g],
&block,
&primary,
m,
&v,
);
for (a, &p) in acc.iter_mut().zip(partial.iter()) {
*a += p;
}
}
acc
},
|mut a, b| {
for (ax, bx) in a.iter_mut().zip(b.iter()) {
*ax += *bx;
}
a
},
)
.unwrap_or_else(|| vec![0.0_f64; p_total])
};
let warm = cpu_hvp_parallel();
assert_eq!(warm.len(), p_total);
let mut cpu_us: Vec<u128> = Vec::with_capacity(iters);
for _ in 0..iters {
let t0 = std::time::Instant::now();
let out = cpu_hvp_parallel();
cpu_us.push(t0.elapsed().as_micros());
assert_eq!(out.len(), p_total);
}
cpu_us.sort_unstable();
let cpu_median = cpu_us[iters / 2];
let speedup = (cpu_median as f64) / (gpu_median.max(1) as f64);
eprintln!(
"[bms_flex_row hvp hill-climb] large-scale n={n} r={r} p={p_total}: \
cpu_median={cpu_median}us gpu_median={gpu_median}us \
speedup={speedup:.2}× (perf record; the gate is the policy decision)"
);
assert_row_batch_dispatch_worthy(
"large-scale HVP dispatch-worthiness gate",
runtime.policy(),
n,
);
}
}
#[test]
pub(crate) fn bms_flex_row_dense_block_dispatch_worthiness_at_large_scale() {
#[cfg(not(target_os = "linux"))]
{
eprintln!(
"[bms_flex_row dense_block hill-climb] non-Linux host — skipping V100 perf gate"
);
}
#[cfg(target_os = "linux")]
{
let Some(runtime) = cuda_runtime_for_test("bms_flex_row dense_block hill-climb")
else {
return;
};
let n = 195_000_usize;
let p_m = 14_usize;
let p_g = 12_usize;
let p_h_dim = 10_usize;
let p_w_dim = 8_usize;
let r = 2 + p_h_dim + p_w_dim;
let p_total = p_m + p_g + p_h_dim + p_w_dim;
let block = BmsFlexBlockLayout {
p_m,
p_g,
h: Some(p_m + p_g..p_m + p_g + p_h_dim),
w: Some(p_m + p_g + p_h_dim..p_m + p_g + p_h_dim + p_w_dim),
p_total,
};
let primary = BmsFlexPrimaryLayout {
h: Some(2..2 + p_h_dim),
w: Some(2 + p_h_dim..2 + p_h_dim + p_w_dim),
r,
};
let mut row_hessians = vec![0.0_f64; n * r * r];
for row in 0..n {
let base = row * r * r;
for u in 0..r {
for vv in 0..r {
let seed = (row as f64) * 0.137 + (u as f64) * 1.901 + (vv as f64) * 0.317;
let a = (seed.sin() * 1.7 + (seed * 0.5).cos() * 0.9) * 0.5;
row_hessians[base + u * r + vv] = a;
}
}
for u in 0..r {
for vv in (u + 1)..r {
let upper = row_hessians[base + u * r + vv];
let lower = row_hessians[base + vv * r + u];
let sym = 0.5 * (upper + lower);
row_hessians[base + u * r + vv] = sym;
row_hessians[base + vv * r + u] = sym;
}
row_hessians[base + u * r + u] += r as f64;
}
}
let mut marginal = vec![0.0_f64; n * p_m];
for row in 0..n {
for j in 0..p_m {
let seed = (row as f64) * 0.073 + (j as f64) * 0.211 + 0.4;
marginal[row * p_m + j] = seed.sin() * 0.8 - (seed * 0.7).cos() * 0.3;
}
}
let mut logslope = vec![0.0_f64; n * p_g];
for row in 0..n {
for j in 0..p_g {
let seed = (row as f64) * 0.091 + (j as f64) * 0.179 - 0.2;
logslope[row * p_g + j] = seed.cos() * 0.7 + (seed * 0.3).sin() * 0.25;
}
}
if p_total > DENSE_BLOCK_MAX_P {
eprintln!(
"[bms_flex_row dense_block hill-climb] p_total={p_total} > MAX={DENSE_BLOCK_MAX_P}, skipping"
);
return;
}
let backend = match HvpKernelBackend::probe() {
Ok(b) => b,
Err(err) => {
eprintln!("[bms_flex_row dense_block hill-climb] backend probe failed: {err}");
return;
}
};
let stream = backend.stream.clone();
let d_h = match stream.clone_htod(&row_hessians) {
Ok(s) => s,
Err(err) => {
eprintln!("[bms_flex_row dense_block hill-climb] upload h failed: {err}");
return;
}
};
let d_m = match stream.clone_htod(&marginal) {
Ok(s) => s,
Err(err) => {
eprintln!("[bms_flex_row dense_block hill-climb] upload marg failed: {err}");
return;
}
};
let d_g = match stream.clone_htod(&logslope) {
Ok(s) => s,
Err(err) => {
eprintln!(
"[bms_flex_row dense_block hill-climb] upload logslope failed: {err}"
);
return;
}
};
let storage = DeviceResidentRowHess {
neglog: stream
.alloc_zeros::<f64>(n)
.expect("[bms_flex_row dense_block hill-climb] alloc neglog"),
grad: stream
.alloc_zeros::<f64>(n * r)
.expect("[bms_flex_row dense_block hill-climb] alloc grad"),
hess: d_h,
marginal_design: d_m,
logslope_design: d_g,
n,
r,
block: block.clone(),
primary: primary.clone(),
bytes: ((n + n * r + n * r * r + n * p_m + n * p_g) * std::mem::size_of::<f64>())
as u64,
};
let warmup: usize = 2;
let iters: usize = 5;
for _ in 0..warmup {
let out = launch_bms_flex_row_dense_block(&storage)
.expect("warmup GPU dense_block must launch");
assert_eq!(out.len(), p_total * p_total);
}
let mut gpu_us: Vec<u128> = Vec::with_capacity(iters);
for _ in 0..iters {
let t0 = std::time::Instant::now();
let out =
launch_bms_flex_row_dense_block(&storage).expect("GPU dense_block must launch");
gpu_us.push(t0.elapsed().as_micros());
assert_eq!(out.len(), p_total * p_total);
}
gpu_us.sort_unstable();
let gpu_median = gpu_us[iters / 2];
const CHUNK_ROWS: usize = 2048;
let h_block_start = block.h.as_ref().map(|r| r.start).unwrap_or(0);
let h_block_len = block.h.as_ref().map(|r| r.len()).unwrap_or(0);
let w_block_start = block.w.as_ref().map(|r| r.start).unwrap_or(0);
let w_block_len = block.w.as_ref().map(|r| r.len()).unwrap_or(0);
let h_primary_start = primary.h.as_ref().map(|r| r.start).unwrap_or(0);
let w_primary_start = primary.w.as_ref().map(|r| r.start).unwrap_or(0);
let cpu_build_parallel = || -> Vec<f64> {
let nchunks = n.div_ceil(CHUNK_ROWS);
gam_linalg::pairwise_reduce::par_deterministic_block_fold(
nchunks,
|ci_range| {
let mut acc = vec![0.0_f64; p_total * p_total];
let mut phi: Vec<Vec<f64>> = vec![vec![0.0_f64; p_total]; r];
for ci in ci_range {
let lo = ci * CHUNK_ROWS;
let hi = (lo + CHUNK_ROWS).min(n);
for row in lo..hi {
for col in phi.iter_mut() {
col.iter_mut().for_each(|v| *v = 0.0);
}
let mrow = &marginal[row * p_m..(row + 1) * p_m];
let grow = &logslope[row * p_g..(row + 1) * p_g];
for k in 0..p_m {
phi[0][k] = mrow[k];
}
for k in 0..p_g {
phi[1][p_m + k] = grow[k];
}
for k in 0..h_block_len {
phi[h_primary_start + k][h_block_start + k] = 1.0;
}
for k in 0..w_block_len {
phi[w_primary_start + k][w_block_start + k] = 1.0;
}
let hrow = &row_hessians[row * r * r..(row + 1) * r * r];
for u in 0..r {
for v_idx in 0..r {
let huv = hrow[u * r + v_idx];
if huv == 0.0 {
continue;
}
for m in 0..p_total {
let pm = phi[u][m];
if pm == 0.0 {
continue;
}
let scaled = huv * pm;
for nn in 0..p_total {
acc[m * p_total + nn] += scaled * phi[v_idx][nn];
}
}
}
}
}
}
acc
},
|mut a, b| {
for (ax, bx) in a.iter_mut().zip(b.iter()) {
*ax += *bx;
}
a
},
)
.unwrap_or_else(|| vec![0.0_f64; p_total * p_total])
};
let warm_cpu = cpu_build_parallel();
assert_eq!(warm_cpu.len(), p_total * p_total);
let mut cpu_us: Vec<u128> = Vec::with_capacity(iters);
for _ in 0..iters {
let t0 = std::time::Instant::now();
let out = cpu_build_parallel();
cpu_us.push(t0.elapsed().as_micros());
assert_eq!(out.len(), p_total * p_total);
}
cpu_us.sort_unstable();
let cpu_median = cpu_us[iters / 2];
let speedup = (cpu_median as f64) / (gpu_median.max(1) as f64);
eprintln!(
"[bms_flex_row dense_block hill-climb] large-scale n={n} r={r} p={p_total}: \
cpu_median={cpu_median}us gpu_median={gpu_median}us \
speedup={speedup:.2}× (perf record; the gate is the policy decision)"
);
assert_row_batch_dispatch_worthy(
"large-scale dense-H dispatch-worthiness gate",
runtime.policy(),
n,
);
}
}
}