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

Function PrepareFolds

catboost/libs/train_lib/eval_feature.cpp:501–558  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

499}
500
501static void PrepareFolds(
502 typename TTrainingDataProviders::TDataPtr srcData,
503 const TCvDataPartitionParams& cvParams,
504 const NCatboostOptions::TFeatureEvalOptions& featureEvalOptions,
505 ui64 cpuUsedRamLimit,
506 TVector<TTrainingDataProviders>* foldsData,
507 TVector<TTrainingDataProviders>* testFoldsData,
508 NPar::ILocalExecutor* localExecutor
509) {
510 const int foldCount = cvParams.Initialized() ? cvParams.FoldCount : featureEvalOptions.FoldCount.Get();
511 CB_ENSURE(foldCount > 0, "Fold count must be positive integer");
512 const auto& objectsGrouping = *srcData->ObjectsGrouping;
513 TVector<NCB::TArraySubsetIndexing<ui32>> testSubsets;
514 if (cvParams.Initialized()) {
515 // group subsets, groups maybe trivial
516 testSubsets = NCB::Split(objectsGrouping, foldCount);
517 // always inverted
518 CB_ENSURE(cvParams.Type == ECrossValidation::Inverted, "Feature evaluation requires inverted cross-validation");
519 } else {
520 const ui32 foldSize = featureEvalOptions.FoldSize;
521 CB_ENSURE(foldSize > 0, "Fold size must be positive integer");
522 // group subsets, groups maybe trivial
523 const auto isObjectwise = IsObjectwiseEval(featureEvalOptions);
524 testSubsets = isObjectwise
525 ? NCB::SplitByObjects(objectsGrouping, foldSize)
526 : NCB::SplitByGroups(objectsGrouping, foldSize);
527 const ui32 offsetInRange = featureEvalOptions.Offset;
528 CB_ENSURE_INTERNAL(offsetInRange + foldCount <= testSubsets.size(), "Dataset permutation logic failed");
529 }
530 const ui32 offsetInRange = !cvParams.Initialized() ? featureEvalOptions.Offset : 0;
531
532 TVector<NCB::TArraySubsetIndexing<ui32>> trainSubsets
533 = CalcTrainSubsetsRange(testSubsets, objectsGrouping.GetGroupCount(), TIndexRange<ui32>(offsetInRange, offsetInRange + foldCount));
534
535 if (!cvParams.Initialized()) {
536 TakeMiddleElements(offsetInRange, foldCount, &trainSubsets);
537 TakeMiddleElements(offsetInRange, foldCount, &testSubsets);
538 }
539 testSubsets.swap(trainSubsets);
540
541 CB_ENSURE(foldsData->empty(), "Need empty vector of folds data");
542 foldsData->resize(foldCount);
543 if (testFoldsData != nullptr) {
544 CB_ENSURE(testFoldsData->empty(), "Need empty vector of test folds data");
545 testFoldsData->resize(foldCount);
546 } else {
547 testFoldsData = foldsData;
548 }
549
550 CreateFoldData(
551 srcData,
552 cpuUsedRamLimit,
553 trainSubsets,
554 testSubsets,
555 foldsData,
556 testFoldsData,
557 localExecutor);
558}

Callers 1

EvaluateFeaturesImplFunction · 0.85

Calls 12

IsObjectwiseEvalFunction · 0.85
CalcTrainSubsetsRangeFunction · 0.85
TakeMiddleElementsFunction · 0.85
CreateFoldDataFunction · 0.85
GetGroupCountMethod · 0.80
SplitFunction · 0.50
InitializedMethod · 0.45
GetMethod · 0.45
sizeMethod · 0.45
swapMethod · 0.45
emptyMethod · 0.45
resizeMethod · 0.45

Tested by

no test coverage detected