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

Function CalcSubtreeValuesForTree

catboost/libs/fstr/shap_prepared_trees.cpp:96–175  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

94}
95
96static TVector<TVector<TVector<double>>> CalcSubtreeValuesForTree(
97 const TModelTrees& forest,
98 const TVector<TVector<double>>& subtreeWeights,
99 const TVector<double>& leafWeights,
100 size_t treeIdx
101) {
102 const size_t approxDimension = forest.GetDimensionsCount();
103 TVector<TVector<TVector<double>>> subtreeValues;
104 if (forest.IsOblivious()) {
105 auto firstLeafPtr = forest.GetFirstLeafPtrForTree(treeIdx);
106 const size_t treeDepth = forest.GetModelTreeData()->GetTreeSizes()[treeIdx];
107 subtreeValues.resize(treeDepth + 1);
108 size_t leafNum = size_t(1) << treeDepth;
109 subtreeValues[treeDepth].resize(leafNum);
110 for (size_t leafIdx = 0; leafIdx < leafNum; ++leafIdx) {
111 subtreeValues[treeDepth][leafIdx].resize(approxDimension, 0.0);
112 for (size_t dimension = 0; dimension < approxDimension; ++dimension) {
113 subtreeValues[treeDepth][leafIdx][dimension] = firstLeafPtr[leafIdx * approxDimension + dimension];
114 }
115 }
116 for (int depth = treeDepth - 1; depth >= 0; --depth) {
117 size_t subtreeNum = size_t(1) << depth;
118 subtreeValues[depth].resize(subtreeNum);
119 for (size_t subtreeIdx = 0; subtreeIdx < subtreeNum; ++subtreeIdx) {
120 subtreeValues[depth][subtreeIdx].resize(approxDimension, 0.0);
121 if (!FuzzyEquals(1 + subtreeWeights[depth][subtreeIdx], 1 + 0.0)) {
122 for (size_t dimension = 0; dimension < approxDimension; ++dimension) {
123 subtreeValues[depth][subtreeIdx][dimension] =
124 subtreeValues[depth + 1][subtreeIdx * 2][dimension]
125 * subtreeWeights[depth + 1][subtreeIdx * 2] +
126 subtreeValues[depth + 1][subtreeIdx * 2 + 1][dimension]
127 * subtreeWeights[depth + 1][subtreeIdx * 2 + 1];
128 subtreeValues[depth][subtreeIdx][dimension] /= subtreeWeights[depth][subtreeIdx];
129 }
130 }
131 }
132 }
133 } else {
134 const size_t startOffset = forest.GetModelTreeData()->GetTreeStartOffsets()[treeIdx];
135 auto firstLeafPtr = &forest.GetModelTreeData()->GetLeafValues()[0];
136 TVector<size_t> reversedTree = GetReversedSubtreeForNonObliviousTree(forest, treeIdx);
137 subtreeValues.resize(1);
138 subtreeValues[0].resize(reversedTree.size(), TVector<double>(approxDimension, 0.0));
139 if (reversedTree.size() == 1) {
140 size_t leafIdx = forest.GetModelTreeData()->GetNonSymmetricNodeIdToLeafId()[startOffset];
141 for (size_t dimension = 0; dimension < approxDimension; ++dimension) {
142 subtreeValues[0][0][dimension] = firstLeafPtr[leafIdx + dimension];
143 }
144 } else {
145 for (size_t localIdx = reversedTree.size() - 1; localIdx > 0; --localIdx) {
146 size_t leafIdx = forest.GetModelTreeData()->GetNonSymmetricNodeIdToLeafId()[startOffset + localIdx];
147 size_t leafWeightIdx = leafIdx / approxDimension;
148 if (leafWeightIdx < leafWeights.size()) {
149 if (!FuzzyEquals(1 + leafWeights[leafWeightIdx], 1 + 0.0)) {
150 for (size_t dimension = 0; dimension < approxDimension; ++dimension) {
151 subtreeValues[0][localIdx][dimension] +=
152 firstLeafPtr[leafIdx + dimension]
153 * leafWeights[leafWeightIdx];

Callers 1

CalcTreeStatsFunction · 0.85

Calls 11

FuzzyEqualsFunction · 0.85
GetTreeSizesMethod · 0.80
GetTreeStartOffsetsMethod · 0.80
GetDimensionsCountMethod · 0.45
IsObliviousMethod · 0.45
resizeMethod · 0.45
GetLeafValuesMethod · 0.45
sizeMethod · 0.45

Tested by

no test coverage detected