#pragma once
#include "c_api.h"
#include <string>
#include <array>
#include <vector>
#include <functional>
#include <memory>
class ModelCalcerWrapper {
public:
ModelCalcerWrapper()
: CalcerHolder(CalcerHolderType(ModelCalcerCreate(), ModelCalcerDelete))
{}
explicit ModelCalcerWrapper(const std::string& filename) {
CalcerHolder = CalcerHolderType(ModelCalcerCreate(), ModelCalcerDelete);
if (!LoadFullModelFromFile(CalcerHolder.get(), filename.c_str()) ) {
throw std::runtime_error(GetErrorString());
}
InitProps();
}
explicit ModelCalcerWrapper(const void* binaryBuffer, size_t binaryBufferSize) {
CalcerHolder = CalcerHolderType(ModelCalcerCreate(), ModelCalcerDelete);
if (!LoadFullModelFromBuffer(CalcerHolder.get(), binaryBuffer, binaryBufferSize) ) {
throw std::runtime_error(GetErrorString());
}
InitProps();
}
void EnableGPUEvaluation(int deviceId = 0) {
if (!::EnableGPUEvaluation(CalcerHolder.get(), deviceId)) {
throw std::runtime_error(GetErrorString());
}
}
double CalcFlat(const std::vector<float>& features) const {
double result;
const float* ptr = features.data();
if (!CalcModelPredictionFlat(CalcerHolder.get(), 1, &ptr, features.size(), &result, 1)) {
throw std::runtime_error(GetErrorString());
}
return result;
}
std::vector<double> CalcFlatMulti(const std::vector<float>& features) const {
std::vector<double> result(DimensionsCount, 0.0);
const float* ptr = features.data();
if (!CalcModelPredictionFlat(CalcerHolder.get(), 1, &ptr, features.size(), result.data(), DimensionsCount)) {
throw std::runtime_error(GetErrorString());
}
return result;
}
double Calc(const std::vector<float>& floatFeatures, const std::vector<std::string>& catFeatures) const {
double result;
const float* floatPtr = floatFeatures.data();
std::vector<const char*> catFeaturesPtrs;
FromStringToCharVector(catFeatures, &catFeaturesPtrs);
const char** catFeaturesPtr = catFeaturesPtrs.data();
if (!CalcModelPrediction(CalcerHolder.get(), 1, &floatPtr, floatFeatures.size(), &catFeaturesPtr, catFeatures.size(), &result, 1)) {
throw std::runtime_error(GetErrorString());
}
return result;
}
std::vector<double> CalcMulti(const std::vector<float>& floatFeatures, const std::vector<std::string>& catFeatures) const {
std::vector<double> result(DimensionsCount);
const float* floatPtr = floatFeatures.data();
std::vector<const char*> catFeaturesPtrs;
FromStringToCharVector(catFeatures, &catFeaturesPtrs);
const char** catFeaturesPtr = catFeaturesPtrs.data();
if (!CalcModelPrediction(CalcerHolder.get(), 1, &floatPtr, floatFeatures.size(), &catFeaturesPtr, catFeatures.size(), result.data(), DimensionsCount)) {
throw std::runtime_error(GetErrorString());
}
return result;
}
double Calc(
const std::vector<float>& floatFeatures,
const std::vector<std::string>& catFeatures,
const std::vector<std::string>& textFeatures
) const {
double result;
const float* floatPtr = floatFeatures.data();
std::vector<const char*> catFeaturesPtrs;
FromStringToCharVector(catFeatures, &catFeaturesPtrs);
const char** catFeaturesPtr = catFeaturesPtrs.data();
std::vector<const char*> textFeaturesPtrs;
FromStringToCharVector(textFeatures, &textFeaturesPtrs);
const char** textFeaturesPtr = textFeaturesPtrs.data();
if (!CalcModelPredictionText(
CalcerHolder.get(), 1,
&floatPtr, floatFeatures.size(),
&catFeaturesPtr, catFeatures.size(),
&textFeaturesPtr, textFeatures.size(),
&result, 1
)) {
throw std::runtime_error(GetErrorString());
}
return result;
}
std::vector<double> CalcMulti(
const std::vector<float>& floatFeatures,
const std::vector<std::string>& catFeatures,
const std::vector<std::string>& textFeatures
) const {
std::vector<double> result(DimensionsCount);
const float* floatPtr = floatFeatures.data();
std::vector<const char*> catFeaturesPtrs;
FromStringToCharVector(catFeatures, &catFeaturesPtrs);
const char** catFeaturesPtr = catFeaturesPtrs.data();
std::vector<const char*> textFeaturesPtrs;
FromStringToCharVector(textFeatures, &textFeaturesPtrs);
const char** textFeaturesPtr = textFeaturesPtrs.data();
if (!CalcModelPredictionText(
CalcerHolder.get(), 1,
&floatPtr, floatFeatures.size(),
&catFeaturesPtr, catFeatures.size(),
&textFeaturesPtr, textFeatures.size(),
result.data(), DimensionsCount
)) {
throw std::runtime_error(GetErrorString());
}
return result;
}
std::vector<double> CalcFlat(const std::vector<std::vector<float>>& features) const {
std::vector<double> result(features.size() * DimensionsCount);
std::vector<const float*> ptrsVector;
size_t flatVecSize = 0;
for (const auto& flatVec : features) {
flatVecSize = flatVec.size();
ptrsVector.push_back(flatVec.data());
}
if (!CalcModelPredictionFlat(CalcerHolder.get(), features.size(), ptrsVector.data(), flatVecSize, result.data(), result.size())) {
throw std::runtime_error(GetErrorString());
}
return result;
}
std::vector<double> Calc(const std::vector<std::vector<float>>& floatFeatures,
const std::vector<std::vector<std::string>>& catFeatures) const {
std::vector<double> result(floatFeatures.size() * DimensionsCount);
std::vector<const float*> floatPtrsVector;
size_t floatFeatureCount = 0;
for (const auto& floatFeatureVec : floatFeatures) {
if (floatFeatureCount == 0) {
floatFeatureCount = floatFeatureVec.size();
}
floatPtrsVector.push_back(floatFeatureVec.data());
}
size_t catFeatureCount = 0;
std::vector<const char*> catFeaturesPtrsVector;
std::vector<const char**> charPtrPtrsVector;
FromStringToCharVectors(catFeatures, &catFeatureCount, &catFeaturesPtrsVector, &charPtrPtrsVector);
if (!CalcModelPrediction(
CalcerHolder.get(),
result.size(),
floatPtrsVector.data(), floatFeatureCount,
charPtrPtrsVector.data(), catFeatureCount,
result.data(), result.size())
) {
throw std::runtime_error(GetErrorString());
}
return result;
}
std::vector<double> Calc(
const std::vector<std::vector<float>>& floatFeatures,
const std::vector<std::vector<std::string>>& catFeatures,
const std::vector<std::vector<std::string>>& textFeatures
) const {
std::vector<double> result(floatFeatures.size() * DimensionsCount);
std::vector<const float*> floatPtrsVector;
size_t floatFeatureCount = 0;
for (const auto& floatFeatureVec : floatFeatures) {
if (floatFeatureCount == 0) {
floatFeatureCount = floatFeatureVec.size();
}
floatPtrsVector.push_back(floatFeatureVec.data());
}
size_t catFeatureCount = 0;
std::vector<const char*> catFeaturesPtrsVector;
std::vector<const char**> charPtrPtrsVector;
FromStringToCharVectors(catFeatures, &catFeatureCount, &catFeaturesPtrsVector, &charPtrPtrsVector);
size_t textFeatureCount = 0;
std::vector<const char*> textFeaturesPtrsVector;
std::vector<const char**> charTextPtrPtrsVector;
FromStringToCharVectors(textFeatures, &textFeatureCount, &textFeaturesPtrsVector, &charTextPtrPtrsVector);
if (!CalcModelPredictionText(
CalcerHolder.get(),
result.size(),
floatPtrsVector.data(), floatFeatureCount,
charPtrPtrsVector.data(), catFeatureCount,
charTextPtrPtrsVector.data(), textFeatureCount,
result.data(), result.size()
)) {
throw std::runtime_error(GetErrorString());
}
return result;
}
std::vector<double> CalcHashed(const std::vector<std::vector<float>>& floatFeatures,
const std::vector<std::vector<int>>& catFeatureHashes) const {
std::vector<double> result(floatFeatures.size() * DimensionsCount);
std::vector<const float*> floatPtrsVector;
std::vector<const int*> hashPtrsVector;
size_t floatFeatureCount = 0;
for (const auto& floatFeatureVec : floatFeatures) {
floatFeatureCount = floatFeatureVec.size();
floatPtrsVector.push_back(floatFeatureVec.data());
}
size_t catFeatureCount = 0;
for (const auto& hashVec : catFeatureHashes) {
catFeatureCount = hashVec.size();
hashPtrsVector.push_back(hashVec.data());
}
if (!CalcModelPredictionWithHashedCatFeatures(
CalcerHolder.get(),
result.size(),
floatPtrsVector.data(), floatFeatureCount,
hashPtrsVector.data(), catFeatureCount,
result.data(), result.size())
) {
throw std::runtime_error(GetErrorString());
}
return result;
}
bool InitFromFile(const std::string& filename) {
if (!LoadFullModelFromFile(CalcerHolder.get(), filename.c_str())) {
return false;
}
InitProps();
return true;
}
bool InitFromMemory(const void* pointer, size_t size) {
if (!LoadFullModelFromBuffer(CalcerHolder.get(), pointer, size)) {
return false;
}
InitProps();
return true;
}
bool init_from_file(const std::string& filename) { return InitFromFile(filename);
}
size_t GetTreeCount() const {
return ::GetTreeCount(CalcerHolder.get());
}
size_t GetFloatFeaturesCount() const {
return ::GetFloatFeaturesCount(CalcerHolder.get());
}
size_t GetCatFeaturesCount() const {
return ::GetCatFeaturesCount(CalcerHolder.get());
}
bool CheckMetadataHasKey(const std::string& key) const {
return ::CheckModelMetadataHasKey(CalcerHolder.get(), key.c_str(), key.size());
}
std::string GetMetadataKeyValue(const std::string& key) const {
if (!CheckMetadataHasKey(key)) {
return "";
}
size_t value_size = GetModelInfoValueSize(CalcerHolder.get(), key.c_str(), key.size());
const char* value_ptr = GetModelInfoValue(CalcerHolder.get(), key.c_str(), key.size());
return std::string(value_ptr, value_size);
}
private:
void InitProps() {
DimensionsCount = GetDimensionsCount(CalcerHolder.get());
}
void FromStringToCharVector(const std::vector<std::string>& stringFeatures, std::vector<const char*>* charFeatures) const {
charFeatures->clear();
charFeatures->reserve(stringFeatures.size());
for (const auto& str : stringFeatures) {
charFeatures->push_back(str.data());
}
}
void FromStringToCharVectors(
const std::vector<std::vector<std::string>>& stringFeatures,
size_t* featureCount,
std::vector<const char*>* featuresPtrsVector,
std::vector<const char**>* charPtrPtrsVector
) const {
size_t currentTextOffset = 0;
for (const auto& stringVec : stringFeatures) {
if (*featureCount == 0) {
*featureCount = stringVec.size();
}
if (*featureCount != stringVec.size()) {
throw std::runtime_error("All text feature vectors should be of the same length");
}
}
if (*featureCount != 0) {
featuresPtrsVector->reserve(stringFeatures.size() * (*featureCount));
for (const auto& stringVec : stringFeatures) {
for (const auto& string : stringVec) {
featuresPtrsVector->push_back(string.data());
}
charPtrPtrsVector->push_back(featuresPtrsVector->data() + currentTextOffset);
currentTextOffset += *featureCount;
}
}
}
using CalcerHolderType = std::unique_ptr<ModelCalcerHandle, std::function<void(ModelCalcerHandle*)>>;
CalcerHolderType CalcerHolder;
size_t DimensionsCount = 0;
};