| 250 | } |
| 251 | |
| 252 | void TrainEvalSplit( |
| 253 | const NCB::TDataProvider& srcDataProvider, |
| 254 | NCB::TDataProviderPtr* trainDataProvider, |
| 255 | NCB::TDataProviderPtr* evalDataProvider, |
| 256 | const TTrainTestSplitParams& splitParams, |
| 257 | bool saveEvalDataset, |
| 258 | int threadCount, |
| 259 | ui64 cpuUsedRamLimit |
| 260 | ) { |
| 261 | NPar::TLocalExecutor executor; |
| 262 | executor.RunAdditionalThreads(threadCount - 1); |
| 263 | |
| 264 | bool shuffle = splitParams.Shuffle && srcDataProvider.ObjectsData->GetOrder() != NCB::EObjectsOrder::RandomShuffled; |
| 265 | NCB::TObjectsGroupingSubset postShuffleGroupingSubset; |
| 266 | if (shuffle) { |
| 267 | TRestorableFastRng64 rand(splitParams.PartitionRandSeed); |
| 268 | postShuffleGroupingSubset = NCB::Shuffle(srcDataProvider.ObjectsGrouping, 1, &rand); |
| 269 | } else { |
| 270 | postShuffleGroupingSubset = NCB::GetSubset( |
| 271 | srcDataProvider.ObjectsGrouping, |
| 272 | NCB::TArraySubsetIndexing<ui32>(NCB::TFullSubset<ui32>(srcDataProvider.ObjectsGrouping->GetGroupCount())), |
| 273 | NCB::EObjectsOrder::Ordered |
| 274 | ); |
| 275 | } |
| 276 | auto postShuffleGrouping = postShuffleGroupingSubset.GetSubsetGrouping(); |
| 277 | |
| 278 | // for groups |
| 279 | NCB::TArraySubsetIndexing<ui32> postShuffleTrainIndices; |
| 280 | NCB::TArraySubsetIndexing<ui32> postShuffleTestIndices; |
| 281 | |
| 282 | if (splitParams.Stratified) { |
| 283 | auto maybeOneDimensionalTarget = srcDataProvider.RawTargetData.GetOneDimensionalTarget(); |
| 284 | CB_ENSURE(maybeOneDimensionalTarget, "Cannot do stratified split without one-dimensional target data"); |
| 285 | |
| 286 | auto doStratifiedSplit = [&](auto targetArrayRef) { |
| 287 | typedef std::remove_const_t<typename decltype(targetArrayRef)::value_type> TDst; |
| 288 | TVector<TDst> shuffledTarget; |
| 289 | if (shuffle) { |
| 290 | shuffledTarget = NCB::GetSubset<TDst>(targetArrayRef, postShuffleGroupingSubset.GetObjectsIndexing(), &executor); |
| 291 | targetArrayRef = shuffledTarget; |
| 292 | } |
| 293 | NCB::StratifiedTrainTestSplit( |
| 294 | *postShuffleGrouping, |
| 295 | targetArrayRef, |
| 296 | splitParams.TrainPart, |
| 297 | &postShuffleTrainIndices, |
| 298 | &postShuffleTestIndices |
| 299 | ); |
| 300 | }; |
| 301 | |
| 302 | std::visit( |
| 303 | TOverloaded{ |
| 304 | [&](const NCB::ITypedSequencePtr<float>& floatTarget) { doStratifiedSplit(TConstArrayRef<float>(NCB::ToVector(*floatTarget))); }, |
| 305 | [&](const TVector<TString>& stringTarget) { doStratifiedSplit(TConstArrayRef<TString>(stringTarget)); } |
| 306 | }, |
| 307 | **maybeOneDimensionalTarget |
| 308 | ); |
| 309 | } else { |
nothing calls this directly
no test coverage detected