MCPcopy Create free account
hub / github.com/catboost/catboost / WriteModel

Method WriteModel

catboost/libs/model/model_export/cpp_exporter.cpp:157–252  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

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 }

Callers

nothing calls this directly

Calls 15

GetBinaryFeatureCountFunction · 0.85
OutputArrayInitializerFunction · 0.85
OutputBorderCountsFunction · 0.85
OutputBordersFunction · 0.85
OutputLeafValuesFunction · 0.85
ResetCountMethod · 0.80
GetFloatFeaturesMethod · 0.80
GetTreeSizesMethod · 0.80
GetTreeSplitsMethod · 0.80
GetRepackedBinsMethod · 0.80
GetCatFeaturesMethod · 0.80

Tested by

no test coverage detected