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

Function PrepareTimeSplitFolds

catboost/libs/train_lib/eval_feature.cpp:438–499  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

436}
437
438static void PrepareTimeSplitFolds(
439 typename TTrainingDataProviders::TDataPtr srcData,
440 const NCatboostOptions::TFeatureEvalOptions& featureEvalOptions,
441 ui64 cpuUsedRamLimit,
442 TVector<TTrainingDataProviders>* foldsData,
443 TVector<TTrainingDataProviders>* testFoldsData,
444 NPar::ILocalExecutor* localExecutor
445) {
446 CB_ENSURE(srcData->ObjectsData->GetGroupIds(), "Timesplit feature evaluation requires dataset with groups");
447 CB_ENSURE(srcData->ObjectsData->GetTimestamp(), "Timesplit feature evaluation requires dataset with timestamps");
448
449 const ui32 foldSize = featureEvalOptions.FoldSize;
450 CB_ENSURE(foldSize > 0, "Fold size must be positive integer");
451 // group subsets, groups maybe trivial
452 const auto& objectsGrouping = *srcData->ObjectsGrouping;
453
454 const auto timesplitQuantileTimestamp = FindQuantileTimestamp(
455 *srcData->ObjectsData->GetGroupIds(),
456 *srcData->ObjectsData->GetTimestamp(),
457 featureEvalOptions.TimeSplitQuantile);
458 TVector<NCB::TArraySubsetIndexing<ui32>> trainTestSubsets; // [0, offset + foldCount) -- train, [offset + foldCount] -- test
459 if (IsObjectwiseEval(featureEvalOptions)) {
460 trainTestSubsets = NCB::QuantileSplitByObjects(
461 objectsGrouping,
462 *srcData->ObjectsData->GetTimestamp(),
463 timesplitQuantileTimestamp,
464 foldSize);
465 } else {
466 trainTestSubsets = NCB::QuantileSplitByGroups(
467 objectsGrouping,
468 *srcData->ObjectsData->GetTimestamp(),
469 timesplitQuantileTimestamp,
470 foldSize);
471 }
472 const ui32 offsetInRange = featureEvalOptions.Offset;
473 const ui32 trainSubsetsCount = trainTestSubsets.size() - 1;
474 const ui32 foldCount = featureEvalOptions.FoldCount;
475 CB_ENSURE_INTERNAL(offsetInRange + foldCount <= trainSubsetsCount, "Dataset permutation logic failed");
476
477 CB_ENSURE(foldsData->empty(), "Need empty vector of folds data");
478 foldsData->resize(foldCount);
479 if (testFoldsData != nullptr) {
480 CB_ENSURE(testFoldsData->empty(), "Need empty vector of test folds data");
481 testFoldsData->resize(foldCount);
482 } else {
483 testFoldsData = foldsData;
484 }
485
486 TVector<NCB::TArraySubsetIndexing<ui32>> trainSubsets(trainTestSubsets.begin(), trainTestSubsets.begin() + trainSubsetsCount);
487 TakeMiddleElements(offsetInRange, foldCount, &trainSubsets);
488
489 TVector<NCB::TArraySubsetIndexing<ui32>> testSubsets(foldCount, trainTestSubsets.back());
490
491 CreateFoldData(
492 srcData,
493 cpuUsedRamLimit,
494 trainSubsets,
495 testSubsets,

Callers 1

EvaluateFeaturesImplFunction · 0.85

Calls 11

FindQuantileTimestampFunction · 0.85
IsObjectwiseEvalFunction · 0.85
TakeMiddleElementsFunction · 0.85
CreateFoldDataFunction · 0.85
GetTimestampMethod · 0.80
GetGroupIdsMethod · 0.45
sizeMethod · 0.45
emptyMethod · 0.45
resizeMethod · 0.45
beginMethod · 0.45
backMethod · 0.45

Tested by

no test coverage detected