use crate::model::Iterations;
use std::num::NonZeroUsize;
use serde::{Deserialize, Serialize};
use super::process::{keyed_normal, try_filled};
use super::{DiffusionFormat, Sde};
use super::{check_regressor, validate_regressor_params};
use crate::check::{ensure, fraction};
use crate::config::{TrainingParams, TreeMethod};
use crate::data::DMatrix;
use crate::error::{HessboostError, Result};
use crate::model::BoostedModel;
use crate::rng::{keyed_unit, splitmix64};
use encoding::encode_row;
mod encoding;
mod fit;
mod format;
const EPS: f64 = 1e-3;
fn level_time(n_t: usize, level: usize) -> f64 {
EPS + (1.0 - EPS) * level as f64 / (n_t - 1) as f64
}
const NOISE_STREAM: u64 = 0xF0E5_0001;
const PRIOR_STREAM: u64 = 0xF0E5_0002;
const STEP_STREAM: u64 = 0xF0E5_0003;
const KNOWN_STREAM: u64 = 0xF0E5_0004;
const REPAINT_STREAM: u64 = 0xF0E5_0005;
const LABEL_STREAM: u64 = 0xF0E5_0006;
#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Hash, Serialize, Deserialize)]
#[serde(try_from = "usize", into = "usize")]
pub struct NoiseLevels(usize);
impl NoiseLevels {
pub const fn new(n: usize) -> Option<Self> {
if n >= 2 { Some(NoiseLevels(n)) } else { None }
}
pub const fn get(self) -> usize {
self.0
}
}
impl TryFrom<usize> for NoiseLevels {
type Error = HessboostError;
fn try_from(n: usize) -> Result<Self> {
NoiseLevels::new(n).ok_or_else(|| {
HessboostError::invalid_param("n_t", format!("needs at least 2 noise levels, got {n}"))
})
}
}
impl From<NoiseLevels> for usize {
fn from(n: NoiseLevels) -> usize {
n.0
}
}
#[derive(Debug, Clone, Copy, PartialEq, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
#[non_exhaustive]
pub enum ForestMethod {
Flow,
Diffusion {
beta_min: f64,
beta_max: f64,
},
}
impl ForestMethod {
pub fn forest_diffusion() -> Self {
ForestMethod::Diffusion {
beta_min: 0.1,
beta_max: 8.0,
}
}
fn sde(self) -> Option<Sde> {
match self {
ForestMethod::Flow => None,
ForestMethod::Diffusion { beta_min, beta_max } => {
Some(Sde::VariancePreserving { beta_min, beta_max })
}
}
}
fn validate(self) -> Result<()> {
if let Some(sde) = self.sde() {
super::Method::Score(super::ScoreConfig {
sde,
..super::ScoreConfig::treeffuser()
})
.validate()
.map_err(|e| match e {
HessboostError::InvalidParameter { reason, .. } => {
HessboostError::invalid_param("method", reason)
}
other => other,
})?;
}
Ok(())
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
#[non_exhaustive]
pub enum ColumnKind {
Continuous,
Integer,
Categorical,
}
#[derive(Debug, Clone, Copy, Default, PartialEq)]
#[non_exhaustive]
pub struct ImputeOptions {
pub seed: u64,
pub repaint: Option<Repaint>,
}
impl ImputeOptions {
pub fn seeded(seed: u64) -> Self {
ImputeOptions {
seed,
repaint: None,
}
}
#[must_use]
pub fn with_repaint(mut self, repaint: Repaint) -> Self {
self.repaint = Some(repaint);
self
}
}
#[derive(Debug, Clone, Copy, PartialEq)]
#[non_exhaustive]
pub struct Repaint {
pub resample: NonZeroUsize,
pub jump: f64,
}
impl Default for Repaint {
fn default() -> Self {
Repaint {
resample: const { NonZeroUsize::new(5).unwrap() },
jump: 0.1,
}
}
}
#[derive(Debug, Clone)]
#[non_exhaustive]
pub struct ForestParams {
pub method: ForestMethod,
pub n_t: NoiseLevels,
pub duplicate_k: NonZeroUsize,
pub column_kinds: Option<Vec<ColumnKind>>,
pub training: TrainingParams,
pub num_boost_round: NonZeroUsize,
pub seed: u64,
}
impl Default for ForestParams {
fn default() -> Self {
ForestParams {
method: ForestMethod::Flow,
n_t: const { NoiseLevels::new(50).unwrap() },
duplicate_k: const { NonZeroUsize::new(100).unwrap() },
column_kinds: None,
training: TrainingParams {
tree_method: TreeMethod::Hist,
max_depth: NonZeroUsize::new(7),
eta: 0.3,
lambda: 0.0,
..TrainingParams::default()
},
num_boost_round: const { NonZeroUsize::new(100).unwrap() },
seed: 0,
}
}
}
impl ForestParams {
pub fn forest_flow() -> Self {
ForestParams::default()
}
pub fn forest_diffusion() -> Self {
ForestParams {
method: ForestMethod::forest_diffusion(),
..ForestParams::default()
}
}
pub fn validate(&self) -> Result<()> {
self.method.validate()?;
validate_regressor_params("training", &self.training)
}
}
#[derive(Debug, Clone, Copy, PartialEq, Serialize, Deserialize)]
struct Scale {
min: f64,
range: f64,
}
impl Scale {
fn forward(self, v: f64) -> f64 {
(v - self.min) * 2.0 / self.range - 1.0
}
fn inverse(self, v: f64) -> f64 {
(v + 1.0) * self.range / 2.0 + self.min
}
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
struct Column {
kind: ColumnKind,
min: f64,
max: f64,
categories: Vec<f64>,
}
impl Column {
fn width(&self) -> usize {
match self.kind {
ColumnKind::Categorical => self.categories.len().saturating_sub(1),
ColumnKind::Continuous | ColumnKind::Integer => 1,
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
#[serde(from = "bool", into = "bool")]
enum OutputLayout {
Joint,
PerColumn,
}
impl OutputLayout {
fn gbdts_per_level(self, c: usize) -> usize {
match self {
OutputLayout::Joint => 1,
OutputLayout::PerColumn => c,
}
}
fn gbdt_outputs(self, c: usize) -> usize {
match self {
OutputLayout::Joint => c,
OutputLayout::PerColumn => 1,
}
}
}
impl From<bool> for OutputLayout {
fn from(per_output: bool) -> Self {
if per_output {
OutputLayout::PerColumn
} else {
OutputLayout::Joint
}
}
}
impl From<OutputLayout> for bool {
fn from(layout: OutputLayout) -> bool {
layout == OutputLayout::PerColumn
}
}
#[derive(Debug, Clone, PartialEq)]
pub struct Synthetic {
values: Vec<f32>,
labels: Option<Vec<f32>>,
n_columns: usize,
}
impl Synthetic {
pub fn row(&self, row: usize) -> Option<&[f32]> {
let start = row.checked_mul(self.n_columns)?;
self.values.get(start..start.checked_add(self.n_columns)?)
}
pub fn rows(&self) -> impl ExactSizeIterator<Item = &[f32]> {
self.values.chunks_exact(self.n_columns)
}
pub fn as_slice(&self) -> &[f32] {
&self.values
}
pub fn into_parts(self) -> (Vec<f32>, Option<Vec<f32>>) {
(self.values, self.labels)
}
pub fn labels(&self) -> Option<&[f32]> {
self.labels.as_deref()
}
pub fn n_rows(&self) -> usize {
self.values.len() / self.n_columns
}
pub fn n_columns(&self) -> usize {
self.n_columns
}
pub fn to_dmatrix(&self) -> Result<DMatrix> {
let m = DMatrix::from_dense(&self.values, self.n_rows(), self.n_columns)?;
match &self.labels {
Some(labels) => m.with_labels(labels),
None => Ok(m),
}
}
}
#[derive(Debug, Clone, PartialEq)]
pub struct Imputations {
values: Vec<f32>,
draws: usize,
n_rows: usize,
n_columns: usize,
}
impl Imputations {
pub fn n_imputations(&self) -> usize {
self.draws
}
pub fn n_rows(&self) -> usize {
self.n_rows
}
pub fn n_columns(&self) -> usize {
self.n_columns
}
pub fn get(&self, imputation: usize, row: usize) -> Option<&[f32]> {
if imputation >= self.draws || row >= self.n_rows {
return None;
}
let start = (imputation * self.n_rows + row) * self.n_columns;
Some(&self.values[start..start + self.n_columns])
}
pub fn as_slice(&self) -> &[f32] {
&self.values
}
pub fn into_vec(self) -> Vec<f32> {
self.values
}
}
impl AsRef<[f32]> for Imputations {
fn as_ref(&self) -> &[f32] {
&self.values
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(try_from = "format::UncheckedForestModel")]
pub struct ForestModel {
method: ForestMethod,
n_t: NoiseLevels,
columns: Vec<Column>,
scales: Vec<Scale>,
classes: Vec<f64>,
class_probs: Vec<f64>,
#[serde(rename = "per_output")]
layout: OutputLayout,
models: Vec<BoostedModel>,
}
impl ForestModel {
pub fn fit(params: &ForestParams, data: &DMatrix) -> Result<Self> {
fit::fit(params, data)
}
pub fn sample(&self, n_rows: usize, seed: u64) -> Result<Synthetic> {
positive_count("n_rows", n_rows)?;
let mut class_of: Vec<usize> = try_filled(n_rows, 0, "n_rows")?;
if !self.classes.is_empty() {
let key = splitmix64(seed ^ LABEL_STREAM);
for (r, class) in class_of.iter_mut().enumerate() {
let u = keyed_unit(key, r as u64);
let mut acc = 0.0;
*class = self
.class_probs
.iter()
.position(|&p| {
acc += p;
u < acc
})
.unwrap_or(self.classes.len() - 1);
}
}
self.sample_classes(&class_of, seed)
}
pub fn sample_for_labels(&self, labels: &[f32], seed: u64) -> Result<Synthetic> {
positive_count("labels", labels.len())?;
let class_of = self.class_indices(labels)?;
self.sample_classes(&class_of, seed)
}
pub fn impute(
&self,
data: &DMatrix,
n_imputations: usize,
options: &ImputeOptions,
) -> Result<Imputations> {
let ImputeOptions { seed, repaint } = *options;
let Some(sde) = self.method.sde() else {
return Err(HessboostError::incompatible_model(
"impute",
"imputation needs ForestMethod::Diffusion (flow matching cannot condition \
on the observed entries)",
));
};
positive_count("n_imputations", n_imputations)?;
let (resample, jump) = match repaint {
None => (1, self.n_t.get()),
Some(r) => {
fraction("repaint.jump", r.jump)?;
(
r.resample.get(),
((r.jump * self.n_t.get() as f64).ceil() as usize).max(1),
)
}
};
let p = self.columns.len();
if data.n_cols() != p {
return Err(HessboostError::dimension_mismatch(
"imputation column count",
p,
data.n_cols(),
));
}
refuse_metadata(data)?;
let class_of = if self.classes.is_empty() {
vec![0; data.n_rows()]
} else {
let labels = data.labels().ok_or_else(|| {
HessboostError::invalid_data(
"labels",
"a class-conditional model imputes rows of known class: attach labels",
)
})?;
self.class_indices(labels)?
};
let raw = super::fit::dense_features(data);
let c = self.scales.len();
let mut known = try_filled(data.n_rows() * c, 0.0, "data")?;
for (row, out) in raw.chunks_exact(p).zip(known.chunks_exact_mut(c)) {
for (j, (column, &v)) in self.columns.iter().zip(row).enumerate() {
let v = f64::from(v);
if column.kind == ColumnKind::Categorical
&& !v.is_nan()
&& column
.categories
.binary_search_by(|x| x.total_cmp(&v))
.is_err()
{
return Err(HessboostError::invalid_data(
"data",
format!("column {j} has category {v}, unseen in training"),
));
}
}
encode_row(&self.columns, row, out);
for (v, s) in out.iter_mut().zip(&self.scales) {
*v = s.forward(*v);
}
}
let n_rows = data.n_rows();
let total = n_imputations
.checked_mul(n_rows)
.and_then(|m| m.checked_mul(p))
.ok_or_else(|| {
HessboostError::invalid_param("n_imputations", "the imputed table overflows usize")
})?;
let mut out = try_filled(total, 0.0f32, "n_imputations")?;
let mut sampler = Sampler::new(self, &class_of, 0)?;
for (i, dest) in out.chunks_exact_mut(n_rows * p).enumerate() {
sampler.key = splitmix64(seed ^ splitmix64(i as u64));
let x = sampler.reverse_sde(sde, Some(&known), (resample, jump))?;
self.decode_rows(&x, dest)?;
}
Ok(Imputations {
values: out,
draws: n_imputations,
n_rows,
n_columns: p,
})
}
pub fn method(&self) -> ForestMethod {
self.method
}
pub fn n_t(&self) -> NoiseLevels {
self.n_t
}
pub fn n_columns(&self) -> usize {
self.columns.len()
}
pub fn classes(&self) -> &[f64] {
&self.classes
}
pub fn encode(&self, format: DiffusionFormat) -> Result<Vec<u8>> {
match format {
DiffusionFormat::Binary => format::write(self),
DiffusionFormat::Json => Ok(serde_json::to_vec_pretty(self)?),
}
}
pub fn decode(bytes: impl AsRef<[u8]>, format: DiffusionFormat) -> Result<Self> {
let bytes = bytes.as_ref();
match format {
DiffusionFormat::Binary => format::read(bytes),
DiffusionFormat::Json => Ok(serde_json::from_slice(bytes)?),
}
}
pub fn save(&self, path: impl AsRef<std::path::Path>, format: DiffusionFormat) -> Result<()> {
Ok(std::fs::write(path, self.encode(format)?)?)
}
pub fn load(path: impl AsRef<std::path::Path>, format: DiffusionFormat) -> Result<Self> {
Self::decode(std::fs::read(path)?, format)
}
pub fn gbdts(&self) -> &[BoostedModel] {
&self.models
}
fn class_indices(&self, labels: &[f32]) -> Result<Vec<usize>> {
if self.classes.is_empty() {
return Err(HessboostError::incompatible_model(
"labels",
"the model was fitted without class labels",
));
}
labels
.iter()
.map(|&l| {
let l = f64::from(l);
self.classes
.binary_search_by(|c| c.total_cmp(&l))
.map_err(|_| {
HessboostError::invalid_data(
"labels",
format!("label {l} was not a training class"),
)
})
})
.collect()
}
fn sample_classes(&self, class_of: &[usize], seed: u64) -> Result<Synthetic> {
let n_rows = class_of.len();
let p = self.columns.len();
let total = n_rows.checked_mul(p).ok_or_else(|| {
HessboostError::invalid_param("n_rows", "the synthetic table overflows usize")
})?;
let mut values = try_filled(total, 0.0f32, "n_rows")?;
let mut sampler = Sampler::new(self, class_of, seed)?;
let x = match self.method.sde() {
None => sampler.euler_flow()?,
Some(sde) => sampler.reverse_sde(sde, None, (1, self.n_t.get()))?,
};
self.decode_rows(&x, &mut values)?;
let labels = (!self.classes.is_empty())
.then(|| class_of.iter().map(|&k| self.classes[k] as f32).collect());
Ok(Synthetic {
values,
labels,
n_columns: p,
})
}
fn decode_rows(&self, x: &[f64], out: &mut [f32]) -> Result<()> {
let (c, p) = (self.scales.len(), self.columns.len());
for (row, dest) in x.chunks_exact(c).zip(out.chunks_exact_mut(p)) {
let mut at = 0;
for (column, d) in self.columns.iter().zip(dest.iter_mut()) {
let width = column.width();
let value = match column.kind {
ColumnKind::Categorical => {
let mut best = (0, 0.5);
for m in 0..width {
let v = self.scales[at + m].inverse(row[at + m]);
if v > best.1 {
best = (m + 1, v);
}
}
column.categories.get(best.0).copied().unwrap_or(column.min)
}
ColumnKind::Integer => self.scales[at].inverse(row[at]).round_ties_even(),
ColumnKind::Continuous => self.scales[at].inverse(row[at]),
};
at += width;
let v = value.max(column.min).min(column.max) as f32;
if !v.is_finite() {
return Err(diverged());
}
*d = v;
}
}
Ok(())
}
fn level_models(&self, class: usize, level: usize) -> &[BoostedModel] {
let per = self.layout.gbdts_per_level(self.scales.len());
let start = (class * self.n_t.get() + level) * per;
&self.models[start..start + per]
}
fn validate(&self) -> Result<()> {
let bad = |msg: String| Err(HessboostError::model_format(msg));
self.method
.validate()
.map_err(|e| HessboostError::model_format(e.to_string()))?;
if self.columns.is_empty() {
return bad("a forest model needs a column".into());
}
let c: usize = self.columns.iter().map(Column::width).sum();
if c == 0 || self.scales.len() != c {
return bad(format!(
"{} scales for {c} encoded columns",
self.scales.len()
));
}
for column in &self.columns {
let sorted = column.categories.windows(2).all(|w| w[0] < w[1]);
let valid = f32_exact(column.min)
&& f32_exact(column.max)
&& column.min <= column.max
&& column.categories.iter().all(|&v| f32_exact(v))
&& sorted
&& (column.kind == ColumnKind::Categorical) != column.categories.is_empty();
if !valid {
return bad("a column's range or categories are invalid".into());
}
}
if !self
.scales
.iter()
.all(|s| s.min.is_finite() && s.range.is_finite() && s.range > 0.0)
{
return bad("an encoded column's scale is invalid".into());
}
let classes_sorted = self.classes.windows(2).all(|w| w[0] < w[1]);
if !classes_sorted
|| self.classes.len() != self.class_probs.len()
|| !self.classes.iter().all(|&v| f32_exact(v))
|| !self.class_probs.iter().all(|p| p.is_finite() && *p >= 0.0)
{
return bad("the classes are invalid".into());
}
let per = self.layout.gbdts_per_level(c);
let expected = self
.classes
.len()
.max(1)
.checked_mul(self.n_t.get())
.and_then(|m| m.checked_mul(per));
if expected != Some(self.models.len()) {
return bad(format!(
"{} GBDTs, expected one per class, level{}",
self.models.len(),
match self.layout {
OutputLayout::Joint => "",
OutputLayout::PerColumn => " and column",
}
));
}
let outputs = self.layout.gbdt_outputs(c);
for model in &self.models {
check_regressor("forest GBDT", model, c, outputs)?;
}
Ok(())
}
}
fn refuse_metadata(data: &DMatrix) -> Result<()> {
super::fit::refuse_unsupported_metadata(data, "forest")?;
if data.labels().is_some() && data.n_targets() != 1 {
return Err(HessboostError::invalid_data(
"labels",
"forest models do not support label matrices (labels are class labels)",
));
}
Ok(())
}
struct Sampler<'a> {
model: &'a ForestModel,
key: u64,
n_rows: usize,
batches: Vec<ClassBatch>,
}
struct ClassBatch {
class: usize,
rows: Vec<usize>,
input: DMatrix,
}
impl<'a> Sampler<'a> {
fn new(model: &'a ForestModel, class_of: &[usize], key: u64) -> Result<Self> {
let c = model.scales.len();
let mut rows = vec![Vec::new(); model.classes.len().max(1)];
for (r, &class) in class_of.iter().enumerate() {
rows[class].push(r);
}
let batches = rows
.into_iter()
.enumerate()
.filter(|(_, rows)| !rows.is_empty())
.map(|(class, rows)| {
let values = try_filled(rows.len() * c, 0.0f32, "n_rows")?;
let input = DMatrix::from_dense_vec(values, rows.len(), c)?;
Ok(ClassBatch { class, rows, input })
})
.collect::<Result<_>>()?;
Ok(Sampler {
model,
key,
n_rows: class_of.len(),
batches,
})
}
fn noise(&self, stream: u64, row: usize, counter: u64) -> f64 {
keyed_normal(
splitmix64(splitmix64(self.key ^ stream) ^ row as u64),
counter,
)
}
fn prior(&self) -> Result<Vec<f64>> {
let c = self.model.scales.len();
let mut x = try_filled(self.n_rows * c, 0.0, "n_rows")?;
for (row, values) in x.chunks_exact_mut(c).enumerate() {
for (j, v) in values.iter_mut().enumerate() {
*v = self.noise(PRIOR_STREAM, row, j as u64);
}
}
Ok(x)
}
fn predict(&mut self, x: &[f64], level: usize) -> Result<Vec<f64>> {
let model = self.model;
let c = model.scales.len();
let mut out = vec![0.0; x.len()];
for batch in &mut self.batches {
let values = batch
.input
.dense_values_mut()
.ok_or_else(|| HessboostError::model_format("sampling input must be dense"))?;
for (&r, dest) in batch.rows.iter().zip(values.chunks_exact_mut(c)) {
for (d, &v) in dest.iter_mut().zip(&x[r * c..(r + 1) * c]) {
*d = v as f32;
}
}
if values.iter().any(|v| v.is_infinite()) {
return Err(diverged());
}
let models = model.level_models(batch.class, level);
for (m, gbdt) in models.iter().enumerate() {
let pred = gbdt.predict_margin(&batch.input, Iterations::Best)?;
for (&r, p) in batch.rows.iter().zip(pred.rows()) {
match model.layout {
OutputLayout::PerColumn => out[r * c + m] = f64::from(p[0]),
OutputLayout::Joint => {
for (o, &v) in out[r * c..(r + 1) * c].iter_mut().zip(p) {
*o = f64::from(v);
}
}
}
}
}
}
Ok(out)
}
fn euler_flow(&mut self) -> Result<Vec<f64>> {
let n_t = self.model.n_t.get();
let h = 1.0 / (n_t - 1) as f64;
let mut x = self.prior()?;
for step in 0..n_t - 1 {
let v = self.predict(&x, step)?;
for (s, &d) in x.iter_mut().zip(&v) {
*s += h * d;
}
check_finite(&x)?;
}
Ok(x)
}
fn score(
&mut self,
sde: Sde,
x: &mut [f64],
known: Option<&[f64]>,
level: usize,
eval: u64,
) -> Result<Vec<f64>> {
let (alpha, std) = sde.marginal(level_time(self.model.n_t.get(), level));
if let Some(known) = known {
let c = self.model.scales.len();
for (row, (state, obs)) in x.chunks_exact_mut(c).zip(known.chunks_exact(c)).enumerate()
{
for (j, (s, &o)) in state.iter_mut().zip(obs).enumerate() {
if !o.is_nan() {
*s = alpha * o
+ std * self.noise(KNOWN_STREAM, row, eval * c as u64 + j as u64);
}
}
}
}
let mut out = self.predict(x, level)?;
for v in &mut out {
*v = -*v / std;
}
Ok(out)
}
fn reverse_sde(
&mut self,
sde: Sde,
known: Option<&[f64]>,
(resample, jump): (usize, usize),
) -> Result<Vec<f64>> {
let n_t = self.model.n_t.get();
let c = self.model.scales.len() as u64;
let level = |i: usize| n_t - 1 - i;
let step = |i: usize| level_time(n_t, level(i)) - level_time(n_t, level(i) - 1);
let mut x = self.prior()?;
let mut eval = 0u64;
let (mut i, mut passes) = (0usize, 0usize);
while i < n_t - 1 {
let t = level_time(n_t, level(i));
let h = step(i);
let score = self.score(sde, &mut x, known, level(i), eval)?;
let (drift, g2) = sde.drift_diffusion(t);
let noise_scale = (g2 * h).sqrt();
for (e, (v, &s)) in x.iter_mut().zip(&score).enumerate() {
let (row, j) = (e / c as usize, e as u64 % c);
let reverse = drift * *v - g2 * s;
*v = *v - reverse * h + noise_scale * self.noise(STEP_STREAM, row, eval * c + j);
}
check_finite(&x)?;
eval += 1;
if (i + 1).is_multiple_of(jump) && passes + 1 < resample && i + 1 >= jump {
let span: f64 = (i + 1 - jump..=i).map(step).sum();
let (drift, g2) = sde.drift_diffusion(t);
let scale = (g2 * span).sqrt();
for (e, v) in x.iter_mut().enumerate() {
let (row, j) = (e / c as usize, e as u64 % c);
*v += drift * *v * span + scale * self.noise(REPAINT_STREAM, row, eval * c + j);
}
passes += 1;
i = i + 1 - jump;
continue;
}
if (i + 1).is_multiple_of(jump) {
passes = 0;
}
i += 1;
}
let (_, std) = sde.marginal(EPS);
let score = self.score(sde, &mut x, known, 0, eval)?;
for (v, &s) in x.iter_mut().zip(&score) {
*v += std * std * s;
}
if let Some(known) = known {
for (v, &o) in x.iter_mut().zip(known) {
if !o.is_nan() {
*v = o;
}
}
}
check_finite(&x)?;
Ok(x)
}
}
fn positive_count(name: &'static str, v: usize) -> Result<()> {
ensure(name, v != 0, "must be at least 1")
}
fn check_finite(x: &[f64]) -> Result<()> {
if x.iter().all(|v| v.is_finite()) {
Ok(())
} else {
Err(diverged())
}
}
fn diverged() -> HessboostError {
HessboostError::invalid_param("n_t", "sampling diverged to non-finite values")
}
fn f32_exact(v: f64) -> bool {
let narrowed = v as f32;
narrowed.is_finite() && f64::from(narrowed) == v
}
#[cfg(test)]
mod tests {
use super::*;
use crate::training::train;
fn level_indexed_model(n_t: usize) -> ForestModel {
let data = DMatrix::from_dense(&[0.0, 1.0], 2, 1).unwrap();
let params = TrainingParams::default();
let models = (0..n_t)
.map(|level| {
let dtrain = data.clone().with_labels(&[level as f32; 2]).unwrap();
train(¶ms, &dtrain, 1).unwrap()
})
.collect();
ForestModel {
method: ForestMethod::forest_diffusion(),
n_t: NoiseLevels::new(n_t).unwrap(),
columns: vec![Column {
kind: ColumnKind::Continuous,
min: -1e30,
max: 1e30,
categories: Vec::new(),
}],
scales: vec![Scale {
min: -1.0,
range: 2.0,
}],
classes: Vec::new(),
class_probs: Vec::new(),
layout: OutputLayout::Joint,
models,
}
}
#[test]
fn reverse_steps_use_the_level_trained_at_their_time() {
let n_t = 600;
let model = level_indexed_model(n_t);
let sde = model.method.sde().unwrap();
let mut sampler = Sampler::new(&model, &[0, 0, 0], 7).unwrap();
let drawn = sampler.reverse_sde(sde, None, (1, n_t)).unwrap();
let mut x = sampler.prior().unwrap();
for (eval, level) in (1..n_t).rev().enumerate() {
let t = level_time(n_t, level);
let h = t - level_time(n_t, level - 1);
let (_, std) = sde.marginal(t);
let (drift, g2) = sde.drift_diffusion(t);
for (row, v) in x.iter_mut().enumerate() {
let score = -(level as f64) / std;
let noise = sampler.noise(STEP_STREAM, row, eval as u64);
*v = *v - (drift * *v - g2 * score) * h + (g2 * h).sqrt() * noise;
}
}
for (a, b) in drawn.iter().zip(&x) {
assert!((a - b).abs() <= 1e-9 * b.abs().max(1.0), "{a} vs {b}");
}
}
}