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

Function GetPredictionDiff

catboost/libs/fstr/compare_documents.cpp:165–249  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

163}
164
165TVector<double> GetPredictionDiff(
166 const TFullModel& model,
167 const TObjectsDataProviderPtr objectsDataProvider,
168 NPar::ILocalExecutor* localExecutor
169) {
170 CB_ENSURE(model.ModelTrees->GetDimensionsCount() == 1, "Is not supported for multiclass");
171 CB_ENSURE(objectsDataProvider->GetObjectCount() == 2, "PredictionDiff requires 2 documents for compare");
172 CB_ENSURE(model.GetNumCatFeatures() == 0, "Models with categorical features are not supported");
173
174 TVector<ui32> leafIdxes = CalcLeafIndexesMulti(model, objectsDataProvider, 0, 0);
175
176 const auto* const rawObjectsData = dynamic_cast<const TRawObjectsDataProvider*>(objectsDataProvider.Get());
177
178 const auto& binSplits = model.ModelTrees->GetBinFeatures();
179
180 TVector<TVector<double>> floatFeatureValues(objectsDataProvider->GetObjectCount());
181 for (size_t idx = 0; idx < objectsDataProvider->GetObjectCount(); ++idx) {
182 floatFeatureValues[idx].resize(model.GetNumFloatFeatures());
183 }
184
185 TVector<TVector<ui32>> docBorders(objectsDataProvider->GetObjectCount());
186 TVector<ui32> borderIdxForSplit(binSplits.size(), std::numeric_limits<ui32>::infinity());
187 ui32 splitIdx = 0;
188 for (const auto& feature : model.ModelTrees->GetFloatFeatures()) {
189 TMaybeData<const TFloatValuesHolder*> featureData
190 = rawObjectsData->GetFloatFeature(feature.Position.FlatIndex);
191
192 if (const auto* arrayColumn = dynamic_cast<const TFloatArrayValuesHolder*>(*featureData)) {
193 arrayColumn->GetData()->ParallelForEach(
194 [&] (ui32 docId, float value) {
195 docBorders[docId].push_back(0);
196 for (const auto& border: feature.Borders) {
197 if (value > border) {
198 docBorders[docId].back()++;
199 floatFeatureValues[docId][feature.Position.FlatIndex] = value;
200 }
201 }
202 },
203 localExecutor
204 );
205 } else {
206 CB_ENSURE_INTERNAL(false, "GetPredictionDiff: Unsupported column type");
207 }
208 if (splitIdx == binSplits.size() ||
209 binSplits[splitIdx].Type != ESplitType::FloatFeature ||
210 binSplits[splitIdx].FloatFeature.FloatFeature > feature.Position.Index
211 ) {
212 continue;
213 }
214 Y_ASSERT(binSplits[splitIdx].FloatFeature.FloatFeature >= feature.Position.Index);
215 for (ui32 idx = 0; idx < feature.Borders.size() && binSplits[splitIdx].FloatFeature.FloatFeature == feature.Position.Index; ++idx) {
216 if (abs(binSplits[splitIdx].FloatFeature.Split - feature.Borders[idx]) < 1e-15) {
217 borderIdxForSplit[splitIdx] = idx;
218 ++splitIdx;
219 }
220 }
221 }
222

Callers 3

GetPredictionDiffWrapperFunction · 0.85
GetFeatureImportancesFunction · 0.85

Calls 15

CalcLeafIndexesMultiFunction · 0.85
ApplyModelMultiFunction · 0.85
GetPredictionDiffSingleFunction · 0.85
MakeArrayRefFunction · 0.85
GetFloatFeaturesMethod · 0.80
absFunction · 0.50
GetDimensionsCountMethod · 0.45
GetObjectCountMethod · 0.45
GetNumCatFeaturesMethod · 0.45
GetMethod · 0.45
GetBinFeaturesMethod · 0.45
resizeMethod · 0.45

Tested by

no test coverage detected