use std::num::NonZeroUsize;
use super::fit::{Edm, dense_features, quantile_sorted, residual_mean};
use super::process::{T_EPS, keyed_normal, try_filled};
use super::{DiffusionModel, Method, OdeSolver, Parameterization, ScoreConfig};
use crate::data::DMatrix;
use crate::error::{HessboostError, Result};
use crate::model::{Iterations, Predictions};
use crate::rng::splitmix64;
const SAMPLE_STREAM: u64 = 0x5EED_0004;
const CHUNK_PAIRS: usize = 1 << 14;
#[derive(Debug, Clone, PartialEq)]
pub struct Samples {
values: Vec<f32>,
n_rows: usize,
per_row: usize,
n_outputs: usize,
}
fn check_layout(values: &[f32], n_samples: usize, n_outputs: usize) -> Result<usize> {
if n_samples == 0 || n_outputs == 0 {
return Err(HessboostError::invalid_data(
"samples",
format!("needs at least one sample and one output, got {n_samples} and {n_outputs}"),
));
}
let Some(width) = n_samples.checked_mul(n_outputs) else {
return Err(HessboostError::invalid_data(
"samples",
format!("{n_samples} samples × {n_outputs} outputs overflows usize"),
));
};
if !values.len().is_multiple_of(width) {
return Err(HessboostError::invalid_data(
"samples",
format!(
"{} draws are not whole rows of {n_samples} samples × {n_outputs} outputs",
values.len()
),
));
}
if !values.iter().all(|v| v.is_finite()) {
return Err(HessboostError::invalid_data(
"samples",
"all draws must be finite",
));
}
Ok(values.len() / width)
}
impl Samples {
pub fn new(values: Vec<f32>, n_samples: usize, n_outputs: usize) -> Result<Self> {
let n_rows = check_layout(&values, n_samples, n_outputs)?;
Ok(Samples {
values,
n_rows,
per_row: n_samples,
n_outputs,
})
}
pub fn view(&self) -> SamplesView<'_> {
SamplesView {
values: &self.values,
n_rows: self.n_rows,
per_row: self.per_row,
n_outputs: self.n_outputs,
}
}
pub fn as_slice(&self) -> &[f32] {
&self.values
}
pub fn into_vec(self) -> Vec<f32> {
self.values
}
pub fn n_rows(&self) -> usize {
self.n_rows
}
pub fn n_samples(&self) -> usize {
self.per_row
}
pub fn n_outputs(&self) -> usize {
self.n_outputs
}
pub fn row(&self, row: usize) -> Option<&[f32]> {
self.view().row(row)
}
pub fn get(&self, row: usize, sample: usize) -> Option<&[f32]> {
self.view().get(row, sample)
}
pub fn mean(&self) -> Predictions<f64> {
self.view().mean()
}
pub fn quantiles(&self, levels: &[f64]) -> Result<Quantiles> {
self.view().quantiles(levels)
}
pub fn crps(&self, labels: &[f32]) -> Result<Predictions<f64>> {
self.view().crps(labels)
}
}
impl AsRef<[f32]> for Samples {
fn as_ref(&self) -> &[f32] {
&self.values
}
}
#[derive(Debug, Clone, Copy, PartialEq)]
pub struct SamplesView<'a> {
values: &'a [f32],
n_rows: usize,
per_row: usize,
n_outputs: usize,
}
impl<'a> SamplesView<'a> {
pub fn new(values: &'a [f32], n_samples: usize, n_outputs: usize) -> Result<Self> {
let n_rows = check_layout(values, n_samples, n_outputs)?;
Ok(SamplesView {
values,
n_rows,
per_row: n_samples,
n_outputs,
})
}
pub fn as_slice(&self) -> &'a [f32] {
self.values
}
pub fn n_rows(&self) -> usize {
self.n_rows
}
pub fn n_samples(&self) -> usize {
self.per_row
}
pub fn n_outputs(&self) -> usize {
self.n_outputs
}
pub fn row(&self, row: usize) -> Option<&'a [f32]> {
let width = self.per_row * self.n_outputs;
let start = row.checked_mul(width)?;
self.values.get(start..start.checked_add(width)?)
}
pub fn get(&self, row: usize, sample: usize) -> Option<&'a [f32]> {
if sample >= self.per_row {
return None;
}
let draws = self.row(row)?;
Some(&draws[sample * self.n_outputs..(sample + 1) * self.n_outputs])
}
pub fn mean(&self) -> Predictions<f64> {
let mut out = vec![0.0; self.n_rows * self.n_outputs];
let width = self.per_row * self.n_outputs;
for (mean, draws) in out
.chunks_exact_mut(self.n_outputs)
.zip(self.values.chunks_exact(width))
{
for draw in draws.chunks_exact(self.n_outputs) {
for (m, &v) in mean.iter_mut().zip(draw) {
*m += f64::from(v);
}
}
for m in mean {
*m /= self.per_row as f64;
}
}
Predictions::new(out, self.n_rows, self.n_outputs)
}
pub fn quantiles(&self, levels: &[f64]) -> Result<Quantiles> {
if let Some(&level) = levels.iter().find(|l| !(0.0..=1.0).contains(*l)) {
return Err(HessboostError::invalid_param(
"levels",
format!("quantile levels must be in [0, 1], got {level}"),
));
}
let (k, d) = (levels.len(), self.n_outputs);
let len = self
.n_rows
.checked_mul(k)
.and_then(|n| n.checked_mul(d))
.ok_or_else(|| {
HessboostError::invalid_param("levels", "the quantile table overflows usize")
})?;
let mut out = try_filled(len, 0.0, "levels")?;
let mut column = Vec::new();
for row in 0..self.n_rows {
for o in 0..d {
self.sorted_column(row, o, &mut column);
for (l, &level) in levels.iter().enumerate() {
out[(row * k + l) * d + o] = quantile_sorted(&column, level);
}
}
}
Ok(Quantiles {
values: out,
n_rows: self.n_rows,
n_levels: k,
n_outputs: d,
})
}
pub fn crps(&self, labels: &[f32]) -> Result<Predictions<f64>> {
let d = self.n_outputs;
if labels.len() != self.n_rows * d {
return Err(HessboostError::DimensionMismatch {
what: "crps labels",
expected: self.n_rows * d,
got: labels.len(),
});
}
if labels.iter().any(|v| !v.is_finite()) {
return Err(HessboostError::invalid_data(
"labels",
"all labels must be finite",
));
}
let s = self.per_row as f64;
let mut column = Vec::new();
let mut out = vec![0.0; self.n_rows * d];
for row in 0..self.n_rows {
for o in 0..d {
self.sorted_column(row, o, &mut column);
let y = f64::from(labels[row * d + o]);
let spread_to_label = column.iter().map(|x| (x - y).abs()).sum::<f64>() / s;
let pairwise: f64 = column
.iter()
.enumerate()
.map(|(i, x)| (2.0 * i as f64 - s + 1.0) * x)
.sum::<f64>()
* 2.0;
out[row * d + o] = spread_to_label - pairwise / (2.0 * s * s);
}
}
Ok(Predictions::new(out, self.n_rows, d))
}
fn sorted_column(&self, row: usize, output: usize, column: &mut Vec<f64>) {
let base = row * self.per_row * self.n_outputs + output;
column.clear();
column.extend((0..self.per_row).map(|s| f64::from(self.values[base + s * self.n_outputs])));
column.sort_unstable_by(f64::total_cmp);
}
}
#[derive(Debug, Clone, PartialEq)]
pub struct Quantiles {
values: Vec<f64>,
n_rows: usize,
n_levels: usize,
n_outputs: usize,
}
impl Quantiles {
pub fn n_rows(&self) -> usize {
self.n_rows
}
pub fn n_levels(&self) -> usize {
self.n_levels
}
pub fn n_outputs(&self) -> usize {
self.n_outputs
}
pub fn get(&self, row: usize, level: usize) -> Option<&[f64]> {
if row >= self.n_rows || level >= self.n_levels {
return None;
}
let start = (row * self.n_levels + level) * self.n_outputs;
Some(&self.values[start..start + self.n_outputs])
}
pub fn as_slice(&self) -> &[f64] {
&self.values
}
pub fn into_vec(self) -> Vec<f64> {
self.values
}
}
impl AsRef<[f64]> for Quantiles {
fn as_ref(&self) -> &[f64] {
&self.values
}
}
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq)]
#[non_exhaustive]
pub struct SampleOptions {
pub seed: u64,
pub n_steps: Option<NonZeroUsize>,
}
impl SampleOptions {
pub fn seeded(seed: u64) -> Self {
SampleOptions {
seed,
n_steps: None,
}
}
#[must_use]
pub fn with_n_steps(mut self, n_steps: NonZeroUsize) -> Self {
self.n_steps = Some(n_steps);
self
}
}
pub(super) fn sample(
model: &DiffusionModel,
data: &DMatrix,
n_samples: usize,
options: &SampleOptions,
) -> Result<Samples> {
crate::check::ensure("n_samples", n_samples != 0, "must be at least 1")?;
if data.n_cols() != model.n_features {
return Err(HessboostError::DimensionMismatch {
what: "sampling feature count",
expected: model.n_features,
got: data.n_cols(),
});
}
if data.base_margin().is_some() {
return Err(HessboostError::invalid_data(
"base_margin",
"diffusion models do not support base margins",
));
}
let (n_rows, d) = (data.n_rows(), model.n_outputs);
let n_pairs = n_rows.checked_mul(n_samples);
let Some(total) = n_pairs.and_then(|pairs| pairs.checked_mul(d)) else {
return Err(HessboostError::invalid_param(
"n_samples",
"rows × samples × outputs overflows usize",
));
};
let mean = match &model.residualizer {
Some(r) => Some(residual_mean(&r.models, data)?),
None => None,
};
let sampler = Sampler {
model,
features: dense_features(data),
mean,
n_samples,
n_steps: options.n_steps.unwrap_or(model.n_steps).get(),
key: splitmix64(options.seed ^ SAMPLE_STREAM),
};
let mut values = try_filled(total, 0.0f32, "n_samples")?;
let n_pairs = n_rows * n_samples;
for (chunk, out) in values.chunks_mut(CHUNK_PAIRS * d).enumerate() {
let start = chunk * CHUNK_PAIRS;
sampler.chunk(start..(start + CHUNK_PAIRS).min(n_pairs), out)?;
}
Ok(Samples {
values,
n_rows,
per_row: n_samples,
n_outputs: d,
})
}
struct Sampler<'a> {
model: &'a DiffusionModel,
features: Vec<f32>,
mean: Option<Vec<f64>>,
n_samples: usize,
n_steps: usize,
key: u64,
}
struct Batch {
input: DMatrix,
streams: Vec<u64>,
cols: usize,
}
impl Batch {
fn predict(
&mut self,
model: &DiffusionModel,
y: &[f64],
scale: f64,
time: [f32; 2],
) -> Result<Vec<f32>> {
let d = model.n_outputs;
let t_col = d + model.n_features;
let values = self
.input
.dense_values_mut()
.ok_or_else(|| HessboostError::model_format("sampling input must be dense"))?;
for (row, state) in values.chunks_exact_mut(self.cols).zip(y.chunks_exact(d)) {
for (v, &s) in row.iter_mut().zip(state) {
*v = (s * scale) as f32;
}
row[t_col..].copy_from_slice(&time[..self.cols - t_col]);
}
Ok(model
.regressor
.predict_margin(&self.input, Iterations::Best)?
.into_vec())
}
}
impl Sampler<'_> {
fn chunk(&self, pairs: std::ops::Range<usize>, out: &mut [f32]) -> Result<()> {
let model = self.model;
let (d, p) = (model.n_outputs, model.n_features);
let cols = d + p + model.method.time_columns();
let m = pairs.len();
let mut input = vec![0.0f32; m * cols];
let mut streams = Vec::with_capacity(m);
for (pair, row_values) in pairs.clone().zip(input.chunks_exact_mut(cols)) {
let (row, sample) = (pair / self.n_samples, pair % self.n_samples);
row_values[d..d + p].copy_from_slice(&self.features[row * p..(row + 1) * p]);
streams.push(splitmix64(
splitmix64(self.key ^ row as u64) ^ sample as u64,
));
}
let mut batch = Batch {
input: DMatrix::from_dense_vec(input, m, cols)?,
streams,
cols,
};
let prior_std = match model.method {
Method::Score(score) => score.sde.prior_std(),
Method::FlowMatching(_) => 1.0,
};
let mut y = vec![0.0; m * d];
for (state, &stream) in y.chunks_exact_mut(d).zip(&batch.streams) {
for (j, v) in state.iter_mut().enumerate() {
*v = prior_std * keyed_normal(stream, j as u64);
}
}
match model.method {
Method::Score(score) => self.reverse_sde(&score, &mut batch, &mut y)?,
Method::FlowMatching(flow) => self.reverse_ode(flow.solver, &mut batch, &mut y)?,
}
for ((pair, state), draws) in pairs.zip(y.chunks_exact(d)).zip(out.chunks_exact_mut(d)) {
let row = pair / self.n_samples;
for (j, (&u, v)) in state.iter().zip(draws).enumerate() {
let standardized = match (&model.residualizer, &self.mean) {
(Some(r), Some(mean)) => mean[row * d + j] + u * r.scale[j] + r.center[j],
_ => u,
};
let value = (standardized * model.target_scale[j] + model.target_mean[j]) as f32;
if !value.is_finite() {
return Err(diverged());
}
*v = value;
}
}
Ok(())
}
fn reverse_sde(&self, score: &ScoreConfig, batch: &mut Batch, y: &mut [f64]) -> Result<()> {
let model = self.model;
let d = model.n_outputs;
let steps = self.n_steps;
let dt = (1.0 - T_EPS) / steps as f64;
for step in 0..steps {
let t = 1.0 - step as f64 * dt;
let (alpha, std) = score.sde.marginal(t);
let (c, g2) = score.sde.drift_diffusion(t);
let edm = match score.parameterization {
Parameterization::Noise => None,
Parameterization::Edm { sigma_data } => Some(Edm::at(sigma_data, std)),
};
let input_scale = edm.map_or(1.0, |e| e.input);
let pred = batch.predict(model, y, input_scale, [t as f32, std.ln() as f32])?;
let noise_scale = (g2 * dt).sqrt();
let counter = (step as u64 + 1) * d as u64;
for ((state, out), &stream) in y
.chunks_exact_mut(d)
.zip(pred.chunks_exact(d))
.zip(&batch.streams)
{
for (j, (v, &f)) in state.iter_mut().zip(out).enumerate() {
let f = f64::from(f);
let score = match edm {
None => f / std,
Some(e) => (alpha * (e.skip * *v + e.out * f) - *v) / (std * std),
};
let drift = -c * *v + g2 * score;
*v += drift * dt + noise_scale * keyed_normal(stream, counter + j as u64);
}
}
check_finite(y)?;
}
Ok(())
}
fn reverse_ode(&self, solver: OdeSolver, batch: &mut Batch, y: &mut [f64]) -> Result<()> {
let model = self.model;
let steps = self.n_steps;
let ds = (1.0 - T_EPS) / steps as f64;
let time = |t: f64| [t as f32, 0.0];
let mut predictor = match solver {
OdeSolver::Euler => Vec::new(),
OdeSolver::Heun => vec![0.0; y.len()],
};
for step in 0..steps {
let t = 1.0 - step as f64 * ds;
let v0 = batch.predict(model, y, 1.0, time(t))?;
match solver {
OdeSolver::Euler => {
for (s, &v) in y.iter_mut().zip(&v0) {
*s -= f64::from(v) * ds;
}
}
OdeSolver::Heun => {
for ((pred, &s), &v) in predictor.iter_mut().zip(y.iter()).zip(&v0) {
*pred = s - f64::from(v) * ds;
}
check_finite(&predictor)?;
let v1 = batch.predict(model, &predictor, 1.0, time(t - ds))?;
for ((s, &a), &b) in y.iter_mut().zip(&v0).zip(&v1) {
*s -= f64::midpoint(f64::from(a), f64::from(b)) * ds;
}
}
}
check_finite(y)?;
}
Ok(())
}
}
fn check_finite(y: &[f64]) -> Result<()> {
if y.iter().all(|v| v.abs() <= f64::from(f32::MAX)) {
Ok(())
} else {
Err(diverged())
}
}
fn diverged() -> HessboostError {
HessboostError::invalid_param(
"n_steps",
"sampling diverged to non-finite values; use more steps",
)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn crps_and_quantiles_match_their_definitions() {
let values = vec![0.5, 3.0, -1.0, 2.0, 0.0, 1.0, -2.0, 4.0];
let samples = Samples {
values: values.clone(),
n_rows: 2,
per_row: 2,
n_outputs: 2,
};
let labels = [1.0, -0.5, 0.25, 2.0];
let crps = samples.crps(&labels).unwrap();
for row in 0..2 {
for o in 0..2 {
let draws: Vec<f64> = (0..2)
.map(|s| f64::from(values[(row * 2 + s) * 2 + o]))
.collect();
let y = f64::from(labels[row * 2 + o]);
let to_label = draws.iter().map(|x| (x - y).abs()).sum::<f64>() / 2.0;
let pairwise = draws
.iter()
.flat_map(|a| draws.iter().map(move |b| (a - b).abs()))
.sum::<f64>()
/ 8.0;
assert!((crps.get(row, o).unwrap() - (to_label - pairwise)).abs() < 1e-12);
}
}
let q = samples.quantiles(&[0.0, 0.5, 1.0]).unwrap();
assert_eq!(&q.as_slice()[..6], &[-1.0, 2.0, -0.25, 2.5, 0.5, 3.0]);
assert_eq!(samples.mean().row(0).unwrap(), [-0.25, 2.5]);
assert!(samples.quantiles(&[1.5]).is_err());
assert!(samples.crps(&[0.0; 3]).is_err());
}
}