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

Function TrainEvalSplit

catboost/python-package/catboost/helpers.cpp:252–329  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

250}
251
252void 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 {

Callers

nothing calls this directly

Calls 12

GetSubsetFunction · 0.85
StratifiedTrainTestSplitFunction · 0.85
ComposeFunction · 0.85
RunAdditionalThreadsMethod · 0.80
GetOrderMethod · 0.80
GetGroupCountMethod · 0.80
GetSubsetGroupingMethod · 0.80
ShuffleFunction · 0.50
visitFunction · 0.50
ToVectorFunction · 0.50
GetSubsetMethod · 0.45

Tested by

no test coverage detected