| 117 | } |
| 118 | |
| 119 | TVector<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 | |
| 164 | TVector<double> CalculatePartialDependence( |
| 165 | const TFullModel& model, |
no test coverage detected