| 499 | } |
| 500 | |
| 501 | static 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 | } |
no test coverage detected