use std::path::Path;
use std::str::FromStr;
use crate::{Error, Model, Result};
#[derive(Debug, Clone, Copy, PartialEq)]
enum Objective {
Identity,
Sqrt,
Sigmoid(f64),
Log1pExp,
Exp,
}
impl Objective {
fn parse(value: &str, lineno: usize) -> Result<Self> {
let mut tokens = value.split_whitespace();
let name = tokens.next().unwrap_or("");
match name {
"binary" => {
let mut k: f64 = 1.0;
for tok in tokens {
match tok.strip_prefix("sigmoid:") {
Some(v) => {
k = parse_scalar(
v,
"objective sigmoid",
lineno,
)?;
}
None => return Err(Self::unknown_token(name, tok)),
}
}
if !(k.is_finite() && k > 0.0) {
return Err(Error::Parse {
line: lineno,
message: format!(
"sigmoid parameter must be positive and \
finite, got {k}"
),
});
}
Ok(Objective::Sigmoid(k))
}
"regression" | "regression_l1" | "fair" | "quantile"
| "mape" => {
let mut sqrt = false;
for tok in tokens {
match tok {
"sqrt" => sqrt = true,
_ => return Err(Self::unknown_token(name, tok)),
}
}
Ok(if sqrt {
Objective::Sqrt
} else {
Objective::Identity
})
}
"cross_entropy" | "cross_entropy_lambda" | "poisson"
| "gamma" | "tweedie" | "huber" | "lambdarank"
| "rank_xendcg" | "custom" | "none" => {
if let Some(tok) = tokens.next() {
return Err(Self::unknown_token(name, tok));
}
Ok(match name {
"cross_entropy" => Objective::Sigmoid(1.0),
"cross_entropy_lambda" => Objective::Log1pExp,
"poisson" | "gamma" | "tweedie" => Objective::Exp,
_ => Objective::Identity,
})
}
other => Err(Error::Unsupported {
message: format!(
"objective '{other}' is not recognised; loading it \
would produce wrong predictions"
),
}),
}
}
fn unknown_token(name: &str, tok: &str) -> Error {
Error::Unsupported {
message: format!(
"objective '{name}' carries unrecognised token {tok:?}, \
which may change the output transform; loading it would \
risk wrong predictions. The accepted tokens cover every \
LightGBM release through 4.6 — if a newer LightGBM wrote \
this file, please report this token"
),
}
}
fn transform(self, raw: f64) -> f64 {
match self {
Objective::Identity => raw,
Objective::Sqrt => raw * raw.abs(),
Objective::Sigmoid(k) => 1.0 / (1.0 + (-k * raw).exp()),
Objective::Log1pExp => raw.exp().ln_1p(),
Objective::Exp => raw.exp(),
}
}
}
#[derive(Debug)]
struct Tree {
num_leaves: usize,
split_feature: Vec<usize>,
threshold: Vec<f64>,
left_child: Vec<i32>,
right_child: Vec<i32>,
leaf_value: Vec<f64>,
decision_type: Vec<u8>,
cat_boundaries: Vec<usize>,
cat_threshold: Vec<u32>,
}
impl Tree {
fn predict(&self, features: &[f64]) -> f64 {
if self.split_feature.is_empty() {
return self.leaf_value[0];
}
let mut node: i32 = 0;
loop {
if node < 0 {
return self.leaf_value[(!node) as usize];
}
let idx = node as usize;
let feat_idx = self.split_feature[idx];
let val = if feat_idx < features.len() {
features[feat_idx]
} else {
f64::NAN
};
node = if self.decision_type[idx] & 1 != 0 {
self.decide_categorical(idx, val)
} else {
self.decide_numerical(idx, val)
};
}
}
fn decide_numerical(&self, idx: usize, mut val: f64) -> i32 {
let dt = self.decision_type[idx];
let default_left = (dt & 2) != 0;
let missing_type = (dt >> 2) & 3;
if val.is_nan() && missing_type != 2 {
val = 0.0;
}
const K_ZERO_THRESHOLD: f64 = 1e-35_f32 as f64;
let is_missing = (missing_type == 1
&& (-K_ZERO_THRESHOLD..=K_ZERO_THRESHOLD).contains(&val))
|| (missing_type == 2 && val.is_nan());
if is_missing {
if default_left {
self.left_child[idx]
} else {
self.right_child[idx]
}
} else if val <= self.threshold[idx] {
self.left_child[idx]
} else {
self.right_child[idx]
}
}
fn decide_categorical(&self, idx: usize, val: f64) -> i32 {
if val.is_nan() {
return self.right_child[idx];
}
let cat = val as i64; if cat < 0 {
return self.right_child[idx];
}
let cat_idx = self.threshold[idx] as usize;
let bits = &self.cat_threshold[self.cat_boundaries[cat_idx]
..self.cat_boundaries[cat_idx + 1]];
let word = (cat as u64 >> 5) as usize;
let found = word < bits.len()
&& (bits[word] >> (cat as u64 & 31)) & 1 == 1;
if found {
self.left_child[idx]
} else {
self.right_child[idx]
}
}
fn validate(
&self,
header_line: usize,
max_feature_idx: Option<usize>,
) -> Result<()> {
let err = |message: String| Error::Parse {
line: header_line,
message,
};
if self.num_leaves == 0 {
return Err(err("num_leaves must be >= 1".into()));
}
if self.leaf_value.len() != self.num_leaves {
return Err(err(format!(
"num_leaves={} but {} leaf_value entries",
self.num_leaves,
self.leaf_value.len()
)));
}
let internal = self.num_leaves - 1;
for (name, len) in [
("split_feature", self.split_feature.len()),
("threshold", self.threshold.len()),
("decision_type", self.decision_type.len()),
("left_child", self.left_child.len()),
("right_child", self.right_child.len()),
] {
if len != internal {
return Err(err(format!(
"{name} has {len} entries, expected num_leaves-1 = {internal}"
)));
}
}
if self.cat_boundaries.windows(2).any(|w| w[0] > w[1])
|| self.cat_boundaries.last().copied().unwrap_or(0)
> self.cat_threshold.len()
{
return Err(err(
"cat_boundaries must be non-decreasing and within cat_threshold"
.into(),
));
}
let max_idx =
max_feature_idx.unwrap_or(i32::MAX as usize - 1);
for (k, &feat) in self.split_feature.iter().enumerate() {
if feat > max_idx {
return Err(err(format!(
"split_feature[{k}] = {feat} exceeds max_feature_idx = {max_idx}"
)));
}
}
for k in 0..internal {
if self.decision_type[k] & 1 != 0 {
let cat_idx = self.threshold[k];
let max = self.cat_boundaries.len().saturating_sub(1);
if cat_idx.fract() != 0.0
|| cat_idx < 0.0
|| cat_idx as usize >= max
{
return Err(err(format!(
"categorical node {k} references invalid bitset index {cat_idx}"
)));
}
}
for (side, child) in [
("left_child", self.left_child[k]),
("right_child", self.right_child[k]),
] {
let ok = if child >= 0 {
(child as usize) < internal
} else {
((!child) as usize) < self.num_leaves
};
if !ok {
return Err(err(format!(
"{side}[{k}] = {child} is out of range"
)));
}
}
}
if internal > 0 {
let mut node_seen = vec![false; internal];
let mut leaf_seen = vec![false; self.num_leaves];
let mut stack: Vec<i32> = vec![0];
while let Some(child) = stack.pop() {
let seen = if child >= 0 {
&mut node_seen[child as usize]
} else {
&mut leaf_seen[(!child) as usize]
};
if *seen {
return Err(err(format!(
"node {child} is reachable more than once — child \
pointers do not form a tree"
)));
}
*seen = true;
if child >= 0 {
stack.push(self.left_child[child as usize]);
stack.push(self.right_child[child as usize]);
}
}
if node_seen.iter().any(|&v| !v)
|| leaf_seen.iter().any(|&v| !v)
{
return Err(err(
"tree has internal nodes or leaves unreachable from the \
root"
.into(),
));
}
}
Ok(())
}
}
#[derive(Debug)]
pub struct LgbModel {
trees: Vec<Tree>,
objective: Objective,
average_output: bool,
num_features: usize,
}
fn parse_list<T>(s: &str, field: &str, line: usize) -> Result<Vec<T>>
where
T: FromStr,
T::Err: std::fmt::Display,
{
s.split_whitespace()
.map(|tok| {
tok.parse::<T>().map_err(|e| Error::Parse {
line,
message: format!(
"{field}: invalid value {tok:?}: {e}"
),
})
})
.collect()
}
fn parse_scalar<T>(s: &str, field: &str, line: usize) -> Result<T>
where
T: FromStr,
T::Err: std::fmt::Display,
{
s.trim().parse::<T>().map_err(|e| Error::Parse {
line,
message: format!("{field}: {e}"),
})
}
impl LgbModel {
pub fn load(path: &Path) -> Result<Self> {
let content =
std::fs::read_to_string(path).map_err(|source| {
Error::Io {
path: path.to_path_buf(),
source,
}
})?;
Self::parse(&content)
}
pub fn parse(content: &str) -> Result<Self> {
let lines: Vec<&str> = content.lines().collect();
let mut objective = Objective::Identity;
let mut average_output = false;
let mut max_feature_idx: Option<usize> = None;
for (i, line) in lines.iter().enumerate() {
if line.starts_with("Tree=") {
break;
}
if let Some(v) = line.strip_prefix("objective=") {
objective = Objective::parse(v, i + 1)?;
} else if let Some(v) =
line.strip_prefix("max_feature_idx=")
{
let m: usize =
parse_scalar(v, "max_feature_idx", i + 1)?;
if m >= i32::MAX as usize {
return Err(Error::Parse {
line: i + 1,
message: format!(
"max_feature_idx = {m} is out of range"
),
});
}
max_feature_idx = Some(m);
} else if let Some(v) = line.strip_prefix("num_class=") {
let n: usize = parse_scalar(v, "num_class", i + 1)?;
if n > 1 {
return Err(Error::Unsupported {
message: format!(
"multiclass model (num_class={n}); only \
single-output models are supported"
),
});
}
} else if *line == "average_output" {
average_output = true;
}
}
let mut trees = Vec::new();
let mut i = 0;
while i < lines.len() {
if !lines[i].starts_with("Tree=") {
i += 1;
continue;
}
let header_line = i + 1;
i += 1;
let mut num_leaves: Option<usize> = None;
let mut split_feature = Vec::new();
let mut threshold = Vec::new();
let mut left_child = Vec::new();
let mut right_child = Vec::new();
let mut leaf_value = Vec::new();
let mut decision_type = Vec::new();
let mut cat_boundaries = Vec::new();
let mut cat_threshold = Vec::new();
while i < lines.len() && !lines[i].starts_with("Tree=") {
let line = lines[i];
let lineno = i + 1;
if let Some(v) = line.strip_prefix("num_leaves=") {
num_leaves =
Some(parse_scalar(v, "num_leaves", lineno)?);
} else if let Some(v) =
line.strip_prefix("split_feature=")
{
split_feature =
parse_list(v, "split_feature", lineno)?;
} else if let Some(v) =
line.strip_prefix("threshold=")
{
threshold = parse_list(v, "threshold", lineno)?;
} else if let Some(v) =
line.strip_prefix("decision_type=")
{
decision_type =
parse_list(v, "decision_type", lineno)?;
} else if let Some(v) =
line.strip_prefix("left_child=")
{
left_child = parse_list(v, "left_child", lineno)?;
} else if let Some(v) =
line.strip_prefix("right_child=")
{
right_child =
parse_list(v, "right_child", lineno)?;
} else if let Some(v) =
line.strip_prefix("leaf_value=")
{
leaf_value = parse_list(v, "leaf_value", lineno)?;
} else if let Some(v) =
line.strip_prefix("cat_boundaries=")
{
cat_boundaries =
parse_list(v, "cat_boundaries", lineno)?;
} else if let Some(v) =
line.strip_prefix("cat_threshold=")
{
cat_threshold =
parse_list(v, "cat_threshold", lineno)?;
} else if let Some(v) =
line.strip_prefix("is_linear=")
{
let is_linear: u8 =
parse_scalar(v, "is_linear", lineno)?;
if is_linear != 0 {
return Err(Error::Unsupported {
message:
"linear trees (is_linear=1) are \
not supported"
.into(),
});
}
}
i += 1;
}
let Some(num_leaves) = num_leaves else {
continue;
};
let tree = Tree {
num_leaves,
split_feature,
threshold,
left_child,
right_child,
leaf_value,
decision_type,
cat_boundaries,
cat_threshold,
};
tree.validate(header_line, max_feature_idx)?;
trees.push(tree);
}
if trees.is_empty() {
return Err(Error::EmptyModel);
}
let max_split = trees
.iter()
.flat_map(|t| &t.split_feature)
.max()
.copied();
let num_features = match (max_feature_idx, max_split) {
(Some(m), _) => m + 1,
(None, Some(s)) => s + 1,
(None, None) => 0,
};
Ok(Self {
trees,
objective,
average_output,
num_features,
})
}
#[must_use]
pub fn num_features(&self) -> usize {
self.num_features
}
#[must_use]
pub fn num_trees(&self) -> usize {
self.trees.len()
}
#[must_use]
pub fn predict_unchecked(&self, features: &[f64]) -> f64 {
let mut raw: f64 =
self.trees.iter().map(|t| t.predict(features)).sum();
if self.average_output {
raw /= self.trees.len() as f64;
}
self.objective.transform(raw)
}
}
impl Model for LgbModel {
fn predict(&self, features: &[f64]) -> Result<f64> {
if features.len() != self.num_features {
return Err(Error::FeatureCount {
expected: self.num_features,
got: features.len(),
});
}
Ok(self.predict_unchecked(features))
}
}
#[cfg(test)]
mod tests {
use super::*;
fn stump(header: &str, leaf: f64) -> String {
format!("{header}\nTree=0\nnum_leaves=1\nleaf_value={leaf}\n")
}
#[test]
fn test_objective_transforms() {
let cases = [
("objective=regression", 0.5f64),
("objective=regression sqrt", 0.25f64),
(
"objective=binary sigmoid:1",
1.0 / (1.0 + (-0.5f64).exp()),
),
(
"objective=binary sigmoid:2",
1.0 / (1.0 + (-1.0f64).exp()),
),
(
"objective=cross_entropy",
1.0 / (1.0 + (-0.5f64).exp()),
),
("objective=cross_entropy_lambda", 0.5f64.exp().ln_1p()),
("objective=poisson", 0.5f64.exp()),
("objective=tweedie", 0.5f64.exp()),
];
for (header, expected) in cases {
let model = LgbModel::parse(&stump(header, 0.5)).unwrap();
let got = model.predict(&[]).unwrap();
assert!(
(got - expected).abs() < 1e-15,
"{header}: got {got}, expected {expected}"
);
}
let model = LgbModel::parse(&stump(
"objective=regression sqrt",
-0.5,
))
.unwrap();
assert_eq!(model.predict(&[]).unwrap(), -0.25);
}
#[test]
fn test_sqrt_objective_selection() {
for name in [
"regression", "regression_l1", "fair", "quantile", "mape",
] {
let sqrt = LgbModel::parse(&stump(
&format!("objective={name} sqrt"),
0.5,
))
.unwrap();
assert_eq!(
sqrt.objective,
Objective::Sqrt,
"{name} sqrt"
);
let plain = LgbModel::parse(&stump(
&format!("objective={name}"),
0.5,
))
.unwrap();
assert_eq!(
plain.objective,
Objective::Identity,
"{name}"
);
}
}
#[test]
fn test_average_output() {
let text = "objective=regression\naverage_output\n\
Tree=0\nnum_leaves=1\nleaf_value=0.4\n\
Tree=1\nnum_leaves=1\nleaf_value=0.8\n";
let model = LgbModel::parse(text).unwrap();
assert!((model.predict(&[]).unwrap() - 0.6).abs() < 1e-15);
}
#[test]
fn test_unsupported_models_are_rejected() {
for (name, text) in [
("multiclass", stump("objective=multiclass num_class:3\nnum_class=3", 0.5)),
("unknown objective", stump("objective=who_knows", 0.5)),
("unknown binary token", stump("objective=binary who_knows", 0.5)),
("unknown regression token", stump("objective=regression who_knows", 0.5)),
("sqrt on non-sqrt objective", stump("objective=huber sqrt", 0.5)),
(
"linear tree",
"objective=regression\nTree=0\nnum_leaves=1\nleaf_value=0.5\nis_linear=1\n"
.to_string(),
),
] {
let err = LgbModel::parse(&text).unwrap_err();
assert!(
matches!(err, Error::Unsupported { .. }),
"{name}: expected Unsupported, got {err:?}"
);
}
}
#[test]
fn test_invalid_sigmoid_parameter_is_rejected() {
for bad in ["nan", "inf", "0", "-1"] {
let text = stump(
&format!("objective=binary sigmoid:{bad}"),
0.5,
);
let err = LgbModel::parse(&text).unwrap_err();
assert!(
matches!(err, Error::Parse { .. }),
"sigmoid:{bad}: got {err:?}"
);
}
}
#[test]
fn test_lightgbm_parity() {
let model_path = Path::new("tests/fixtures/tiny_binary.lgb");
let model = LgbModel::load(model_path)
.expect("Failed to load fixture model");
assert_eq!(
model.num_trees(),
8,
"fixture should have 8 trees"
);
assert_eq!(
model.num_features(),
5,
"fixture header declares max_feature_idx=4"
);
assert_eq!(
model.objective,
Objective::Sigmoid(1.0),
"Expected sigmoid (binary) objective"
);
let features = [0.5, -0.2, 0.7, 1.1, -0.9];
let pred = model.predict(&features).unwrap();
let expected = 0.879687246542221;
let diff = (pred - expected).abs();
assert!(
diff < 1e-9,
"Prediction mismatch vs LightGBM: rust={pred}, python={expected}, diff={diff:.2e}"
);
}
#[test]
fn test_lightgbm_nan_parity() {
let model_path = Path::new("tests/fixtures/tiny_nan.lgb");
let model = LgbModel::load(model_path)
.expect("Failed to load NaN fixture model");
let features = [f64::NAN, 0.3, f64::NAN, -0.5, 1.2];
let pred = model.predict(&features).unwrap();
let expected = 0.18514897124790036;
let diff = (pred - expected).abs();
assert!(
diff < 1e-9,
"NaN prediction mismatch vs LightGBM: rust={pred}, python={expected}, diff={diff:.2e}"
);
}
#[test]
fn test_lightgbm_zero_as_missing_parity() {
let model =
LgbModel::load(Path::new("tests/fixtures/tiny_zero.lgb"))
.expect(
"Failed to load zero-as-missing fixture model",
);
let n_zero_splits: usize = model
.trees
.iter()
.flat_map(|t| &t.decision_type)
.filter(|&&dt| (dt >> 2) & 3 == 1)
.count();
assert!(
n_zero_splits > 0,
"fixture must contain missing_type=Zero splits"
);
let cases: [([f64; 5], f64); 5] = [
([0.0, 0.3, 0.0, -0.5, 1.2], 0.30046007383088824),
([1e-40, 0.3, 1e-40, -0.5, 1.2], 0.30046007383088824),
([1e-30, 0.3, 1e-30, -0.5, 1.2], 0.3412099306576426),
([0.5, -0.2, 0.7, 1.1, -0.9], 0.9178425581052003),
([0.0, 0.0, 0.0, 0.0, 0.0], 0.5302747959214583),
];
for (features, expected) in cases {
let pred = model.predict(&features).unwrap();
let diff = (pred - expected).abs();
assert!(
diff < 1e-9,
"zero-as-missing mismatch on {features:?}: rust={pred}, python={expected}, diff={diff:.2e}"
);
}
}
#[test]
fn test_lightgbm_sqrt_parity() {
let model =
LgbModel::load(Path::new("tests/fixtures/tiny_sqrt.lgb"))
.expect("Failed to load sqrt fixture model");
assert_eq!(model.objective, Objective::Sqrt);
let cases: [([f64; 5], f64); 3] = [
([0.5, -0.2, 0.7, 1.1, -0.9], 0.6882935825372578),
([-1.0, 0.3, -0.4, 0.2, 0.6], -1.305067557709267),
([0.0, 0.0, 0.0, 0.0, 0.0], -0.04682388288592408),
];
for (features, expected) in cases {
let pred = model.predict(&features).unwrap();
let diff = (pred - expected).abs();
assert!(
diff < 1e-9,
"reg_sqrt mismatch on {features:?}: rust={pred}, python={expected}, diff={diff:.2e}"
);
}
}
#[test]
fn test_lightgbm_categorical_parity() {
let model_path = Path::new("tests/fixtures/tiny_cat.lgb");
let model = LgbModel::load(model_path)
.expect("Failed to load categorical fixture model");
let n_cat_splits: usize = model
.trees
.iter()
.flat_map(|t| &t.decision_type)
.filter(|&&dt| dt & 1 != 0)
.count();
assert!(
n_cat_splits > 0,
"fixture must contain categorical splits"
);
for (features, expected) in CAT_FIXTURE_CASES {
let pred = model.predict(features).unwrap();
let diff = (pred - expected).abs();
assert!(
diff < 1e-9,
"categorical mismatch on {features:?}: rust={pred}, python={expected}, diff={diff:.2e}"
);
}
}
const CAT_FIXTURE_CASES: &[([f64; 5], f64)] = &[
([0.5, -0.2, 0.7, 1.0, 3.0], CAT_EXPECTED[0]),
([0.5, -0.2, 0.7, 15.0, 0.0], CAT_EXPECTED[1]),
([-1.0, 0.3, -0.4, 29.0, 11.0], CAT_EXPECTED[2]),
([-1.0, 0.3, -0.4, 40.0, 25.0], CAT_EXPECTED[3]), ([0.0, 0.0, 0.0, -1.0, -2.0], CAT_EXPECTED[4]), ([0.0, 0.0, 0.0, f64::NAN, f64::NAN], CAT_EXPECTED[5]), ];
const CAT_EXPECTED: [f64; 6] = [
0.973622176489283, 0.9740570762683468, 0.02734900954762666,
0.027716944663299832, 0.038079953138787107,
0.038079953138787107,
];
fn one_split(left: i32, right: i32) -> String {
format!(
"Tree=0\nnum_leaves=2\nsplit_feature=0\nthreshold=0.5\n\
decision_type=0\nleft_child={left}\nright_child={right}\n\
leaf_value=0.1 0.2\n"
)
}
#[test]
fn test_cyclic_tree_is_rejected() {
let err = LgbModel::parse(&one_split(0, 0)).unwrap_err();
assert!(matches!(err, Error::Parse { .. }), "got {err:?}");
}
#[test]
fn test_unreachable_leaf_is_rejected() {
let err = LgbModel::parse(&one_split(-1, -1)).unwrap_err();
assert!(matches!(err, Error::Parse { .. }), "got {err:?}");
}
#[test]
fn test_feature_count_is_checked() {
let model = LgbModel::load(Path::new(
"tests/fixtures/tiny_binary.lgb",
))
.unwrap();
let err = model.predict(&[0.5, -0.2, 0.7, 1.1]).unwrap_err();
assert!(
matches!(
err,
Error::FeatureCount {
expected: 5,
got: 4
}
),
"got {err:?}"
);
let lenient = model.predict_unchecked(&[0.5, -0.2, 0.7, 1.1]);
assert!(lenient.is_finite(), "got {lenient}");
assert_eq!(
lenient,
model.predict(&[0.5, -0.2, 0.7, 1.1, f64::NAN]).unwrap()
);
}
#[test]
fn test_huge_feature_indices_are_rejected() {
let huge_header = format!(
"max_feature_idx={}\nTree=0\nnum_leaves=1\nleaf_value=0.5\n",
usize::MAX
);
let huge_split = one_split(-1, -2).replace(
"split_feature=0",
&format!("split_feature={}", usize::MAX),
);
for text in [huge_header, huge_split] {
let err = LgbModel::parse(&text).unwrap_err();
assert!(
matches!(err, Error::Parse { .. }),
"got {err:?}"
);
}
}
#[test]
fn test_split_feature_beyond_header_bound_is_rejected() {
let text =
format!("max_feature_idx=0\n{}", one_split(-1, -2));
LgbModel::parse(&text).unwrap();
let bad = text.replace("split_feature=0", "split_feature=3");
let err = LgbModel::parse(&bad).unwrap_err();
assert!(matches!(err, Error::Parse { .. }), "got {err:?}");
}
#[test]
fn test_predict_batch_default_impl() {
let model = LgbModel::load(Path::new(
"tests/fixtures/tiny_binary.lgb",
))
.unwrap();
let a = [0.5, -0.2, 0.7, 1.1, -0.9];
let b = [-1.0, 0.3, -0.4, 0.2, 0.6];
let flat: Vec<f64> =
a.iter().chain(b.iter()).copied().collect();
let batch = model.predict_batch(&flat, 5).unwrap();
assert_eq!(batch.len(), 2);
assert_eq!(batch[0], model.predict(&a).unwrap());
assert_eq!(batch[1], model.predict(&b).unwrap());
for (flat, n) in [(&flat[..7], 5), (&flat[..], 0)] {
let err = model.predict_batch(flat, n).unwrap_err();
assert!(
matches!(err, Error::BatchShape { .. }),
"got {err:?}"
);
}
}
#[test]
fn test_parse_error_is_reported() {
let bad = "Tree=0\nnum_leaves=2\nsplit_feature=0\nthreshold=not_a_number\n\
decision_type=2\nleft_child=-1\nright_child=-2\nleaf_value=0.1 0.2\n";
let err = LgbModel::parse(bad).unwrap_err();
assert!(matches!(err, Error::Parse { .. }), "got {err:?}");
}
#[test]
fn test_empty_model_is_reported() {
assert!(matches!(
LgbModel::parse("objective=binary sigmoid:1\n")
.unwrap_err(),
Error::EmptyModel
));
}
}