| 155 | } |
| 156 | |
| 157 | void TCatboostModelToCppConverter::WriteModel(bool forCatFeatures, const TFullModel& model, const THashMap<ui32, TString>* catFeaturesHashToString) { |
| 158 | TIndent indent(0); |
| 159 | TSequenceCommaSeparator comma; |
| 160 | |
| 161 | Out << "/* Model data */" << '\n'; |
| 162 | |
| 163 | int binaryFeatureCount = forCatFeatures |
| 164 | ? model.ModelTrees->GetEffectiveBinaryFeaturesBucketsCount() |
| 165 | : GetBinaryFeatureCount(model); |
| 166 | |
| 167 | Out << indent++ << "static const struct CatboostModel {" << '\n'; |
| 168 | Out << indent << "CatboostModel() = default;" << '\n'; |
| 169 | Out << indent << "unsigned int FloatFeatureCount = " << model.GetNumFloatFeatures() << ";" << '\n'; |
| 170 | Out << indent << "unsigned int CatFeatureCount = " << model.GetNumCatFeatures() << ";" << '\n'; |
| 171 | Out << indent << "unsigned int BinaryFeatureCount = " << binaryFeatureCount << ";" << '\n'; |
| 172 | Out << indent << "unsigned int TreeCount = " << model.GetTreeCount() << ";" << '\n'; |
| 173 | |
| 174 | Out << indent++ << "std::vector<std::vector<float>> FloatFeatureBorders = {" << '\n'; |
| 175 | comma.ResetCount(model.ModelTrees->GetFloatFeatures().size()); |
| 176 | for (const auto& floatFeature : model.ModelTrees->GetFloatFeatures()) { |
| 177 | Out << indent << "{" |
| 178 | << OutputArrayInitializer([&floatFeature](size_t i) { return FloatToString(floatFeature.Borders[i], PREC_NDIGITS, 9); }, floatFeature.Borders.size()) |
| 179 | << "}" << comma << '\n'; |
| 180 | } |
| 181 | Out << --indent << "};" << '\n'; |
| 182 | |
| 183 | Out << indent << "unsigned int TreeDepth[" << model.ModelTrees->GetModelTreeData()->GetTreeSizes().size() << "] = {" << OutputArrayInitializer(model.ModelTrees->GetModelTreeData()->GetTreeSizes()) << "};" << '\n'; |
| 184 | Out << indent << "unsigned int TreeSplits[" << model.ModelTrees->GetModelTreeData()->GetTreeSplits().size() << "] = {" << OutputArrayInitializer(model.ModelTrees->GetModelTreeData()->GetTreeSplits()) << "};" << '\n'; |
| 185 | |
| 186 | if (forCatFeatures) { |
| 187 | const auto& bins = model.ModelTrees->GetRepackedBins(); |
| 188 | Out << indent << "unsigned char TreeSplitIdxs[" << bins.size() << "] = {" << OutputArrayInitializer([&bins](size_t i) { return (int)bins[i].SplitIdx; }, bins.size()) << "};" << '\n'; |
| 189 | Out << indent << "unsigned short TreeSplitFeatureIndex[" << bins.size() << "] = {" << OutputArrayInitializer([&bins](size_t i) { return (int)bins[i].FeatureIndex; }, bins.size()) << "};" << '\n'; |
| 190 | Out << indent << "unsigned char TreeSplitXorMask[" << bins.size() << "] = {" << OutputArrayInitializer([&bins](size_t i) { return (int)bins[i].XorMask; }, bins.size()) << "};" << '\n'; |
| 191 | |
| 192 | Out << indent << "unsigned int CatFeaturesIndex[" << model.ModelTrees->GetCatFeatures().size() << "] = {" |
| 193 | << OutputArrayInitializer([&model](size_t i) { return model.ModelTrees->GetCatFeatures()[i].Position.Index; }, model.ModelTrees->GetCatFeatures().size()) << "};" << '\n'; |
| 194 | |
| 195 | Out << indent << "std::vector<unsigned int> OneHotCatFeatureIndex = {" |
| 196 | << OutputArrayInitializer([&model](size_t i) { return model.ModelTrees->GetOneHotFeatures()[i].CatFeatureIndex; }, model.ModelTrees->GetOneHotFeatures().size()) |
| 197 | << "};" << '\n'; |
| 198 | |
| 199 | Out << indent++ << "std::vector<std::vector<int>> OneHotHashValues = {" << '\n'; |
| 200 | comma.ResetCount(model.ModelTrees->GetOneHotFeatures().size()); |
| 201 | for (const auto& oneHotFeature : model.ModelTrees->GetOneHotFeatures()) { |
| 202 | Out << indent << "{" |
| 203 | << OutputArrayInitializer([&oneHotFeature](size_t i) { return oneHotFeature.Values[i]; }, oneHotFeature.Values.size()) |
| 204 | << "}" << comma << '\n'; |
| 205 | } |
| 206 | Out << --indent << "};" << '\n'; |
| 207 | |
| 208 | Out << indent++ << "std::vector<std::vector<float>> CtrFeatureBorders = {" << '\n'; |
| 209 | comma.ResetCount(model.ModelTrees->GetCtrFeatures().size()); |
| 210 | for (const auto& ctrFeature : model.ModelTrees->GetCtrFeatures()) { |
| 211 | Out << indent << "{" |
| 212 | << OutputArrayInitializer([&ctrFeature](size_t i) { return FloatToString(ctrFeature.Borders[i], PREC_NDIGITS, 9) + "f"; }, ctrFeature.Borders.size()) |
| 213 | << "}" << comma << '\n'; |
| 214 | } |
nothing calls this directly
no test coverage detected