use super::parse::{scalar_count, scalar_f64};
use crate::error::{HessboostError, Result};
use crate::model::objective::{ModelObjective, StoredObjectiveParams};
use crate::objective::{AftDistribution, Loss, Objective};
use serde::Deserialize;
use serde_json::{Map, Value, json};
use std::sync::Arc;
pub(super) fn build_objective(
objective: &ModelObjective,
n_targets: usize,
max_delta_step: f64,
) -> Result<Option<Arc<dyn Loss>>> {
match crate::model::rebuild_objective(objective, max_delta_step, n_targets) {
None => Ok(None),
Some(Ok(loss)) => Ok(Some(loss)),
Some(Err(error)) if n_targets > 1 => Err(HessboostError::model_format(format!(
"`num_target` {n_targets}: {error}"
))),
Some(Err(error)) => Err(HessboostError::model_format(format!(
"objective `{}`: {error}",
objective.name()
))),
}
}
pub(super) fn format_base_score(margins: &[f32], objective: &dyn Loss) -> String {
let mut stored = margins.to_vec();
if !matches!(objective.name(), "multi:softmax" | "multi:softprob") {
objective.margins_to_probs(&mut stored);
}
format_float_vector(stored)
}
pub(super) const MAX_BROADCAST_OUTPUTS: usize = 1 << 16;
pub(super) fn parse_base_score(
stored: &str,
objective: Option<&dyn Loss>,
n_outputs: usize,
) -> Result<Vec<f32>> {
let invalid = || HessboostError::model_format(format!("invalid `base_score` `{stored}`"));
let inner = stored
.trim()
.strip_prefix('[')
.and_then(|s| s.strip_suffix(']'))
.ok_or_else(invalid)?;
let values = inner
.split(',')
.map(|v| v.trim().parse::<f32>().ok())
.collect::<Option<Vec<f32>>>()
.ok_or_else(invalid)?;
let mut values = match values.len() {
1 if n_outputs <= MAX_BROADCAST_OUTPUTS => vec![values[0]; n_outputs],
len if len == n_outputs => values,
1 => {
return Err(HessboostError::model_format(format!(
"`base_score` has one entry for {n_outputs} outputs; above \
{MAX_BROADCAST_OUTPUTS} outputs it must list one entry per output"
)));
}
len => {
return Err(HessboostError::model_format(format!(
"`base_score` has {len} entries for {n_outputs} outputs"
)));
}
};
if let Some(obj) = objective {
obj.probs_to_margins(&mut values);
}
Ok(values)
}
pub(super) const SCALE_POS_WEIGHT: (&str, &str) = ("reg_loss_param", "scale_pos_weight");
pub(super) const MAX_DELTA_STEP: (&str, &str) = ("poisson_regression_param", "max_delta_step");
pub(super) const TWEEDIE_VARIANCE_POWER: (&str, &str) =
("tweedie_regression_param", "tweedie_variance_power");
pub(super) const HUBER_SLOPE: (&str, &str) = ("pseudo_huber_param", "huber_slope");
pub(super) const LAMBDARANK_NUM_PAIR: (&str, &str) =
("lambdarank_param", "lambdarank_num_pair_per_sample");
pub(super) const SOFTMAX_NUM_CLASS: (&str, &str) = ("softmax_multiclass_param", "num_class");
pub(super) const QUANTILE_ALPHA: (&str, &str) = ("quantile_loss_param", "quantile_alpha");
pub(super) const EXPECTILE_ALPHA: (&str, &str) = ("expectile_loss_param", "expectile_alpha");
pub(super) const AFT_LOSS_PARAM: &str = "aft_loss_param";
pub(super) fn format_float_vector(values: impl IntoIterator<Item = f32>) -> String {
use std::fmt::Write;
let mut out = String::from("[");
for (i, v) in values.into_iter().enumerate() {
if i > 0 {
out.push(',');
}
let _ = write!(out, "{v}");
}
out.push(']');
out
}
pub(super) fn parse_param_array(text: &str) -> Option<Vec<f64>> {
let text = text.trim();
let text = match text.strip_prefix('(').and_then(|t| t.strip_suffix(')')) {
Some(inner) => format!("[{inner}]"),
None => text.to_string(),
};
let as_f32 = |v: &Value| v.as_f64().map(|v| f64::from(v as f32));
match serde_json::from_str::<Value>(&text).ok()? {
Value::Array(entries) => entries.iter().map(as_f32).collect(),
number @ Value::Number(_) => as_f32(&number).map(|v| vec![v]),
_ => None,
}
}
pub(super) fn objective_to_json(objective: &Objective, max_delta_step: f64) -> Value {
let mut out = Map::with_capacity(2);
out.insert(
"name".to_string(),
Value::String(objective.name().to_string()),
);
let ((block, key), value) = match objective {
Objective::Aft(aft) => {
let fields = json!({
"aft_loss_distribution": aft.distribution(),
"aft_loss_distribution_scale": aft.scale().to_string(),
});
out.insert(AFT_LOSS_PARAM.to_string(), fields);
return Value::Object(out);
}
Objective::Cox
| Objective::SquaredLogError
| Objective::BinaryHinge
| Objective::AbsoluteError
| Objective::Dist(_)
| Objective::RankXendcg
| Objective::Custom(_) => {
return Value::Object(out);
}
Objective::Quantile(q) => (
QUANTILE_ALPHA,
format_float_vector(q.alpha().iter().map(|&v| v as f32)),
),
Objective::Expectile(e) => (
EXPECTILE_ALPHA,
format_float_vector(e.alpha().iter().map(|&v| v as f32)),
),
Objective::Softmax(c) | Objective::Softprob(c) => {
(SOFTMAX_NUM_CLASS, c.num_class().to_string())
}
Objective::Poisson => (MAX_DELTA_STEP, max_delta_step.to_string()),
Objective::Tweedie(t) => (TWEEDIE_VARIANCE_POWER, t.variance_power().to_string()),
Objective::PseudoHuber(h) => (HUBER_SLOPE, h.slope().to_string()),
Objective::RankPairwise(r) | Objective::RankNdcg(r) | Objective::RankMap(r) => {
(LAMBDARANK_NUM_PAIR, r.num_pair_per_sample().to_string())
}
Objective::SquaredError(r)
| Objective::RegLogistic(r)
| Objective::BinaryLogistic(r)
| Objective::BinaryLogitRaw(r)
| Objective::Gamma(r) => (SCALE_POS_WEIGHT, r.scale_pos_weight().to_string()),
};
let mut fields = Map::new();
if block == LAMBDARANK_NUM_PAIR.0 {
for (k, v) in [
("lambdarank_bias_norm", "1"),
("lambdarank_normalization", "1"),
("lambdarank_pair_method", "topk"),
("lambdarank_score_normalization", "1"),
("lambdarank_unbiased", "0"),
("ndcg_exp_gain", "1"),
] {
fields.insert(k.to_string(), Value::String(v.to_string()));
}
}
fields.insert(key.to_string(), Value::String(value));
out.insert(block.to_string(), Value::Object(fields));
Value::Object(out)
}
pub(super) fn objective_params_from_json(
objective: &str,
obj: Option<&Value>,
) -> Result<StoredObjectiveParams> {
let mut params = StoredObjectiveParams::defaults_for(objective);
let Some(obj) = obj else {
return Ok(params);
};
let invalid =
|key: &str, value: &Value| HessboostError::model_format(format!("invalid `{key}` {value}"));
for (param, value) in [
(SCALE_POS_WEIGHT, &mut params.scale_pos_weight),
(MAX_DELTA_STEP, &mut params.max_delta_step),
(TWEEDIE_VARIANCE_POWER, &mut params.tweedie_variance_power),
(HUBER_SLOPE, &mut params.huber_slope),
(
(AFT_LOSS_PARAM, "aft_loss_distribution_scale"),
&mut params.aft_loss_distribution_scale,
),
] {
if let Some(v) = objective_param(obj, param)? {
*value = scalar_f64(v).ok_or_else(|| invalid(param.1, v))?;
}
}
if let Some(v) = objective_param(obj, LAMBDARANK_NUM_PAIR)? {
let count = scalar_count(v).ok_or_else(|| invalid(LAMBDARANK_NUM_PAIR.1, v))?;
if count != u32::MAX as usize {
params.lambdarank_num_pair_per_sample = count;
}
}
for (param, alpha) in [
(QUANTILE_ALPHA, &mut params.quantile_alpha),
(EXPECTILE_ALPHA, &mut params.expectile_alpha),
] {
if let Some(v) = objective_param(obj, param)? {
*alpha = v
.as_str()
.and_then(parse_param_array)
.ok_or_else(|| invalid(param.1, v))?;
}
}
let distribution = (AFT_LOSS_PARAM, "aft_loss_distribution");
if let Some(v) = objective_param(obj, distribution)? {
params.aft_loss_distribution =
AftDistribution::deserialize(v).map_err(|_| invalid(distribution.1, v))?;
}
Ok(params)
}
pub(super) fn objective_param<'a>(
obj: &'a Value,
(block, key): (&str, &str),
) -> Result<Option<&'a Value>> {
match obj.get(block) {
None => Ok(None),
Some(Value::Object(fields)) => Ok(fields.get(key)),
Some(other) => Err(HessboostError::model_format(format!(
"objective parameter block `{block}` is not an object: {other}"
))),
}
}