#[cfg(all(feature = "mesalock_sgx", not(target_env = "sgx")))]
use std::prelude::v1::*;
use crate::config::{Config, Loss};
use crate::decision_tree::DecisionTree;
#[cfg(feature = "enable_training")]
use crate::decision_tree::TrainingCache;
use crate::decision_tree::{DataVec, PredVec, ValueType, VALUE_TYPE_MIN, VALUE_TYPE_UNKNOWN};
use crate::errors::Result;
#[cfg(feature = "enable_training")]
use crate::fitness::{label_average, logit_loss_gradient, weighted_label_median, AUC, MAE, RMSE};
#[cfg(feature = "enable_training")]
use rand::prelude::SliceRandom;
#[cfg(feature = "enable_training")]
use rand::thread_rng;
#[cfg(not(feature = "mesalock_sgx"))]
use std::fs::File;
#[cfg(feature = "mesalock_sgx")]
use std::untrusted::fs::File;
use std::io::prelude::*;
use std::io::BufReader;
use serde_derive::{Deserialize, Serialize};
#[cfg(feature = "profiling")]
use std::time::Instant;
#[derive(Default, Serialize, Deserialize)]
pub struct GBDT {
conf: Config,
trees: Vec<DecisionTree>,
bias: ValueType,
}
impl GBDT {
pub fn new(conf: &Config) -> GBDT {
GBDT {
conf: conf.clone(),
trees: Vec::new(),
bias: 0.0,
}
}
#[cfg(feature = "enable_training")]
fn check_valid_data(&self, dv: &DataVec) -> bool {
dv.iter().all(|x| x.feature.len() == self.conf.feature_size)
}
#[cfg(feature = "enable_training")]
fn init(&mut self, len: usize, dv: &DataVec) {
assert!(dv.len() >= len);
if !self.check_valid_data(dv) {
panic!("There are invalid data in data vector, check your data please.");
}
if self.conf.initial_guess_enabled {
return;
}
self.bias = match self.conf.loss {
Loss::SquaredError => label_average(dv, len),
Loss::LogLikelyhood => {
let v: ValueType = label_average(dv, len);
((1.0 + v) / (1.0 - v)).ln() / 2.0
}
Loss::LAD => weighted_label_median(dv, len),
_ => label_average(dv, len),
}
}
#[cfg(feature = "enable_training")]
pub fn fit(&mut self, train_data: &mut DataVec) {
self.trees = Vec::with_capacity(self.conf.iterations);
for i in 0..self.conf.iterations {
self.trees.push(DecisionTree::new());
self.trees[i].set_feature_size(self.conf.feature_size);
self.trees[i].set_max_depth(self.conf.max_depth);
self.trees[i].set_min_leaf_size(self.conf.min_leaf_size);
self.trees[i].set_feature_sample_ratio(self.conf.feature_sample_ratio);
self.trees[i].set_loss(self.conf.loss.clone());
}
let nr_samples: usize = if self.conf.data_sample_ratio < 1.0 {
((train_data.len() as f64) * self.conf.data_sample_ratio) as usize
} else {
train_data.len()
};
self.init(train_data.len(), train_data);
let mut rng = thread_rng();
let mut predicted_cache: PredVec = self.predict_n(train_data, 0, 0, train_data.len());
#[cfg(feature = "profiling")]
let t1 = Instant::now();
let mut cache = TrainingCache::get_cache(
self.conf.feature_size,
train_data,
self.conf.training_optimization_level,
);
#[cfg(feature = "profiling")]
let t2 = Instant::now();
#[cfg(feature = "profiling")]
println!("cache {:?}", t2 - t1);
for i in 0..self.conf.iterations {
#[cfg(feature = "profiling")]
let t1 = Instant::now();
let mut samples: Vec<usize> = (0..train_data.len()).collect();
let (subset, remaining) = if nr_samples < train_data.len() {
samples.shuffle(&mut rng);
let (left, right) = samples.split_at(nr_samples);
let mut left = left.to_vec();
let mut right = right.to_vec();
left.sort();
right.sort();
(left, right)
} else {
(samples, Vec::new())
};
match self.conf.loss {
Loss::SquaredError => {
self.square_loss_process(train_data, train_data.len(), &predicted_cache)
}
Loss::LogLikelyhood => {
self.log_loss_process(train_data, train_data.len(), &predicted_cache)
}
Loss::LAD => self.lad_loss_process(train_data, train_data.len(), &predicted_cache),
_ => self.square_loss_process(train_data, train_data.len(), &predicted_cache),
}
self.trees[i].fit_n(train_data, &subset, &mut cache);
let train_preds = cache.get_preds();
for index in subset.iter() {
predicted_cache[*index] += train_preds[*index] * self.conf.shrinkage;
}
let predicted_tmp = self.trees[i].predict_n(train_data, &remaining);
for index in remaining.iter() {
predicted_cache[*index] += predicted_tmp[*index] * self.conf.shrinkage;
}
#[cfg(feature = "profiling")]
let t2 = Instant::now();
#[cfg(feature = "profiling")]
println!(
"iteration {} {:?} nodes: {}",
i,
t2 - t1,
self.trees[i].len()
);
}
}
fn predict_n(&self, test_data: &DataVec, begin: usize, iters: usize, n: usize) -> PredVec {
assert!((begin + iters) <= self.trees.len());
assert!(n <= test_data.len());
if self.trees.is_empty() {
return vec![VALUE_TYPE_UNKNOWN; test_data.len()];
}
let mut predicted: PredVec = if !self.conf.initial_guess_enabled {
vec![self.bias; n]
} else {
test_data.iter().take(n).map(|x| x.initial_guess).collect()
};
let subset: Vec<usize> = (0..n).collect();
for i in begin..(iters + begin) {
let v: PredVec = self.trees[i].predict_n(test_data, &subset);
for (e, v) in predicted.iter_mut().take(n).zip(v.iter()) {
*e += self.conf.shrinkage * v;
}
}
predicted
}
pub fn predict(&self, test_data: &DataVec) -> PredVec {
assert_eq!(self.conf.iterations, self.trees.len());
let predicted = self.predict_n(test_data, 0, self.conf.iterations, test_data.len());
match self.conf.loss {
Loss::LogLikelyhood => predicted
.iter()
.map(|x| {
1.0 / (1.0 + ((-2.0 * x).exp()))
})
.collect(),
Loss::BinaryLogistic | Loss::RegLogistic => {
predicted.iter().map(|x| 1.0 / (1.0 + (-x).exp())).collect()
}
_ => predicted,
}
}
pub fn predict_multiclass(
&self,
test_data: &DataVec,
class_num: usize,
) -> (Vec<usize>, Vec<Vec<ValueType>>) {
assert_eq!(self.conf.iterations, self.trees.len());
assert_eq!(self.trees.len() % class_num, 0);
let mut probs: Vec<Vec<ValueType>> = Vec::with_capacity(test_data.len());
for _index in 0..test_data.len() {
probs.push(vec![self.bias; class_num]);
}
for (index, tree) in self.trees.iter().enumerate() {
let preds = tree.predict(test_data);
for (x, y) in probs.iter_mut().zip(preds.iter()) {
x[index % class_num] += y;
}
}
let mut labels = vec![0; test_data.len()];
for (elem_index, elem) in probs.iter_mut().enumerate() {
let mut sum: ValueType = 0.0;
let mut max_value = VALUE_TYPE_MIN;
let mut max_index = 0;
let mut prob_vec = vec![0.0; class_num];
for (index, item) in elem.iter().enumerate() {
let v = item.exp();
prob_vec[index] = v;
sum += v;
if v > max_value {
max_index = index;
max_value = v;
}
}
for item in prob_vec.iter_mut() {
*item /= sum;
}
*elem = prob_vec;
labels[elem_index] = max_index;
}
(labels, probs)
}
pub fn print_trees(&self) {
for i in 0..self.trees.len() {
self.trees[i].print();
}
}
#[cfg(feature = "enable_training")]
fn square_loss_process(&self, dv: &mut DataVec, samples: usize, predicted: &PredVec) {
for i in 0..samples {
dv[i].target = dv[i].label - predicted[i];
}
if self.conf.debug {
println!("RMSE = {}", RMSE(dv, predicted, samples));
}
}
#[cfg(feature = "enable_training")]
fn log_loss_process(&self, dv: &mut DataVec, samples: usize, predicted: &PredVec) {
for i in 0..samples {
dv[i].target = logit_loss_gradient(dv[i].label, predicted[i]);
}
if self.conf.debug {
let normalized_preds = predicted
.iter()
.map(|x| 1.0 / (1.0 + ((-2.0 * x).exp())))
.collect();
println!("AUC = {}", AUC(dv, &normalized_preds, dv.len()));
}
}
#[cfg(feature = "enable_training")]
fn lad_loss_process(&self, dv: &mut DataVec, samples: usize, predicted: &PredVec) {
for i in 0..samples {
dv[i].residual = dv[i].label - predicted[i];
dv[i].target = if dv[i].residual >= 0.0 { 1.0 } else { -1.0 };
}
if self.conf.debug {
println!("MAE {}", MAE(dv, predicted, samples));
}
}
pub fn save_model(&self, filename: &str) -> Result<()> {
let mut file = File::create(filename)?;
let serialized = serde_json::to_string(self)?;
file.write_all(serialized.as_bytes())?;
Ok(())
}
pub fn load_model(filename: &str) -> Result<Self> {
let mut file = File::open(filename)?;
let mut contents = String::new();
file.read_to_string(&mut contents)?;
let ret: Self = serde_json::from_str(&contents)?;
Ok(ret)
}
pub fn from_xgboost_dump(model_file: &str, objective: &str) -> Result<Self> {
let tree_file = File::open(model_file)?;
let reader = BufReader::new(tree_file);
Self::from_xgboost_reader(reader, objective)
}
pub fn from_xgboost_reader<R>(reader: R, objective: &str) -> Result<Self>
where
R: std::io::BufRead,
{
let mut all_lines: Vec<String> = Vec::new();
let mut has_read_score = false;
let mut base_score: ValueType = 0.0;
for line in reader.lines() {
if !has_read_score {
has_read_score = true;
base_score = line?.parse::<ValueType>()?;
continue;
}
let value: String = line?;
all_lines.push(value);
}
let single_line = all_lines.join("");
let json_obj: serde_json::Value = serde_json::from_str(&single_line)?;
let nodes = json_obj.as_array().ok_or("parse trees error")?;
let mut cfg = Config::new();
cfg.set_loss(objective);
cfg.set_iterations(nodes.len());
cfg.shrinkage = 1.0;
let mut gbdt = GBDT::new(&cfg);
gbdt.bias = base_score;
for node in nodes.iter() {
let tree = DecisionTree::get_from_xgboost(node)?;
gbdt.trees.push(tree);
}
Ok(gbdt)
}
}