5#include <MedProcessTools/MedProcessTools/MedProcessUtils.h>
6#include <xgboost/c_api.h>
7#include "MedProcessTools/MedProcessTools/MedSamples.h"
18 vector<string> eval_metric;
22 float colsample_bytree;
23 float colsample_bylevel;
25 float scale_pos_weight;
32 string split_penalties;
33 string monotone_constraints;
38 objective =
"binary:logistic";
45 eval_metric.push_back(
"auc");
46 missing_value = MED_MAT_MISSING_VALUE;
48 colsample_bytree = 1.0;
49 colsample_bylevel = 1.0;
51 scale_pos_weight = 1.0;
61 ADD_SERIALIZATION_FUNCS(booster, objective, eta, gamma, min_child_weight, max_depth, num_round, eval_metric, silent, missing_value, num_class,
62 colsample_bytree, colsample_bylevel, subsample, scale_pos_weight, tree_method, lambda, alpha, seed, verbose_eval, validate_frac, split_penalties, monotone_constraints)
68 BoosterHandle my_learner = NULL;
71 int feat_contrib_flags = 0;
72 virtual int init(
void *classifier_params)
79 virtual int set_params(map<string, string> &initialization_map);
82 MedXGB() { init_defaults(); };
85 int Learn(
float *x,
float *y,
const float *w,
int nsamples,
int nftrs);
86 int Learn(
float *x,
float *y,
int nsamples,
int nftrs);
87 int Predict(
float *x,
float *&preds,
int nsamples,
int nftrs)
const;
88 void prepare_mat_handle(
float *x,
float *y,
const float *w,
int nsamples,
int nftrs, DMatrixHandle &matrix_handle);
90 virtual void print(FILE *fp,
const string &prefix,
int level = 0)
const;
92 void calc_feature_importance(vector<float> &features_importance_scores,
93 const string &general_params,
const MedFeatures *features);
99 void export_predictor(
const string &output_fname);
103 void pre_serialization()
105 const char *out_dptr;
107 string cfg_js =
"{ \"format\":\"json\" }";
108 if (my_learner != NULL)
110 if (XGBoosterSaveModelToBuffer(my_learner, cfg_js.c_str(), &len, &out_dptr) != 0)
111 throw runtime_error(
"failed XGBoosterSaveModelToBuffer\n");
112 serial_xgb.resize(len);
113 memcpy(&serial_xgb[0], out_dptr, len);
119 void post_deserialization()
121 if (this->my_learner != NULL)
122 XGBoosterFree(this->my_learner);
123 if (!serial_xgb.empty())
125 DMatrixHandle h_train_empty[1];
126 if (XGBoosterCreate(h_train_empty, 0, &my_learner) != 0)
127 throw runtime_error(
"failed XGBoosterCreate\n");
128 if (XGBoosterLoadModelFromBuffer(my_learner, &serial_xgb[0], serial_xgb.size()) != 0)
129 throw runtime_error(
"failed XGBoosterLoadModelFromBuffer\n");
134 void prepare_predict_single();
135 void predict_single(
const vector<float> &x, vector<float> &preds)
const;
137 void get_json(
const char ***json,
int &len,
string type)
139 if (my_learner != NULL)
143 int succ = XGBoosterDumpModelEx(my_learner, no_fmap.c_str(), 1, type.c_str(), &_len, json);
145 HMTHROW_AND_ERR(
"Error MedXGB::get_json - can't get model\n");
156 bool _mark_learn_done;
157 bool prepared_single;
158 vector<BoosterHandle> learner_per_thread;
160 void translate_split_penalties(
string &split_penalties_s);
161 void translate_monotone_constraints(
string &monotone_constraints_s);
162 void calc_feature_importance_local(vector<float> &features_importance_scores,
string &importance_type);
163 vector<char> serial_xgb;
MedAlgo - APIs to different algorithms: Linear Models, RF, GBM, KNN, and more.
#define ADD_SERIALIZATION_FUNCS(...)
Definition SerializableObject.h:156
#define MEDSERIALIZE_SUPPORT(Type)
Definition SerializableObject.h:142
A class for holding features data as a virtual matrix
Definition MedFeatures.h:47
Base Interface for predictor.
Definition MedAlgo.h:72
int features_count
The model features count used in Learn, to validate when caling predict.
Definition MedAlgo.h:90
MedPredictorTypes classifier_type
The Predicotr enum type.
Definition MedAlgo.h:74
vector< string > model_features
The model features used in Learn, to validate when caling predict.
Definition MedAlgo.h:87
virtual int set_params(map< string, string > &initialization_map)
The parsed fields from init command.
Definition MedXGB.cpp:488
void calc_feature_contribs(MedMat< float > &x, MedMat< float > &contribs)
Feature contributions explains the prediction on each sample (aka BUT_WHY)
Definition MedXGB.cpp:83
int Predict(float *x, float *&preds, int nsamples, int nftrs) const
Predict should be implemented for each model.
Definition MedXGB.cpp:62
int n_preds_per_sample() const
Number of predictions per sample. typically 1 - but some models return several per sample (for exampl...
Definition MedXGB.cpp:38
int Learn(float *x, float *y, const float *w, int nsamples, int nftrs)
Learn should be implemented for each model.
Definition MedXGB.cpp:179
Definition SerializableObject.h:33