mod bitstream;
mod decode;
mod encode;
use super::native::{
OBJECTIVE_SECTIONS, SHRINKAGE_SECTIONS, read_objective_params, read_shrinkage,
write_objective_params, write_shrinkage,
};
use super::objective::{ModelObjective, StoredObjectiveParams};
use super::sections::{Sections, Writer};
use crate::data::DMatrix;
use crate::error::{HessboostError, Result};
use crate::model::ModelFormat;
use crate::model::{
BoostedModel, RowBlock, Shrinkage, initial_margins, shrink_margins, transform_model_margins,
validate_prediction_data,
};
use crate::tree::{RegTree, scalar_tree_output};
use bitstream::read_bits;
use encode::encode;
use rayon::prelude::*;
const MAGIC: &[u8; 4] = b"HBTD";
const VERSION: u8 = 1;
const PREFIX_BYTES: usize = MAGIC.len() + 1 + 4;
const MAX_HEAP_DEPTH: u32 = 24;
const STREAM_PAD: usize = 8;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
#[repr(u32)]
enum ThresholdKind {
Unsigned = 0,
Signed = 1,
Float = 2,
Categorical = 3,
}
impl TryFrom<u32> for ThresholdKind {
type Error = HessboostError;
fn try_from(raw: u32) -> Result<Self> {
Ok(match raw {
0 => ThresholdKind::Unsigned,
1 => ThresholdKind::Signed,
2 => ThresholdKind::Float,
3 => ThresholdKind::Categorical,
_ => return Err(format_error("invalid numeric type")),
})
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
#[repr(u32)]
enum DefaultDirection {
AllLeft = 0,
AllRight = 1,
PerNode = 2,
}
impl TryFrom<u32> for DefaultDirection {
type Error = HessboostError;
fn try_from(raw: u32) -> Result<Self> {
Ok(match raw {
0 => DefaultDirection::AllLeft,
1 => DefaultDirection::AllRight,
2 => DefaultDirection::PerNode,
_ => return Err(format_error("invalid default-direction mode")),
})
}
}
fn bits(n: u64) -> u32 {
u64::BITS - n.leading_zeros()
}
fn format_error(msg: impl Into<String>) -> HessboostError {
HessboostError::model_format(format!("compact model: {}", msg.into()))
}
#[derive(Debug, Clone)]
struct Meta {
objective: ModelObjective,
max_delta_step: f64,
num_class: usize,
n_targets: usize,
num_parallel_tree: usize,
shrinkage: Option<Shrinkage>,
}
const META_SECTIONS: &[&str] = &["objective", "num_class", "n_targets", "num_parallel_tree"];
impl Meta {
fn section_table(&self) -> Writer {
let mut w = Writer::default();
let name = self.objective.name();
w.str("objective", name);
w.u64("num_class", self.num_class as u64);
w.u64("n_targets", self.n_targets as u64);
w.u64("num_parallel_tree", self.num_parallel_tree as u64);
let params = StoredObjectiveParams::of(&self.objective, self.max_delta_step);
if params != StoredObjectiveParams::defaults_for(name) {
write_objective_params(&mut w, ¶ms);
}
if let Some(shrinkage) = &self.shrinkage {
write_shrinkage(&mut w, shrinkage);
}
w
}
fn decode(bytes: &[u8]) -> Result<Self> {
let (s, rest) = Sections::parse(bytes, |name| {
META_SECTIONS.contains(&name)
|| OBJECTIVE_SECTIONS.contains(&name)
|| SHRINKAGE_SECTIONS.contains(&name)
})
.map_err(|e| format_error(format!("metadata: {e}")))?;
if !rest.is_empty() {
return Err(format_error("metadata has trailing bytes"));
}
let name = s.str("objective")?;
let params = read_objective_params(&s, StoredObjectiveParams::defaults_for(name))?;
let num_class = s.usize("num_class")?;
Ok(Meta {
objective: ModelObjective::from_stored(name, ¶ms, num_class)?,
max_delta_step: params.max_delta_step,
num_class,
n_targets: s.usize("n_targets")?,
num_parallel_tree: s.usize("num_parallel_tree")?,
shrinkage: read_shrinkage(&s)?,
})
}
}
#[derive(Debug, Clone)]
enum Dictionary {
Numeric(Vec<f32>),
Categorical(Vec<Vec<u32>>),
}
impl Dictionary {
fn len(&self) -> usize {
match self {
Dictionary::Numeric(t) => t.len(),
Dictionary::Categorical(s) => s.len(),
}
}
}
#[derive(Debug, Clone)]
struct FeatureEntry {
input: usize,
dict: Dictionary,
}
#[derive(Debug, Clone, Copy)]
struct Widths {
feature_ref: u32,
threshold_ref: u32,
leaf_ref: u32,
default_bit: u32,
}
impl Widths {
fn new(
n_used: usize,
max_thresholds: usize,
n_leaves: usize,
default_mode: DefaultDirection,
) -> Self {
Widths {
feature_ref: bits(n_used.saturating_sub(1) as u64),
threshold_ref: bits(max_thresholds.saturating_sub(1) as u64),
leaf_ref: bits(n_leaves.saturating_sub(1) as u64),
default_bit: u32::from(default_mode == DefaultDirection::PerNode),
}
}
fn split(self) -> u32 {
self.feature_ref + self.threshold_ref + self.default_bit
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum Layout {
Heap { depth: u32, complete: bool },
Preorder { nodes: u32 },
}
impl Layout {
fn slot_width(self, w: Widths) -> u32 {
match self {
Layout::Heap { complete: true, .. } => w.split(),
Layout::Heap { .. } => 1 + w.split().max(w.leaf_ref),
Layout::Preorder { nodes } => {
1 + (w.split() + bits(u64::from(nodes) - 1)).max(w.leaf_ref)
}
}
}
fn total_bits(self, w: Widths) -> u128 {
let slot = u128::from(self.slot_width(w));
match self {
Layout::Heap { depth, .. } => {
let leaves = 1u128 << depth;
(leaves - 1) * slot + leaves * u128::from(w.leaf_ref)
}
Layout::Preorder { nodes } => u128::from(nodes) * slot,
}
}
}
#[derive(Debug, Clone, Copy)]
struct PackedTree {
layout: Layout,
offset: usize,
}
enum Slot {
Leaf(u32),
Split {
feature: u32,
threshold: u32,
default_left: bool,
right: u32,
},
}
#[derive(Debug, Clone)]
pub struct CompactModel {
bytes: Vec<u8>,
meta: Meta,
n_features: usize,
base_score: Vec<f32>,
tree_weights: Option<Vec<f32>>,
default_mode: DefaultDirection,
features: Vec<FeatureEntry>,
leaf_values: Vec<f32>,
widths: Widths,
trees: Vec<PackedTree>,
}
impl CompactModel {
fn serialized(&self) -> &[u8] {
&self.bytes[..self.bytes.len() - STREAM_PAD]
}
#[inline]
fn slot(&self, tree: PackedTree, i: u32, leaf_level: bool) -> Slot {
let w = self.widths;
let s = &self.bytes;
let (mut pos, flagged) = match tree.layout {
Layout::Heap { depth, complete } => {
let internal = (1usize << depth) - 1;
let slot = tree.layout.slot_width(w) as usize;
if leaf_level {
let j = i as usize - internal;
let pos = tree.offset + internal * slot + j * w.leaf_ref as usize;
return Slot::Leaf(read_bits(s, pos, w.leaf_ref));
}
(tree.offset + i as usize * slot, !complete)
}
Layout::Preorder { .. } => {
let slot = tree.layout.slot_width(w) as usize;
(tree.offset + i as usize * slot, true)
}
};
if flagged {
let is_leaf = read_bits(s, pos, 1) == 1;
pos += 1;
if is_leaf {
return Slot::Leaf(read_bits(s, pos, w.leaf_ref));
}
}
let feature = read_bits(s, pos, w.feature_ref);
pos += w.feature_ref as usize;
let threshold = read_bits(s, pos, w.threshold_ref);
pos += w.threshold_ref as usize;
let default_left = match self.default_mode {
DefaultDirection::AllLeft => true,
DefaultDirection::AllRight => false,
DefaultDirection::PerNode => {
let bit = read_bits(s, pos, 1) == 1;
pos += 1;
bit
}
};
let right = match tree.layout {
Layout::Preorder { nodes } => read_bits(s, pos, bits(u64::from(nodes) - 1)),
Layout::Heap { .. } => 0,
};
Slot::Split {
feature,
threshold,
default_left,
right,
}
}
#[inline]
fn goes_left(&self, feature: u32, threshold: u32, default_left: bool, row: &[f32]) -> bool {
let entry = &self.features[feature as usize];
let v = row[entry.input];
if v.is_nan() {
return default_left;
}
match &entry.dict {
Dictionary::Numeric(t) => v < t[threshold as usize],
Dictionary::Categorical(sets) => {
sets[threshold as usize].binary_search(&(v as u32)).is_ok()
}
}
}
fn tree_leaf(&self, t: usize, row: &[f32]) -> f32 {
let tree = self.trees[t];
let first_leaf = match tree.layout {
Layout::Heap { depth, .. } => (1u32 << depth) - 1,
Layout::Preorder { .. } => u32::MAX,
};
let mut i = 0u32;
let leaf = loop {
match self.slot(tree, i, i >= first_leaf) {
Slot::Leaf(leaf) => break leaf,
Slot::Split {
feature,
threshold,
default_left,
right,
} => {
let left = self.goes_left(feature, threshold, default_left, row);
i = match tree.layout {
Layout::Heap { .. } => 2 * i + if left { 1 } else { 2 },
Layout::Preorder { .. } => i + if left { 1 } else { right },
};
}
}
};
self.leaf_values[leaf as usize]
}
pub fn predict_margin(&self, data: &DMatrix) -> Result<super::Predictions> {
let k = self.n_outputs();
validate_prediction_data(self.n_features, k, data)?;
let shrinkage = self.meta.shrinkage.as_ref();
let (mut out, weights) = match shrinkage {
Some(shrinkage) => (shrinkage.start_margins(data), None),
None => (
initial_margins(&self.base_score, data),
self.tree_weights.as_deref(),
),
};
let weight = |t: usize| weights.map_or(1.0, |w| w[t]);
let parallel = self.meta.num_parallel_tree;
let per = parallel * k;
out.par_chunks_mut(k)
.enumerate()
.with_min_len(256)
.for_each_init(
|| RowBlock::single_rows(data),
|block, (r, margins)| {
block.load(r, 1);
let row = block.row(0).expect("single-row blocks are dense");
for t in 0..self.trees.len() {
if let Some(shrinkage) = shrinkage
&& t % per == 0
{
shrink_margins(margins, shrinkage.factors()[t / per]);
}
margins[scalar_tree_output(t, parallel, k)] +=
weight(t) * self.tree_leaf(t, row);
}
},
);
if let Some(shrinkage) = shrinkage {
shrinkage.finish_margins(data, &mut out);
}
Ok(super::Predictions::new(out, data.n_rows(), k))
}
pub fn predict(&self, data: &DMatrix) -> Result<super::Predictions> {
let margin = self.predict_margin(data)?;
Ok(transform_model_margins(
&self.meta.objective,
self.meta.max_delta_step,
self.meta.n_targets,
margin,
))
}
pub fn encode(&self) -> Vec<u8> {
self.serialized().to_vec()
}
pub fn size_bytes(&self) -> usize {
self.serialized().len()
}
pub fn save(&self, path: impl AsRef<std::path::Path>) -> Result<()> {
std::fs::write(path, self.serialized())?;
Ok(())
}
pub fn load(path: impl AsRef<std::path::Path>) -> Result<Self> {
Self::decode(std::fs::read(path)?)
}
pub fn objective(&self) -> &ModelObjective {
&self.meta.objective
}
pub fn num_trees(&self) -> usize {
self.trees.len()
}
pub fn n_features(&self) -> usize {
self.n_features
}
pub fn n_outputs(&self) -> usize {
self.base_score.len()
}
pub fn used_features(&self) -> Vec<usize> {
self.features.iter().map(|e| e.input).collect()
}
pub fn num_thresholds(&self) -> usize {
self.features.iter().map(|e| e.dict.len()).sum()
}
pub fn num_leaf_values(&self) -> usize {
self.leaf_values.len()
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
#[non_exhaustive]
pub struct ModelSizeReport {
pub native_bytes: usize,
pub compact_bytes: usize,
pub trees: usize,
pub splits: usize,
pub leaves: usize,
pub used_features: usize,
pub thresholds: usize,
pub leaf_values: usize,
}
impl ModelSizeReport {
pub fn compression_ratio(&self) -> f64 {
self.native_bytes as f64 / self.compact_bytes as f64
}
pub fn reuse_factor(&self) -> f64 {
(self.splits + self.leaves) as f64 / (self.thresholds + self.leaf_values).max(1) as f64
}
}
impl BoostedModel {
pub fn to_compact(&self) -> Result<CompactModel> {
CompactModel::decode(encode(self)?)
}
pub fn to_compact_bytes(&self) -> Result<Vec<u8>> {
encode(self)
}
pub fn size_report(&self) -> Result<ModelSizeReport> {
let compact = self.to_compact()?;
let trees = &self.trees()[..compact.num_trees()];
let splits: usize = trees
.iter()
.map(|t| t.nodes().iter().filter(|n| !n.is_leaf()).count())
.sum();
let nodes: usize = trees.iter().map(RegTree::num_nodes).sum();
Ok(ModelSizeReport {
native_bytes: self.encode(ModelFormat::Binary)?.len(),
compact_bytes: compact.size_bytes(),
trees: trees.len(),
splits,
leaves: nodes - splits,
used_features: compact.features.len(),
thresholds: compact.num_thresholds(),
leaf_values: compact.num_leaf_values(),
})
}
}
#[cfg(test)]
mod tests;