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

Function MergeBucketRanges

catboost/libs/fstr/partial_dependence.cpp:119–162  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

117}
118
119TVector<double> MergeBucketRanges(
120 const TFullModel& model,
121 const TVector<int>& features,
122 const TDataProvider& dataProvider,
123 const TVector<TVector<TFloatFeatureBucketRange>>& leafBucketRanges,
124 const TVector<double> leafWeights
125) {
126 const auto& leafValues = model.ModelTrees->GetModelTreeData()->GetLeafValues();
127 TVector<TFloatFeatureBucketRange> defaultRanges = PrepareFeatureRanges(model, features);
128 CB_ENSURE(defaultRanges.size() == 2, "Number of features must be 2");
129
130 int columnNum = defaultRanges[1].NumOfBuckets;
131 int rowNum = defaultRanges[0].NumOfBuckets;
132 int numOfBucketsTotal = rowNum * columnNum;
133 TVector<double> edges(numOfBucketsTotal);
134
135 for (size_t leafIdx = 0; leafIdx < leafValues.size(); ++leafIdx) {
136 const auto& ranges = leafBucketRanges[leafIdx];
137 double leafValue = leafValues[leafIdx];
138 for (int rowIdx = ranges[0].Start; rowIdx < ranges[0].End; ++rowIdx) {
139 if (ranges[1].Start < ranges[1].End) {
140 edges[rowIdx * columnNum + ranges[1].Start] += leafValue * leafWeights[leafIdx];
141 if (ranges[1].End != columnNum) {
142 edges[rowIdx * columnNum + ranges[1].End] -= leafValue * leafWeights[leafIdx];
143 } else if (rowIdx < rowNum - 1) {
144 edges[(rowIdx + 1) * columnNum] -= leafValue * leafWeights[leafIdx];
145 }
146 }
147 }
148 }
149
150 TVector<double> predictionsByBuckets(numOfBucketsTotal);
151 double acc = 0;
152 for (int idx = 0; idx < numOfBucketsTotal; ++idx) {
153 acc += edges[idx];
154 predictionsByBuckets[idx] = acc;
155 }
156
157 size_t numOfDocuments = dataProvider.GetObjectCount();
158 for (size_t idx = 0; idx < predictionsByBuckets.size(); ++idx) {
159 predictionsByBuckets[idx] /= numOfDocuments;
160 }
161 return predictionsByBuckets;
162}
163
164TVector<double> CalculatePartialDependence(
165 const TFullModel& model,

Callers 1

Calls 4

PrepareFeatureRangesFunction · 0.85
GetLeafValuesMethod · 0.45
sizeMethod · 0.45
GetObjectCountMethod · 0.45

Tested by

no test coverage detected