1#![allow(unused)]
2
3use serde::Deserialize;
4use std::{fs::File, io::BufReader, path::Path};
5use thiserror::Error;
6
7#[derive(Debug, Deserialize)]
8#[serde(rename_all = "snake_case")]
9pub struct CatBoost {
10 features_info: Features,
11 oblivious_trees: Vec<ObliviousTree>,
12 scale_and_bias: (f32, Vec<f32>),
13}
14
15impl CatBoost {
16 pub fn load(path: &Path) -> Result<Self, std::io::Error> {
17 let file = File::open(path)?;
18 let reader = BufReader::new(file);
19 let model: CatBoost = serde_json::from_reader(reader)?;
20 Ok(model)
21 }
22
23 pub fn try_from_json(model_str: &str) -> Result<Self, serde_json::Error> {
24 let model: CatBoost = serde_json::from_str(model_str)?;
25 Ok(model)
26 }
27
28 fn num_features(&self) -> usize {
29 self.features_info
30 .float_features
31 .iter()
32 .map(|f| f.flat_feature_index)
33 .max()
34 .map_or(0, |m| m + 1)
35 }
36}
37
38#[derive(Debug, Deserialize)]
39#[serde(rename_all = "snake_case")]
40struct Features {
41 float_features: Vec<FloatFeature>,
42}
43
44#[derive(Debug, Deserialize)]
45#[serde(rename_all = "snake_case")]
46struct FloatFeature {
47 feature_index: usize,
48 flat_feature_index: usize,
49 borders: Vec<f32>,
50 has_nans: bool,
51 nan_value_treatment: NanValueTreatment,
52}
53
54#[derive(Debug, Deserialize, PartialEq, Clone, Copy)]
55enum NanValueTreatment {
56 #[serde(rename = "AsIs")]
57 Unspecified,
58 #[serde(rename = "AsTrue")]
59 Left,
60 #[serde(rename = "AsFalse")]
61 Right,
62}
63
64#[derive(Debug, Deserialize)]
65#[serde(rename_all = "snake_case")]
66struct ObliviousTree {
67 leaf_values: Vec<f32>,
68 splits: Vec<Split>,
69}
70
71#[derive(Debug, Deserialize)]
72#[serde(rename_all = "snake_case", tag = "split_type")]
73enum Split {
74 #[serde(rename = "FloatFeature")]
75 FloatFeature {
76 #[expect(unused)]
77 float_feature_index: usize,
78 split_index: usize, #[expect(unused)]
80 border: f32,
81 },
82}
83
84#[derive(Debug, Error)]
85pub enum InferenceError {
86 #[error("Incorrect number of features provided. Expected {expected}, got {actual}.")]
87 NumFeaturesMismatch { expected: usize, actual: usize },
88}
89
90impl CatBoost {
91 pub fn predict_raw(&self, features: &[f32]) -> Result<f32, InferenceError> {
95 {
97 let expected_features = self.num_features();
98
99 if features.len() != expected_features {
100 return Err(InferenceError::NumFeaturesMismatch {
101 expected: expected_features,
102 actual: features.len(),
103 });
104 }
105 }
106
107 let go_lefts = self.features_info.float_features.iter().flat_map(|FloatFeature {
108 feature_index,
109 flat_feature_index,
110 borders,
111 has_nans,
112 nan_value_treatment,
113 }| {
114 assert_eq!(*feature_index, *flat_feature_index, "This will always be true if there are only float features (i.e. no categorical features). Categorical features are not supported.");
117 let feature_value = features[*flat_feature_index];
118 borders.iter().map(move |border| {
119 if feature_value.is_nan() {
120 if !has_nans {
121 eprintln!(
122 "Warning: Encountered NaN for feature {} which had no NaNs during training. Treating as <= border.",
123 feature_index
124 );
125 false
126 } else {
127 match nan_value_treatment {
129 NanValueTreatment::Unspecified => {
130 eprintln!(
131 "Warning: Encountered NaN for feature {} with NanValueTreatment::AsIs. Treating as <= border.",
132 feature_index
133 );
134 false
135 }
136 NanValueTreatment::Left => true, NanValueTreatment::Right => false, }
139 }
140 } else {
141 feature_value > *border
143 }
144 }
145 )}).collect::<Vec<bool>>();
146
147 let logits = self
148 .oblivious_trees
149 .iter()
150 .map(|tree| {
151 assert_eq!(
152 2_usize.pow(tree.splits.len() as u32),
153 tree.leaf_values.len(),
154 "The number of leaf values must be equal to 2^number_of_splits"
155 );
156
157 let mut current_leaf_index: usize = 0;
158 let depth = tree.splits.len();
159
160 if depth == 0 {
161 if !tree.leaf_values.is_empty() {
163 return tree.leaf_values[0];
164 }
165 panic!("No leaf values!?");
167 }
168
169 for (level, Split::FloatFeature { split_index, .. }) in
170 tree.splits.iter().enumerate()
171 {
172 let go_left = go_lefts[*split_index];
173
174 current_leaf_index |= (go_left as usize) << level;
176 }
177
178 tree.leaf_values[current_leaf_index]
179 })
180 .sum::<f32>();
181
182 let scale = self.scale_and_bias.0;
184 let bias = self.scale_and_bias.1.first().unwrap_or(&0.0); let prediction = logits * scale + bias;
187
188 Ok(prediction)
189 }
190
191 pub fn predict(&self, features: &[f32]) -> Result<f32, InferenceError> {
195 let prediction = self.predict_raw(features)?;
196 let probability = 1.0 / (1.0 + (-prediction).exp());
198 Ok(probability)
199 }
200}
201
202#[test]
203fn test_against_tiny_model() {
204 let model = CatBoost::load(Path::new("models/test/tiny-binary-catboost.json")).unwrap();
205 let test_features: Vec<f32> = vec![0.1276993, 0.9918129, 0.16597846, 0.98612934];
206 let probability = model.predict(&test_features).unwrap();
207
208 assert!(
209 (probability - 0.5245).abs() < 0.01,
210 "Probability does not match expected value."
211 );
212}
213
214#[test]
215fn test_against_big_model() {
216 let model = CatBoost::load(Path::new("models/test/big-binary-catboost.json")).unwrap();
217
218 let test_features: Vec<f32> = serde_json::from_str(
220 "[-7.60986700e-04,-1.16379880e-02,-1.18961320e-02,2.97898050e-01,-1.04892480e-01,-1.98598710e-01,-1.47249590e-02,
221 1.38537230e-01,8.87154600e-02,4.81008140e-02,2.59864870e-02,-1.16422900e-01,6.40196900e-02,9.56853400e-02,-1.17455475e-01,
222 -1.70977310e-01,-1.43097770e-01,-8.89736000e-02,-1.75322740e-02,-6.27612370e-03,6.12661540e-02,2.41794800e-01,
223 5.64474700e-02,1.10313500e-01,-1.16272320e-02,-3.90173750e-03,-4.98648000e-02,3.72372570e-02,9.55132500e-02,
224 4.76275500e-02,9.35341400e-02,-7.05593400e-02,-5.39090520e-02,-8.08850900e-02,-5.20859100e-03,-1.27458550e-02,
225 -1.56865450e-01,-2.96650380e-02,9.99769850e-03,-6.87093100e-02,3.25046220e-02,1.51788620e-01,6.56115800e-02,
226 -2.34910960e-02,-9.78365400e-02,1.23909080e-02,-3.94314830e-02,-8.03257800e-02,1.14529850e-01,2.28887600e-01,
227 -1.26167830e-02,3.24831100e-02,-1.31223150e-03,1.72440130e-01,-4.61217130e-02,-5.99754340e-02,-1.60393420e-01,
228 1.37332560e-01,-1.32083640e-01,4.97787500e-02,-7.21512200e-02,-3.61616570e-02,-7.18070300e-02,8.66072800e-02,-1.83454280e-01,
229 -2.79655900e-02,-6.01045080e-02,-1.57725930e-01,1.21671826e-01,4.59065920e-02,2.10172160e-02,-9.08666550e-02,
230 -6.02335780e-02,3.82698330e-02,3.70006260e-03,-7.22372700e-02,1.00417980e-01,1.46970600e-04,
231 1.44302440e-01,4.17978000e-02,1.33804590e-01,-7.68408300e-02,-3.29993960e-02,1.02224990e-01,-1.41721010e-01,
232 1.25027700e-01,-1.29502190e-01,-5.90719320e-02,-7.84757500e-02,-6.27289700e-02,-2.28199210e-01,1.31739440e-01,-2.71051100e-02,
233 -5.61463000e-02,1.48174600e-01,1.09539060e-01,7.42163700e-02,-1.00729900e-02,3.70221400e-02,7.27535560e-02,
234 -8.97094300e-02,6.24005540e-03,1.35485190e-01,-7.96700800e-02,-1.05367130e-01,-4.21836970e-02,9.26107100e-02,
235 4.85388820e-02,-6.13413560e-02,-1.53906020e-01,3.15686950e-02,-1.97217990e-02,-5.83019200e-02,-3.23515800e-02,
236 3.77396700e-02,6.65912900e-02,-9.17817700e-02,-3.21443450e-02,-8.52423800e-02,5.77953460e-02,-5.77492940e-02,2.00211370e-02,
237 5.01507040e-02,1.35945710e-01,-1.11538110e-01,-5.57690560e-02,-5.82558660e-02,1.98576520e-01,8.29858260e-02,-2.80917600e-01,
238 1.01130344e-01,-1.45340340e-01,4.36113100e-02,-3.87125200e-04,5.07954320e-02,1.22406400e-01,9.71698700e-02,
239 7.49267200e-02,-1.03064530e-01,-1.75918900e-01,1.06288180e-01,-2.05576430e-01,1.21945880e-01,-3.35259070e-02,
240 -4.77099460e-02,4.69270570e-02,-1.01775070e-01,8.87078000e-03,1.51603420e-01,1.30879980e-01,-1.06380284e-01,
241 1.34356920e-02,-2.48450920e-02,1.33781270e-02,-2.55128460e-02,-2.23467670e-02,-1.91116090e-03,1.12735465e-01,-6.37821200e-02,
242 5.47559100e-02,1.28301070e-01,-7.57556560e-02,-1.68233970e-03,8.44134500e-02,-8.63936840e-02,1.58879640e-01,
243 3.69855670e-04,2.59042890e-02,-7.91520000e-03,6.05584700e-02,-1.23074160e-02,1.17248570e-01,4.87691420e-02,
244 -5.97755870e-02,1.03893470e-01,-9.91846400e-03,-3.04404180e-02,1.81353050e-01,1.82337410e-03,-1.62103290e-02,
245 3.71870470e-02,-6.09729400e-02,-6.26768700e-02,-1.42024580e-01,-1.39169350e-01,1.43498240e-01,-3.35811700e-01,
246 -9.47751550e-02,2.49141700e-02,1.44258100e-02,-2.00787020e-02,1.68550580e-01,-5.71333500e-04,3.26739440e-02,
247 1.65511130e-01,3.88679470e-02,-8.53114600e-03,6.17558250e-02,4.11244970e-02,2.50339060e-01
248 ]",
249 )
250 .unwrap();
251
252 let probability = model.predict(&test_features).unwrap();
253
254 assert!(
255 (probability - 0.74518714).abs() < 0.01,
256 "Probability does not match expected value."
257 );
258}