Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
30 changes: 15 additions & 15 deletions include/treelite/logging.h
Original file line number Diff line number Diff line change
Expand Up @@ -62,25 +62,25 @@ DEFINE_CHECK_FUNC(_NE, !=)
#pragma GCC diagnostic pop


#define CHECK_BINARY_OP(name, op, x, y) \
#define TREELITE_CHECK_BINARY_OP(name, op, x, y) \
if (auto __treelite__log__err = ::treelite::LogCheck##name(x, y)) \
::treelite::LogMessageFatal(__FILE__, __LINE__).stream() \
::treelite::LogMessageFatal(__FILE__, __LINE__).stream() \
<< "Check failed: " << #x " " #op " " #y << *__treelite__log__err << ": "
#define CHECK(x) \
#define TREELITE_CHECK(x) \
if (!(x)) \
::treelite::LogMessageFatal(__FILE__, __LINE__).stream() \
::treelite::LogMessageFatal(__FILE__, __LINE__).stream() \
<< "Check failed: " #x << ": "
#define CHECK_LT(x, y) CHECK_BINARY_OP(_LT, <, x, y)
#define CHECK_GT(x, y) CHECK_BINARY_OP(_GT, >, x, y)
#define CHECK_LE(x, y) CHECK_BINARY_OP(_LE, <=, x, y)
#define CHECK_GE(x, y) CHECK_BINARY_OP(_GE, >=, x, y)
#define CHECK_EQ(x, y) CHECK_BINARY_OP(_EQ, ==, x, y)
#define CHECK_NE(x, y) CHECK_BINARY_OP(_NE, !=, x, y)

#define LOG_INFO ::treelite::LogMessage(__FILE__, __LINE__)
#define LOG_ERROR LOG_INFO
#define LOG_FATAL ::treelite::LogMessageFatal(__FILE__, __LINE__)
#define LOG(severity) LOG_##severity.stream()
#define TREELITE_CHECK_LT(x, y) TREELITE_CHECK_BINARY_OP(_LT, <, x, y)
#define TREELITE_CHECK_GT(x, y) TREELITE_CHECK_BINARY_OP(_GT, >, x, y)
#define TREELITE_CHECK_LE(x, y) TREELITE_CHECK_BINARY_OP(_LE, <=, x, y)
#define TREELITE_CHECK_GE(x, y) TREELITE_CHECK_BINARY_OP(_GE, >=, x, y)
#define TREELITE_CHECK_EQ(x, y) TREELITE_CHECK_BINARY_OP(_EQ, ==, x, y)
#define TREELITE_CHECK_NE(x, y) TREELITE_CHECK_BINARY_OP(_NE, !=, x, y)

#define TREELITE_LOG_INFO ::treelite::LogMessage(__FILE__, __LINE__)
#define TREELITE_LOG_ERROR TREELITE_LOG_INFO
#define TREELITE_LOG_FATAL ::treelite::LogMessageFatal(__FILE__, __LINE__)
#define TREELITE_LOG(severity) TREELITE_LOG_##severity.stream()

class DateLogger {
public:
Expand Down
6 changes: 3 additions & 3 deletions include/treelite/predictor.h
Original file line number Diff line number Diff line change
Expand Up @@ -152,7 +152,7 @@ class Predictor {
* \return length of prediction array
*/
inline size_t QueryResultSize(const DMatrix* dmat) const {
CHECK(pred_func_) << "A shared library needs to be loaded first using Load()";
TREELITE_CHECK(pred_func_) << "A shared library needs to be loaded first using Load()";
return dmat->GetNumRow() * num_class_;
}
/*!
Expand All @@ -164,8 +164,8 @@ class Predictor {
* \return length of prediction array
*/
inline size_t QueryResultSize(const DMatrix* dmat, size_t rbegin, size_t rend) const {
CHECK(pred_func_) << "A shared library needs to be loaded first using Load()";
CHECK(rbegin < rend && rend <= dmat->GetNumRow());
TREELITE_CHECK(pred_func_) << "A shared library needs to be loaded first using Load()";
TREELITE_CHECK(rbegin < rend && rend <= dmat->GetNumRow());
return (rend - rbegin) * num_class_;
}
/*!
Expand Down
2 changes: 1 addition & 1 deletion runtime/java/treelite4j/src/native/treelite4j.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -251,7 +251,7 @@ Java_ml_dmlc_treelite4j_java_TreeliteJNI_TreelitePredictorPredictBatchWithUInt32
API_BEGIN();
PredictorHandle predictor = reinterpret_cast<PredictorHandle>(jpredictor);
DMatrixHandle dmat = reinterpret_cast<DMatrixHandle>(jbatch);
CHECK_EQ(sizeof(jint), sizeof(uint32_t));
TREELITE_CHECK_EQ(sizeof(jint), sizeof(uint32_t));
jint* out_result = jenv->GetIntArrayElements(jout_result, nullptr);
jlong* out_result_size = jenv->GetLongArrayElements(jout_result_size, nullptr);
size_t out_result_size_tmp = 0;
Expand Down
22 changes: 11 additions & 11 deletions src/annotator.cc
Original file line number Diff line number Diff line change
Expand Up @@ -72,8 +72,8 @@ inline void ComputeBranchLoopImpl(
const size_t* count_row_ptr, uint64_t* counts_tloc) {
std::vector<Entry<ElementType>> inst(nthread * dmat->num_col, {-1});
const size_t ntree = model.trees.size();
CHECK_LE(rbegin, rend);
CHECK_LT(static_cast<int64_t>(rend), std::numeric_limits<int64_t>::max());
TREELITE_CHECK_LE(rbegin, rend);
TREELITE_CHECK_LT(static_cast<int64_t>(rend), std::numeric_limits<int64_t>::max());
const size_t num_col = dmat->num_col;
const ElementType missing_value = dmat->missing_value;
const bool nan_missing = treelite::math::CheckNAN(missing_value);
Expand All @@ -87,7 +87,7 @@ inline void ComputeBranchLoopImpl(
const size_t off2 = count_row_ptr[ntree] * tid;
for (size_t j = 0; j < num_col; ++j) {
if (treelite::math::CheckNAN(row[j])) {
CHECK(nan_missing)
TREELITE_CHECK(nan_missing)
<< "The missing_value argument must be set to NaN if there is any NaN in the matrix.";
} else if (nan_missing || row[j] != missing_value) {
inst[off + j].fvalue = row[j];
Expand All @@ -109,8 +109,8 @@ inline void ComputeBranchLoopImpl(
const size_t* count_row_ptr, uint64_t* counts_tloc) {
std::vector<Entry<ElementType>> inst(nthread * dmat->num_col, {-1});
const size_t ntree = model.trees.size();
CHECK_LE(rbegin, rend);
CHECK_LT(static_cast<int64_t>(rend), std::numeric_limits<int64_t>::max());
TREELITE_CHECK_LE(rbegin, rend);
TREELITE_CHECK_LT(static_cast<int64_t>(rend), std::numeric_limits<int64_t>::max());
const auto rbegin_i = static_cast<int64_t>(rbegin);
const auto rend_i = static_cast<int64_t>(rend);
#pragma omp parallel for schedule(static) num_threads(nthread)
Expand Down Expand Up @@ -141,7 +141,7 @@ class ComputeBranchLoopDispatcherWithDenseDMatrix {
const treelite::DMatrix* dmat, size_t rbegin, size_t rend, int nthread,
const size_t* count_row_ptr, uint64_t* counts_tloc) {
const auto* dmat_ = static_cast<const treelite::DenseDMatrixImpl<ElementType>*>(dmat);
CHECK(dmat_) << "Dangling data matrix reference detected";
TREELITE_CHECK(dmat_) << "Dangling data matrix reference detected";
ComputeBranchLoopImpl(model, dmat_, rbegin, rend, nthread, count_row_ptr, counts_tloc);
}
};
Expand All @@ -155,7 +155,7 @@ class ComputeBranchLoopDispatcherWithCSRDMatrix {
const treelite::DMatrix* dmat, size_t rbegin, size_t rend, int nthread,
const size_t* count_row_ptr, uint64_t* counts_tloc) {
const auto* dmat_ = static_cast<const treelite::CSRDMatrixImpl<ElementType>*>(dmat);
CHECK(dmat_) << "Dangling data matrix reference detected";
TREELITE_CHECK(dmat_) << "Dangling data matrix reference detected";
ComputeBranchLoopImpl(model, dmat_, rbegin, rend, nthread, count_row_ptr, counts_tloc);
}
};
Expand All @@ -177,7 +177,7 @@ inline void ComputeBranchLoop(const treelite::ModelImpl<ThresholdType, LeafOutpu
break;
}
default:
LOG(FATAL)
TREELITE_LOG(FATAL)
<< "Annotator does not support DMatrix of type " << static_cast<int>(dmat->GetType());
break;
}
Expand Down Expand Up @@ -214,7 +214,7 @@ AnnotateImpl(
const size_t rend = std::min(rbegin + pstep, num_row);
ComputeBranchLoop(model, dmat, rbegin, rend, nthread, &count_row_ptr[0], &counts_tloc[0]);
if (verbose > 0) {
LOG(INFO) << rend << " of " << num_row << " rows processed";
TREELITE_LOG(INFO) << rend << " of " << num_row << " rows processed";
}
}

Expand Down Expand Up @@ -249,10 +249,10 @@ BranchAnnotator::Load(std::istream& fi) {
doc.ParseStream(is);

std::string err_msg = "JSON file must contain a list of lists of integers";
CHECK(doc.IsArray()) << err_msg;
TREELITE_CHECK(doc.IsArray()) << err_msg;
counts_.clear();
for (const auto& node_cnt : doc.GetArray()) {
CHECK(node_cnt.IsArray()) << err_msg;
TREELITE_CHECK(node_cnt.IsArray()) << err_msg;
counts_.emplace_back();
for (const auto& e : node_cnt.GetArray()) {
counts_.back().push_back(e.GetUint64());
Expand Down
52 changes: 26 additions & 26 deletions src/c_api/c_api.cc
Original file line number Diff line number Diff line change
Expand Up @@ -30,7 +30,7 @@ int TreeliteAnnotateBranch(
std::unique_ptr<BranchAnnotator> annotator{new BranchAnnotator()};
const Model* model_ = static_cast<Model*>(model);
const auto* dmat_ = static_cast<const DMatrix*>(dmat);
CHECK(dmat_) << "Found a dangling reference to DMatrix";
TREELITE_CHECK(dmat_) << "Found a dangling reference to DMatrix";
annotator->Annotate(*model_, dmat_, nthread, verbose);
*out = static_cast<AnnotationHandle>(annotator.release());
API_END();
Expand Down Expand Up @@ -64,8 +64,8 @@ int TreeliteCompilerGenerateCodeV2(CompilerHandle compiler,
API_BEGIN();
const Model* model_ = static_cast<Model*>(model);
Compiler* compiler_ = static_cast<Compiler*>(compiler);
CHECK(model_);
CHECK(compiler_);
TREELITE_CHECK(model_);
TREELITE_CHECK(compiler_);
compiler::CompilerParam param = compiler_->QueryParam();

// create directory named dirpath
Expand All @@ -75,12 +75,12 @@ int TreeliteCompilerGenerateCodeV2(CompilerHandle compiler,
/* compile model */
auto compiled_model = compiler_->Compile(*model_);
if (param.verbose > 0) {
LOG(INFO) << "Code generation finished. Writing code to files...";
TREELITE_LOG(INFO) << "Code generation finished. Writing code to files...";
}

for (const auto& it : compiled_model.files) {
if (param.verbose > 0) {
LOG(INFO) << "Writing file " << it.first << "...";
TREELITE_LOG(INFO) << "Writing file " << it.first << "...";
}
const std::string filename_full = dirpath_ + "/" + it.first;
if (it.second.is_binary) {
Expand Down Expand Up @@ -189,7 +189,7 @@ int TreeliteLoadSKLearnGradientBoostingClassifier(
int TreeliteSerializeModel(const char* filename, ModelHandle handle) {
API_BEGIN();
FILE* fp = std::fopen(filename, "wb");
CHECK(fp) << "Failed to open file '" << filename << "'";
TREELITE_CHECK(fp) << "Failed to open file '" << filename << "'";
auto* model_ = static_cast<Model*>(handle);
model_->SerializeToFile(fp);
std::fclose(fp);
Expand All @@ -199,7 +199,7 @@ int TreeliteSerializeModel(const char* filename, ModelHandle handle) {
int TreeliteDeserializeModel(const char* filename, ModelHandle* out) {
API_BEGIN();
FILE* fp = std::fopen(filename, "rb");
CHECK(fp) << "Failed to open file '" << filename << "'";
TREELITE_CHECK(fp) << "Failed to open file '" << filename << "'";
std::unique_ptr<Model> model = Model::DeserializeFromFile(fp);
std::fclose(fp);
*out = static_cast<ModelHandle>(model.release());
Expand Down Expand Up @@ -251,10 +251,10 @@ int TreeliteQueryNumClass(ModelHandle handle, size_t* out) {

int TreeliteSetTreeLimit(ModelHandle handle, size_t limit) {
API_BEGIN();
CHECK_GT(limit, 0) << "limit should be greater than 0!";
TREELITE_CHECK_GT(limit, 0) << "limit should be greater than 0!";
auto* model_ = static_cast<Model*>(handle);
const size_t num_tree = model_->GetNumTree();
CHECK_GE(num_tree, limit) << "Model contains less trees(" << num_tree << ") than limit";
TREELITE_CHECK_GE(num_tree, limit) << "Model contains fewer trees(" << num_tree << ") than limit";
model_->SetTreeLimit(limit);
API_END();
}
Expand Down Expand Up @@ -293,23 +293,23 @@ int TreeliteDeleteTreeBuilder(TreeBuilderHandle handle) {
int TreeliteTreeBuilderCreateNode(TreeBuilderHandle handle, int node_key) {
API_BEGIN();
auto* builder = static_cast<frontend::TreeBuilder*>(handle);
CHECK(builder) << "Detected dangling reference to deleted TreeBuilder object";
TREELITE_CHECK(builder) << "Detected dangling reference to deleted TreeBuilder object";
builder->CreateNode(node_key);
API_END();
}

int TreeliteTreeBuilderDeleteNode(TreeBuilderHandle handle, int node_key) {
API_BEGIN();
auto* builder = static_cast<frontend::TreeBuilder*>(handle);
CHECK(builder) << "Detected dangling reference to deleted TreeBuilder object";
TREELITE_CHECK(builder) << "Detected dangling reference to deleted TreeBuilder object";
builder->DeleteNode(node_key);
API_END();
}

int TreeliteTreeBuilderSetRootNode(TreeBuilderHandle handle, int node_key) {
API_BEGIN();
auto* builder = static_cast<frontend::TreeBuilder*>(handle);
CHECK(builder) << "Detected dangling reference to deleted TreeBuilder object";
TREELITE_CHECK(builder) << "Detected dangling reference to deleted TreeBuilder object";
builder->SetRootNode(node_key);
API_END();
}
Expand All @@ -319,7 +319,7 @@ int TreeliteTreeBuilderSetNumericalTestNode(
ValueHandle threshold, int default_left, int left_child_key, int right_child_key) {
API_BEGIN();
auto* builder = static_cast<frontend::TreeBuilder*>(handle);
CHECK(builder) << "Detected dangling reference to deleted TreeBuilder object";
TREELITE_CHECK(builder) << "Detected dangling reference to deleted TreeBuilder object";
builder->SetNumericalTestNode(node_key, feature_id, opname,
*static_cast<const frontend::Value*>(threshold),
(default_left != 0), left_child_key, right_child_key);
Expand All @@ -332,10 +332,10 @@ int TreeliteTreeBuilderSetCategoricalTestNode(
int left_child_key, int right_child_key) {
API_BEGIN();
auto* builder = static_cast<frontend::TreeBuilder*>(handle);
CHECK(builder) << "Detected dangling reference to deleted TreeBuilder object";
TREELITE_CHECK(builder) << "Detected dangling reference to deleted TreeBuilder object";
std::vector<uint32_t> vec(left_categories_len);
for (size_t i = 0; i < left_categories_len; ++i) {
CHECK(left_categories[i] <= std::numeric_limits<uint32_t>::max());
TREELITE_CHECK(left_categories[i] <= std::numeric_limits<uint32_t>::max());
vec[i] = static_cast<uint32_t>(left_categories[i]);
}
builder->SetCategoricalTestNode(node_key, feature_id, vec, (default_left != 0),
Expand All @@ -346,7 +346,7 @@ int TreeliteTreeBuilderSetCategoricalTestNode(
int TreeliteTreeBuilderSetLeafNode(TreeBuilderHandle handle, int node_key, ValueHandle leaf_value) {
API_BEGIN();
auto* builder = static_cast<frontend::TreeBuilder*>(handle);
CHECK(builder) << "Detected dangling reference to deleted TreeBuilder object";
TREELITE_CHECK(builder) << "Detected dangling reference to deleted TreeBuilder object";
builder->SetLeafNode(node_key, *static_cast<const frontend::Value*>(leaf_value));
API_END();
}
Expand All @@ -355,11 +355,11 @@ int TreeliteTreeBuilderSetLeafVectorNode(TreeBuilderHandle handle, int node_key,
const ValueHandle* leaf_vector, size_t leaf_vector_len) {
API_BEGIN();
auto* builder = static_cast<frontend::TreeBuilder*>(handle);
CHECK(builder) << "Detected dangling reference to deleted TreeBuilder object";
TREELITE_CHECK(builder) << "Detected dangling reference to deleted TreeBuilder object";
std::vector<frontend::Value> vec(leaf_vector_len);
CHECK(leaf_vector) << "leaf_vector argument must not be null";
TREELITE_CHECK(leaf_vector) << "leaf_vector argument must not be null";
for (size_t i = 0; i < leaf_vector_len; ++i) {
CHECK(leaf_vector[i]) << "leaf_vector[" << i << "] contains an empty Value handle";
TREELITE_CHECK(leaf_vector[i]) << "leaf_vector[" << i << "] contains an empty Value handle";
vec[i] = *static_cast<const frontend::Value*>(leaf_vector[i]);
}
builder->SetLeafVectorNode(node_key, vec);
Expand All @@ -381,7 +381,7 @@ int TreeliteModelBuilderSetModelParam(ModelBuilderHandle handle, const char* nam
const char* value) {
API_BEGIN();
auto* builder = static_cast<frontend::ModelBuilder*>(handle);
CHECK(builder) << "Detected dangling reference to deleted ModelBuilder object";
TREELITE_CHECK(builder) << "Detected dangling reference to deleted ModelBuilder object";
builder->SetModelParam(name, value);
API_END();
}
Expand All @@ -396,35 +396,35 @@ int TreeliteModelBuilderInsertTree(ModelBuilderHandle handle, TreeBuilderHandle
int index) {
API_BEGIN();
auto* model_builder = static_cast<frontend::ModelBuilder*>(handle);
CHECK(model_builder) << "Detected dangling reference to deleted ModelBuilder object";
TREELITE_CHECK(model_builder) << "Detected dangling reference to deleted ModelBuilder object";
auto* tree_builder = static_cast<frontend::TreeBuilder*>(tree_builder_handle);
CHECK(tree_builder) << "Detected dangling reference to deleted TreeBuilder object";
TREELITE_CHECK(tree_builder) << "Detected dangling reference to deleted TreeBuilder object";
return model_builder->InsertTree(tree_builder, index);
API_END();
}

int TreeliteModelBuilderGetTree(ModelBuilderHandle handle, int index, TreeBuilderHandle *out) {
API_BEGIN();
auto* model_builder = static_cast<frontend::ModelBuilder*>(handle);
CHECK(model_builder) << "Detected dangling reference to deleted ModelBuilder object";
TREELITE_CHECK(model_builder) << "Detected dangling reference to deleted ModelBuilder object";
auto* tree_builder = model_builder->GetTree(index);
CHECK(tree_builder) << "Detected dangling reference to deleted TreeBuilder object";
TREELITE_CHECK(tree_builder) << "Detected dangling reference to deleted TreeBuilder object";
*out = static_cast<TreeBuilderHandle>(tree_builder);
API_END();
}

int TreeliteModelBuilderDeleteTree(ModelBuilderHandle handle, int index) {
API_BEGIN();
auto* builder = static_cast<frontend::ModelBuilder*>(handle);
CHECK(builder) << "Detected dangling reference to deleted ModelBuilder object";
TREELITE_CHECK(builder) << "Detected dangling reference to deleted ModelBuilder object";
builder->DeleteTree(index);
API_END();
}

int TreeliteModelBuilderCommitModel(ModelBuilderHandle handle, ModelHandle* out) {
API_BEGIN();
auto* builder = static_cast<frontend::ModelBuilder*>(handle);
CHECK(builder) << "Detected dangling reference to deleted ModelBuilder object";
TREELITE_CHECK(builder) << "Detected dangling reference to deleted ModelBuilder object";
std::unique_ptr<Model> model = builder->CommitModel();
*out = static_cast<ModelHandle>(model.release());
API_END();
Expand Down
2 changes: 1 addition & 1 deletion src/c_api/c_api_runtime.cc
Original file line number Diff line number Diff line change
Expand Up @@ -45,7 +45,7 @@ int TreelitePredictorPredictBatch(
const std::string err_msg
= std::string("Too many columns (features) in the given batch. "
"Number of features must not exceed ") + std::to_string(num_feature);
CHECK_LE(dmat->GetNumCol(), num_feature) << err_msg;
TREELITE_CHECK_LE(dmat->GetNumCol(), num_feature) << err_msg;
*out_result_size = predictor->PredictBatch(dmat, verbose, (pred_margin != 0), out_result);
API_END();
}
Expand Down
2 changes: 1 addition & 1 deletion src/compiler/ast/dump.cc
Original file line number Diff line number Diff line change
Expand Up @@ -13,7 +13,7 @@ void get_dump_from_node(std::ostringstream* oss,
int indent) {
(*oss) << std::string(indent, ' ') << node->GetDump() << "\n";
for (const treelite::compiler::ASTNode* child : node->children) {
CHECK(child);
TREELITE_CHECK(child);
get_dump_from_node(oss, child, indent + 2);
}
}
Expand Down
2 changes: 1 addition & 1 deletion src/compiler/ast/fold_code.cc
Original file line number Diff line number Diff line change
Expand Up @@ -62,7 +62,7 @@ bool fold_code(ASTNode* node, CodeFoldingContext* context,
break;
}
}
CHECK_NE(node_loc, -1); // parent should have a link to current node
TREELITE_CHECK_NE(node_loc, -1); // parent should have a link to current node
parent_node->children[node_loc]
= context->create_new_translation_unit ? tu_node : folder_node;
folder_node->children.push_back(node);
Expand Down
Loading