mod fit;
pub mod forest;
mod format;
mod io;
mod process;
mod sample;
use std::num::NonZeroUsize;
use serde::{Deserialize, Serialize};
use crate::check::{ensure, positive};
use crate::config::{GrowPolicy, ProcessType, TrainingParams, TreeMethod};
use crate::data::DMatrix;
use crate::error::{HessboostError, Result};
use crate::model::BoostedModel;
use crate::objective::Objective;
pub use io::DiffusionFormat;
pub use sample::{Quantiles, SampleOptions, Samples, SamplesView};
#[derive(Debug, Clone, Copy, PartialEq, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
#[non_exhaustive]
pub enum Method {
Score(ScoreConfig),
FlowMatching(FlowMatchingConfig),
}
#[derive(Debug, Clone, Copy, PartialEq, Serialize, Deserialize)]
#[non_exhaustive]
pub struct ScoreConfig {
pub sde: Sde,
pub parameterization: Parameterization,
pub noise_level_feature: bool,
pub time_sampling: TimeSampling,
}
impl Default for ScoreConfig {
fn default() -> Self {
ScoreConfig {
sde: Sde::default(),
parameterization: Parameterization::Edm { sigma_data: 1.0 },
noise_level_feature: true,
time_sampling: TimeSampling::default(),
}
}
}
impl ScoreConfig {
pub fn treeffuser() -> Self {
ScoreConfig {
sde: Sde::default(),
parameterization: Parameterization::Noise,
noise_level_feature: false,
time_sampling: TimeSampling::Uniform,
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Serialize, Deserialize)]
#[non_exhaustive]
pub struct FlowMatchingConfig {
pub path: FlowPath,
pub time_sampling: TimeSampling,
pub solver: OdeSolver,
}
impl Default for FlowMatchingConfig {
fn default() -> Self {
FlowMatchingConfig {
path: FlowPath::VariancePreserving {
beta_min: 0.1,
beta_max: 20.0,
},
time_sampling: TimeSampling::default(),
solver: OdeSolver::Heun,
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
#[non_exhaustive]
pub enum Sde {
VarianceExploding {
sigma_min: f64,
sigma_max: f64,
},
VariancePreserving {
beta_min: f64,
beta_max: f64,
},
SubVariancePreserving {
beta_min: f64,
beta_max: f64,
},
}
impl Default for Sde {
fn default() -> Self {
Sde::VarianceExploding {
sigma_min: 0.01,
sigma_max: 20.0,
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
#[non_exhaustive]
pub enum Parameterization {
Noise,
Edm {
sigma_data: f64,
},
}
#[derive(Debug, Clone, Copy, PartialEq, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
#[non_exhaustive]
pub enum TimeSampling {
Uniform,
LogNoiseNormal {
mean: f64,
std: f64,
},
}
impl Default for TimeSampling {
fn default() -> Self {
TimeSampling::LogNoiseNormal {
mean: -1.2,
std: 1.2,
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
#[non_exhaustive]
pub enum FlowPath {
Linear,
Trigonometric,
VariancePreserving {
beta_min: f64,
beta_max: f64,
},
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
#[non_exhaustive]
pub enum OdeSolver {
Euler,
Heun,
}
#[derive(Debug, Clone, Copy, PartialEq)]
#[non_exhaustive]
pub struct EarlyStopping {
pub rounds: NonZeroUsize,
pub eval_fraction: f64,
}
impl Default for EarlyStopping {
fn default() -> Self {
EarlyStopping {
rounds: const { NonZeroUsize::new(50).unwrap() },
eval_fraction: 0.1,
}
}
}
#[derive(Debug, Clone)]
#[non_exhaustive]
pub struct Residualizer {
pub folds: usize,
pub training: TrainingParams,
pub num_boost_round: NonZeroUsize,
}
impl Default for Residualizer {
fn default() -> Self {
Residualizer {
folds: 5,
training: TrainingParams {
eta: 0.05,
max_depth: NonZeroUsize::new(6),
..lightgbm_like()
},
num_boost_round: const { NonZeroUsize::new(100).unwrap() },
}
}
}
#[derive(Debug, Clone)]
#[non_exhaustive]
pub struct DiffusionParams {
pub method: Method,
pub n_repeats: NonZeroUsize,
pub n_steps: NonZeroUsize,
pub training: TrainingParams,
pub num_boost_round: NonZeroUsize,
pub early_stopping: Option<EarlyStopping>,
pub residualizer: Option<Residualizer>,
pub seed: u64,
}
impl Default for DiffusionParams {
fn default() -> Self {
DiffusionParams {
method: Method::Score(ScoreConfig::default()),
n_repeats: const { NonZeroUsize::new(30).unwrap() },
n_steps: const { NonZeroUsize::new(50).unwrap() },
training: lightgbm_like(),
num_boost_round: const { NonZeroUsize::new(3000).unwrap() },
early_stopping: Some(EarlyStopping::default()),
residualizer: Some(Residualizer::default()),
seed: 0,
}
}
}
impl DiffusionParams {
pub fn treeffuser() -> Self {
DiffusionParams {
method: Method::Score(ScoreConfig::treeffuser()),
residualizer: None,
..DiffusionParams::default()
}
}
pub fn flow_matching() -> Self {
DiffusionParams {
method: Method::FlowMatching(FlowMatchingConfig::default()),
n_steps: const { NonZeroUsize::new(5).unwrap() },
..DiffusionParams::default()
}
}
pub fn validate(&self) -> Result<()> {
self.method.validate()?;
validate_regressor_params("training", &self.training)?;
if let Some(stop) = &self.early_stopping {
let f = stop.eval_fraction;
ensure(
"early_stopping.eval_fraction",
f.is_finite() && f > 0.0 && f < 1.0,
format!("must be in (0, 1), got {f}"),
)?;
}
if let Some(r) = &self.residualizer {
ensure(
"residualizer.folds",
r.folds >= 2,
format!("must be at least 2, got {}", r.folds),
)?;
validate_regressor_params("residualizer.training", &r.training)?;
}
Ok(())
}
}
fn lightgbm_like() -> TrainingParams {
TrainingParams {
tree_method: TreeMethod::Hist,
grow_policy: GrowPolicy::LossGuide,
max_depth: None,
max_leaves: NonZeroUsize::new(31),
eta: 0.1,
min_child_weight: 20.0,
lambda: 0.0,
max_bin: 255,
..TrainingParams::default()
}
}
fn validate_regressor_params(name: &'static str, params: &TrainingParams) -> Result<()> {
params.validate()?;
if !params.objective.is_unweighted_squared_error() {
return Err(HessboostError::invalid_param(
name,
format!(
"the diffusion targets are regressed with squared error: objective must be \
`reg:squarederror` with `scale_pos_weight = 1`, got `{}`{}",
params.objective.name(),
match ¶ms.objective {
Objective::SquaredError(r) => {
format!(" with `scale_pos_weight = {}`", r.scale_pos_weight())
}
_ => String::new(),
}
),
));
}
if matches!(params.process_type, ProcessType::Update(_)) {
return Err(HessboostError::invalid_param(
name,
"the diffusion GBDTs are trained from scratch: process_type must be `default`, \
not `update` (refresh)",
));
}
Ok(())
}
fn check_schedule(name: &'static str, lo: f64, hi: f64) -> Result<()> {
positive(name, lo)?;
positive(name, hi)?;
ensure(
name,
lo < hi,
format!("the minimum ({lo}) must be below the maximum ({hi})"),
)
}
impl Method {
fn validate(&self) -> Result<()> {
match self {
Method::Score(score) => {
match score.sde {
Sde::VarianceExploding {
sigma_min,
sigma_max,
} => check_schedule("sde", sigma_min, sigma_max)?,
Sde::VariancePreserving { beta_min, beta_max }
| Sde::SubVariancePreserving { beta_min, beta_max } => {
check_schedule("sde", beta_min, beta_max)?;
}
}
if let Parameterization::Edm { sigma_data } = score.parameterization {
positive("sigma_data", sigma_data)?;
positive("sigma_data", sigma_data * sigma_data)?;
}
let finite = [process::T_EPS, 1.0].iter().all(|&t| {
let (alpha, std) = score.sde.marginal(t);
let (c, g2) = score.sde.drift_diffusion(t);
[alpha, std, std.ln(), c, g2].iter().all(|v| v.is_finite())
}) && score.sde.prior_std().is_finite();
ensure(
"sde",
finite,
"the schedule's noise scale or drift is zero or overflows on [1e-5, 1]",
)?;
score.time_sampling.validate()
}
Method::FlowMatching(flow) => {
if let FlowPath::VariancePreserving { beta_min, beta_max } = flow.path {
check_schedule("path", beta_min, beta_max)?;
}
let finite = [process::T_EPS, 1.0].iter().all(|&t| {
let c = flow.path.coefficients(t);
[c.a, c.b, c.b.ln(), c.da, c.db]
.iter()
.all(|v| v.is_finite())
});
ensure(
"path",
finite,
"the path's noise scale or velocity is zero or overflows on [1e-5, 1]",
)?;
flow.time_sampling.validate()
}
}
}
fn time_columns(&self) -> usize {
match self {
Method::Score(score) => 1 + usize::from(score.noise_level_feature),
Method::FlowMatching(_) => 1,
}
}
}
impl TimeSampling {
fn validate(self) -> Result<()> {
if let TimeSampling::LogNoiseNormal { mean, std } = self {
ensure(
"time_sampling",
mean.is_finite(),
format!("mean must be finite, got {mean}"),
)?;
positive("time_sampling", std)?;
}
Ok(())
}
}
#[derive(Debug, Clone, Serialize)]
struct FittedResidualizer {
models: Vec<BoostedModel>,
center: Vec<f64>,
scale: Vec<f64>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(try_from = "format::UncheckedDiffusionModel")]
pub struct DiffusionModel {
method: Method,
n_steps: NonZeroUsize,
n_features: usize,
n_outputs: usize,
target_mean: Vec<f64>,
target_scale: Vec<f64>,
residualizer: Option<FittedResidualizer>,
regressor: BoostedModel,
}
impl DiffusionModel {
pub fn fit(params: &DiffusionParams, data: &DMatrix) -> Result<Self> {
fit::fit(params, data)
}
pub fn sample(
&self,
data: &DMatrix,
n_samples: usize,
options: &SampleOptions,
) -> Result<Samples> {
sample::sample(self, data, n_samples, options)
}
pub fn method(&self) -> &Method {
&self.method
}
pub fn n_steps(&self) -> NonZeroUsize {
self.n_steps
}
pub fn n_features(&self) -> usize {
self.n_features
}
pub fn n_outputs(&self) -> usize {
self.n_outputs
}
pub fn is_residualized(&self) -> bool {
self.residualizer.is_some()
}
pub fn regressor(&self) -> &BoostedModel {
&self.regressor
}
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)
}
fn validate(&self) -> Result<()> {
self.method
.validate()
.map_err(|e| HessboostError::model_format(e.to_string()))?;
let bad = |msg: String| Err(HessboostError::model_format(msg));
let d = self.n_outputs;
if d == 0 || self.n_features == 0 {
return bad("a diffusion model needs at least one feature and one output".into());
}
if self.target_mean.len() != d
|| self.target_scale.len() != d
|| !self.target_mean.iter().all(|v| v.is_finite())
|| !self.target_scale.iter().all(|v| v.is_finite() && *v > 0.0)
{
return bad(format!(
"the label standardization needs {d} finite means and positive scales"
));
}
let Some(regressor_features) = d
.checked_add(self.n_features)
.and_then(|n| n.checked_add(self.method.time_columns()))
else {
return bad("the feature count overflows usize".into());
};
check_regressor("regressor", &self.regressor, regressor_features, d)?;
if let Some(r) = &self.residualizer {
if r.models.len() < 2 {
return bad("the residualizer needs at least two fold models".into());
}
for model in &r.models {
check_regressor("residualizer model", model, self.n_features, d)?;
}
if r.center.len() != d
|| r.scale.len() != d
|| !r.center.iter().all(|v| v.is_finite())
|| !r.scale.iter().all(|v| v.is_finite() && *v > 0.0)
{
return bad(format!(
"the residualizer needs {d} finite centers and positive scales"
));
}
}
Ok(())
}
}
fn check_regressor(
what: &str,
model: &BoostedModel,
n_features: usize,
n_outputs: usize,
) -> Result<()> {
if !model
.objective()
.built_in()
.is_some_and(Objective::is_unweighted_squared_error)
|| model.n_features() != n_features
|| model.n_outputs() != n_outputs
{
return Err(HessboostError::model_format(format!(
"the {what} must be a reg:squarederror model (scale_pos_weight 1) with \
{n_features} features and {n_outputs} outputs, got `{}` with {} and {}",
model.objective().name(),
model.n_features(),
model.n_outputs()
)));
}
Ok(())
}