| 278 | } |
| 279 | |
| 280 | void TModelTrees::TruncateTrees(size_t begin, size_t end) { |
| 281 | //TODO(eermishkina): support non symmetric trees |
| 282 | CB_ENSURE(IsOblivious(), "Truncate support only symmetric trees"); |
| 283 | CB_ENSURE(begin <= end, "begin tree index should be not greater than end tree index."); |
| 284 | CB_ENSURE(end <= GetModelTreeData()->GetTreeSplits().size(), "end tree index should be not greater than tree count."); |
| 285 | auto savedScaleAndBias = GetScaleAndBias(); |
| 286 | TObliviousTreeBuilder builder(FloatFeatures, |
| 287 | CatFeatures, |
| 288 | TextFeatures, |
| 289 | EmbeddingFeatures, |
| 290 | ApproxDimension); |
| 291 | auto applyData = GetApplyData(); |
| 292 | const auto& leafOffsets = applyData->TreeFirstLeafOffsets; |
| 293 | |
| 294 | const auto treeSizes = GetModelTreeData()->GetTreeSizes(); |
| 295 | const auto treeSplits = GetModelTreeData()->GetTreeSplits(); |
| 296 | const auto leafValues = GetModelTreeData()->GetLeafValues(); |
| 297 | const auto leafWeights = GetModelTreeData()->GetLeafWeights(); |
| 298 | const auto treeStartOffsets = GetModelTreeData()->GetTreeStartOffsets(); |
| 299 | for (size_t treeIdx = begin; treeIdx < end; ++treeIdx) { |
| 300 | TVector<TModelSplit> modelSplits; |
| 301 | for (int splitIdx = treeStartOffsets[treeIdx]; |
| 302 | splitIdx < treeStartOffsets[treeIdx] + treeSizes[treeIdx]; |
| 303 | ++splitIdx) |
| 304 | { |
| 305 | modelSplits.push_back(GetBinFeatures()[treeSplits[splitIdx]]); |
| 306 | } |
| 307 | TConstArrayRef<double> leafValuesRef( |
| 308 | leafValues.begin() + leafOffsets[treeIdx], |
| 309 | leafValues.begin() + leafOffsets[treeIdx] + ApproxDimension * (1u << treeSizes[treeIdx]) |
| 310 | ); |
| 311 | builder.AddTree( |
| 312 | modelSplits, |
| 313 | leafValuesRef, |
| 314 | leafWeights.empty() ? TConstArrayRef<double>() : TConstArrayRef<double>( |
| 315 | leafWeights.begin() + leafOffsets[treeIdx] / ApproxDimension, |
| 316 | leafWeights.begin() + leafOffsets[treeIdx] / ApproxDimension + (1ull << treeSizes[treeIdx]) |
| 317 | ) |
| 318 | ); |
| 319 | } |
| 320 | builder.Build(this); |
| 321 | this->SetScaleAndBias(savedScaleAndBias); |
| 322 | } |
| 323 | |
| 324 | flatbuffers::Offset<NCatBoostFbs::TModelTrees> |
| 325 | TModelTrees::FBSerialize(TModelPartsCachingSerializer& serializer) const { |
no test coverage detected