| 436 | } |
| 437 | |
| 438 | static 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, |
no test coverage detected