use super::eval::EvalSet;
use crate::config::{BoosterKind, ProcessType, TrainingParams};
use crate::data::DMatrix;
use crate::error::{HessboostError, Result};
use crate::model::BoostedModel;
use crate::objective::Loss;
use crate::training::multi_output;
use std::num::NonZeroUsize;
pub(super) fn validate_trained_model(model: &BoostedModel) -> Result<()> {
model.validate_structure().map_err(|e| {
let reason = match e {
HessboostError::ModelFormat(reason) => reason,
other => other.to_string(),
};
HessboostError::model_format(format!("training produced an invalid model: {reason}"))
})
}
pub(super) fn reject_feature_weights(dtrain: &DMatrix, reason: &'static str) -> Result<()> {
if dtrain.feature_weights().is_some() {
return Err(HessboostError::invalid_data("feature_weights", reason));
}
Ok(())
}
pub(super) fn validate_training_data(params: &TrainingParams, dtrain: &DMatrix) -> Result<()> {
let objective = params.loss(dtrain.n_targets())?;
validate_request(
&TrainRequest {
params,
dtrain,
evals: &[],
early_stopping_rounds: None,
},
objective.as_ref(),
)
}
pub(super) struct TrainRequest<'a> {
pub(super) params: &'a TrainingParams,
pub(super) dtrain: &'a DMatrix,
pub(super) evals: &'a [EvalSet<'a>],
pub(super) early_stopping_rounds: Option<NonZeroUsize>,
}
pub(super) fn validate_request(request: &TrainRequest, objective: &dyn Loss) -> Result<()> {
let &TrainRequest {
params,
dtrain,
evals,
..
} = request;
validate_setup(request, objective)?;
validate_query_bagging(params, dtrain)?;
if objective.requires_labels() && dtrain.labels().is_none() {
return Err(HessboostError::EmptyDataset("train: dtrain has no labels"));
}
validate_balanced_bagging(params, dtrain)?;
validate_datasets(objective, dtrain, evals)?;
validate_constraints(params, dtrain.n_cols())?;
validate_booster(request, objective)?;
validate_shrinkage_margins(request)?;
BoostedModel::check_iteration_size(objective.n_outputs(), params.num_parallel_tree)
}
fn validate_setup(request: &TrainRequest, objective: &dyn Loss) -> Result<()> {
let &TrainRequest {
params,
evals,
early_stopping_rounds,
..
} = request;
params.validate()?;
multi_output::validate(params, objective.n_outputs())?;
if early_stopping_rounds.is_some() && evals.is_empty() {
return Err(HessboostError::invalid_param(
"early_stopping_rounds",
"requires at least one evaluation dataset",
));
}
Ok(())
}
fn validate_query_bagging(params: &TrainingParams, dtrain: &DMatrix) -> Result<()> {
if params.bagging_by_query.is_none() {
return Ok(());
}
let Some(group) = dtrain.group() else {
return Err(HessboostError::invalid_data(
"group_sizes",
"`bagging_by_query` requires query group sizes on the training dataset",
));
};
if !group.partitions(dtrain.n_rows()) || group.iter_ranges().any(|(start, end)| start == end) {
return Err(HessboostError::invalid_data(
"group_sizes",
"`bagging_by_query` requires non-empty query groups covering all training rows",
));
}
Ok(())
}
fn validate_balanced_bagging(params: &TrainingParams, dtrain: &DMatrix) -> Result<()> {
if params.balanced_bagging.is_none() {
return Ok(());
}
if dtrain.n_targets() != 1 {
return Err(HessboostError::invalid_data(
"labels",
"balanced bagging requires exactly one label column",
));
}
let labels = dtrain.labels().ok_or(HessboostError::EmptyDataset(
"train: balanced bagging requires binary labels",
))?;
if labels.iter().any(|&label| label != 0.0 && label != 1.0) {
return Err(HessboostError::invalid_data(
"labels",
"balanced bagging requires labels exactly 0 or 1",
));
}
Ok(())
}
fn validate_constraints(params: &TrainingParams, n_features: usize) -> Result<()> {
if params.monotone_constraints.len() > n_features {
return Err(HessboostError::invalid_param(
"monotone_constraints",
"contains more entries than the training matrix has features",
));
}
for group in ¶ms.interaction_constraints {
if group.is_empty() {
return Err(HessboostError::invalid_param(
"interaction_constraints",
"constraint groups cannot be empty",
));
}
if let Some(&feature) = group.iter().find(|&&f| f as usize >= n_features) {
return Err(HessboostError::FeatureOutOfBounds {
index: feature as usize,
num_features: n_features,
});
}
}
Ok(())
}
fn validate_booster(request: &TrainRequest, objective: &dyn Loss) -> Result<()> {
let &TrainRequest { params, dtrain, .. } = request;
if params.booster == BoosterKind::GbLinear {
reject_feature_weights(dtrain, "gblinear does not sample columns")?;
}
if matches!(params.process_type, ProcessType::Update(_)) {
reject_feature_weights(dtrain, "`process_type=update` does not sample columns")?;
}
if matches!(params.booster, BoosterKind::Boulevard(_)) {
validate_boulevard_request(request, objective, "booster = boulevard")?;
}
if matches!(params.booster, BoosterKind::Ebm(_)) {
validate_ebm_request(request, objective)?;
}
Ok(())
}
fn validate_shrinkage_margins(request: &TrainRequest) -> Result<()> {
let &TrainRequest {
params,
dtrain,
evals,
..
} = request;
if params.model_shrinkage_on()
&& let Some(name) = std::iter::once(EvalSet {
data: dtrain,
name: "dtrain",
})
.chain(evals.iter().copied())
.find_map(|set| set.data.base_margin().map(|_| set.name))
{
return Err(HessboostError::invalid_data(
"base_margin",
"model shrinkage (`model_shrink_rate`) is not supported with a `base_margin`",
)
.in_dataset(name));
}
Ok(())
}
fn validate_boulevard_request(
request: &TrainRequest,
objective: &dyn Loss,
who: &str,
) -> Result<()> {
let refuse = |name: &'static str, reason: &str| {
Err(HessboostError::invalid_param(
name,
format!("`{who}`: {reason}"),
))
};
let refuse_data = |input: &'static str, reason: &str| {
HessboostError::invalid_data(input, format!("`{who}`: {reason}"))
};
if request.early_stopping_rounds.is_some() {
return refuse(
"early_stopping_rounds",
"the model averages every round, so it cannot stop at a best iteration",
);
}
if objective.name() != "reg:squarederror" {
return refuse(
"objective",
&format!("supports reg:squarederror only, got `{}`", objective.name()),
);
}
let dtrain = request.dtrain;
if dtrain.n_targets() != 1 || objective.n_outputs() != 1 {
return Err(refuse_data(
"labels",
&format!("needs one label column, got {}", dtrain.n_targets()),
));
}
if dtrain
.weights()
.is_some_and(|w| w.iter().any(|&v| v != 1.0))
{
return Err(refuse_data(
"weights",
"row weights other than 1 are not supported",
));
}
let dtrain_set = EvalSet {
data: dtrain,
name: "dtrain",
};
for set in std::iter::once(dtrain_set).chain(request.evals.iter().copied()) {
if set.data.base_margin().is_some() {
return Err(
refuse_data("base_margin", "base margins are not supported").in_dataset(set.name)
);
}
}
Ok(())
}
fn validate_ebm_request(request: &TrainRequest, objective: &dyn Loss) -> Result<()> {
let refuse = |name: &'static str, reason: &str| {
Err(HessboostError::invalid_param(
name,
format!("`booster = ebm`: {reason}"),
))
};
let refuse_data = |input: &'static str, reason: &str| {
Err(HessboostError::invalid_data(
input,
format!("`booster = ebm`: {reason}"),
))
};
if request.early_stopping_rounds.is_some() || !request.evals.is_empty() {
return refuse(
"early_stopping_rounds",
"eval sets and early stopping are not supported; evaluate the trained model",
);
}
if request.dtrain.n_targets() != 1 {
return refuse_data(
"labels",
&format!("needs one label column, got {}", request.dtrain.n_targets()),
);
}
if objective.n_outputs() != 1 {
return refuse(
"objective",
&format!(
"needs a single-output objective, got `{}`",
objective.name()
),
);
}
if request.dtrain.base_margin().is_some() {
return refuse_data(
"base_margin",
"base margins are not supported: the terms and their centering assume the \
intercept alone",
);
}
if request.params.ebm_settings().early_stopping().is_some() && request.dtrain.group().is_some()
{
return refuse(
"ebm_early_stopping_rounds",
"early stopping scores each bag's held-out rows, which split the query groups; not \
supported with query groups",
);
}
super::ebm::validate_data(request.dtrain)?;
if request.params.ebm_settings().boulevard() {
validate_boulevard_request(request, objective, "ebm_boulevard")?;
}
Ok(())
}
#[derive(Clone, Copy)]
struct DatasetContract {
targets: usize,
features: usize,
outputs: usize,
}
impl DatasetContract {
fn of(dtrain: &DMatrix, objective: &dyn Loss) -> Self {
DatasetContract {
targets: dtrain.n_targets(),
features: dtrain.n_cols(),
outputs: objective.n_outputs(),
}
}
}
pub(super) fn validate_datasets(
objective: &dyn Loss,
dtrain: &DMatrix,
evals: &[EvalSet],
) -> Result<()> {
let contract = DatasetContract::of(dtrain, objective);
validate_dataset(objective, dtrain, contract, "dtrain")?;
for set in evals {
validate_dataset(objective, set.data, contract, set.name)?;
}
Ok(())
}
fn validate_dataset(
objective: &dyn Loss,
data: &DMatrix,
contract: DatasetContract,
name: &str,
) -> Result<()> {
let DatasetContract {
targets: n_targets,
features: n_features,
outputs: n_out,
} = contract;
match data.labels() {
None if objective.requires_labels() => {
return Err(HessboostError::invalid_data(
"labels",
"missing; the objective needs labels",
)
.in_dataset(name));
}
None => {}
Some(labels) => {
let expected = data.n_rows().checked_mul(n_targets).ok_or_else(|| {
HessboostError::invalid_data("labels", "expected length overflows usize")
})?;
if labels.len() != expected {
return Err(HessboostError::dimension_mismatch(
"labels length (n_rows * training n_targets)",
expected,
labels.len(),
));
}
}
}
if data.n_cols() != n_features {
return Err(HessboostError::dimension_mismatch(
"dataset feature count",
n_features,
data.n_cols(),
));
}
if let Some(margin) = data.base_margin() {
let expected = data.n_rows().checked_mul(n_out).ok_or_else(|| {
HessboostError::invalid_data("base_margin", "expected length overflows usize")
})?;
if margin.len() != data.n_rows() && margin.len() != expected {
return Err(HessboostError::dimension_mismatch(
"base_margin length",
expected,
margin.len(),
));
}
}
objective
.validate_info(&data.info())
.map_err(|error| error.in_dataset(name))
}