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

Function CalcTreesSingleDocImpl

catboost/libs/model/cpu/evaluator_impl.cpp:417–460  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

415
416 template <bool IsSingleClassModel, bool NeedXorMask, bool calcIndexesOnly = false>
417 inline void CalcTreesSingleDocImpl(
418 const TModelTrees& trees,
419 const TModelTrees::TForApplyData& ,
420 const TCPUEvaluatorQuantizedData* quantizedData,
421 size_t,
422 TCalcerIndexType* __restrict indexesVec,
423 size_t treeStart,
424 size_t treeEnd,
425 double* __restrict results) {
426 const ui8* __restrict binFeatures = quantizedData->QuantizedData.data();
427 Y_ASSERT(calcIndexesOnly || (results && AllOf(results, results + trees.GetDimensionsCount(),
428 [](double value) { return value == 0.0; })));
429 const TRepackedBin* __restrict treeSplitsCurPtr =
430 trees.GetRepackedBins().data() + trees.GetModelTreeData()->GetTreeStartOffsets()[treeStart];
431 const double* __restrict treeLeafPtr = trees.GetFirstLeafPtrForTree(treeStart);
432 for (size_t treeId = treeStart; treeId < treeEnd; ++treeId) {
433 const auto curTreeSize = trees.GetModelTreeData()->GetTreeSizes()[treeId];
434 TCalcerIndexType index = 0;
435 for (int depth = 0; depth < curTreeSize; ++depth) {
436 const ui8 borderVal = (ui8)(treeSplitsCurPtr[depth].SplitIdx);
437 const ui32 featureIndex = (treeSplitsCurPtr[depth].FeatureIndex);
438 if constexpr (NeedXorMask) {
439 const ui8 xorMask = (ui8)(treeSplitsCurPtr[depth].XorMask);
440 index |= ((binFeatures[featureIndex] ^ xorMask) >= borderVal) << depth;
441 } else {
442 index |= (binFeatures[featureIndex] >= borderVal) << depth;
443 }
444 }
445 if constexpr (calcIndexesOnly) {
446 *indexesVec++ = index;
447 } else {
448 if constexpr (IsSingleClassModel) { // single class model
449 results[0] += treeLeafPtr[index];
450 } else { // multiclass model
451 const double* __restrict leafValuePtr = treeLeafPtr + index * trees.GetDimensionsCount();
452 for (int classId = 0; classId < (int)trees.GetDimensionsCount(); ++classId) {
453 results[classId] += leafValuePtr[classId];
454 }
455 }
456 treeLeafPtr += (1ull << curTreeSize) * trees.GetDimensionsCount();
457 }
458 treeSplitsCurPtr += curTreeSize;
459 }
460 }
461
462 template <bool NeedXorMask>
463 Y_FORCE_INLINE void CalcIndexesNonSymmetric(

Callers

nothing calls this directly

Calls 7

GetRepackedBinsMethod · 0.80
GetTreeStartOffsetsMethod · 0.80
GetTreeSizesMethod · 0.80
AllOfFunction · 0.50
dataMethod · 0.45
GetDimensionsCountMethod · 0.45

Tested by

no test coverage detected