| 163 | } |
| 164 | |
| 165 | TVector<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 |
no test coverage detected