use crate::config::TrainingParams;
use crate::data::DMatrix;
use crate::error::{HessboostError, Result};
use crate::learner::LinearModel;
use crate::learner::model::for_each_present_value;
use crate::objective::{GradPair, Objective};
struct Column {
rows: Vec<u32>,
vals: Vec<f32>,
}
fn coordinate_delta(sum_grad: f64, sum_hess: f64, w: f64, alpha: f64, lambda: f64) -> f64 {
if sum_hess < 1e-5 {
return 0.0;
}
let sum_grad_l2 = sum_grad + lambda * w;
let sum_hess_l2 = sum_hess + lambda;
let tmp = w - sum_grad_l2 / sum_hess_l2;
if tmp >= 0.0 {
(-(sum_grad_l2 + alpha) / sum_hess_l2).max(-w)
} else {
(-(sum_grad_l2 - alpha) / sum_hess_l2).min(-w)
}
}
pub(crate) fn train_gblinear(
params: &TrainingParams,
dtrain: &DMatrix,
num_round: usize,
initial_margin: &[f32],
n_out: usize,
objective: &dyn Objective,
) -> Result<LinearModel> {
let n = dtrain.n_rows();
let n_features = dtrain.n_cols();
let labels = dtrain.labels().ok_or(HessboostError::EmptyDataset(
"gblinear: dtrain has no labels",
))?;
let weights = dtrain.weights();
let group = dtrain.group();
let eta = params.eta;
let lambda = params.lambda;
let alpha = params.alpha;
let mut cols: Vec<Column> = (0..n_features)
.map(|_| Column {
rows: Vec::new(),
vals: Vec::new(),
})
.collect();
for row in 0..n {
for_each_present_value(dtrain, row, |f, x| {
if x != 0.0 {
let col = &mut cols[f];
col.rows.push(row as u32);
col.vals.push(x);
}
});
}
let mut lin_weights = vec![0.0f32; n_features * n_out];
let mut bias = vec![0.0f32; n_out];
let mut margin = initial_margin.to_vec();
let mut gpair = vec![GradPair::default(); n * n_out];
for _round in 0..num_round {
objective.gradient_grouped(&margin, labels, weights, group, &mut gpair);
for k in 0..n_out {
let mut g = 0.0f64;
let mut h = 0.0f64;
for i in 0..n {
let gp = gpair[i * n_out + k];
g += f64::from(gp.grad);
h += f64::from(gp.hess);
}
let db = eta * coordinate_delta(g, h, 0.0, 0.0, 0.0);
if db != 0.0 {
let db32 = db as f32;
bias[k] += db32;
for i in 0..n {
let gp = &mut gpair[i * n_out + k];
gp.grad += gp.hess * db32;
margin[i * n_out + k] += db32;
}
}
for f in 0..n_features {
let col = &cols[f];
let mut g = 0.0f64;
let mut h = 0.0f64;
for (idx, &row) in col.rows.iter().enumerate() {
let x = f64::from(col.vals[idx]);
let gp = gpair[row as usize * n_out + k];
g += f64::from(gp.grad) * x;
h += f64::from(gp.hess) * x * x;
}
let w = f64::from(lin_weights[f * n_out + k]);
let dw = eta * coordinate_delta(g, h, w, alpha, lambda);
if dw == 0.0 {
continue;
}
let dw32 = dw as f32;
lin_weights[f * n_out + k] += dw32;
for (idx, &row) in col.rows.iter().enumerate() {
let x = col.vals[idx];
let gp = &mut gpair[row as usize * n_out + k];
gp.grad += gp.hess * x * dw32;
margin[row as usize * n_out + k] += x * dw32;
}
}
}
}
Ok(LinearModel::new(lin_weights, bias))
}