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

Function StratifiedTrainTestSplit

catboost/libs/data/objects_grouping.h:293–325  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

291
292 template <class TClassId>
293 void StratifiedTrainTestSplit(
294 const TObjectsGrouping& objectsGrouping,
295 TConstArrayRef<TClassId> objectClasses,
296 double trainPart,
297 TArraySubsetIndexing<ui32>* trainIndices,
298 TArraySubsetIndexing<ui32>* testIndices
299 ) {
300 TVector<TVector<ui32>> splittedByClass = SplitByClass(objectsGrouping, objectClasses);
301 ui32 minLen = GetClassSplitMinLen(objectsGrouping.GetObjectCount(), splittedByClass);
302 if (minLen < 2) {
303 CATBOOST_WARNING_LOG << " Warning: The least populated class in y has only "
304 << minLen << " members, which is too few.";
305 }
306 TVector<ui32> resultTrainIndices;
307 TVector<ui32> resultTestIndices;
308 for (const auto& part : splittedByClass) {
309 for (ui32 idx = 0; idx < part.size() * trainPart; ++idx) {
310 resultTrainIndices.push_back(part[idx]);
311 }
312 for (ui32 idx = part.size() * trainPart; idx < part.size(); ++idx) {
313 resultTestIndices.push_back(part[idx]);
314 }
315 }
316
317 CB_ENSURE(!resultTrainIndices.empty(), "Not enough objects for splitting into train and test subsets");
318 CB_ENSURE(!resultTestIndices.empty(), "Not enough objects for splitting into train and test subsets");
319
320 Sort(resultTrainIndices.begin(), resultTrainIndices.end());
321 *trainIndices = TArraySubsetIndexing<ui32>(std::move(resultTrainIndices));
322
323 Sort(resultTestIndices.begin(), resultTestIndices.end());
324 *testIndices = TArraySubsetIndexing<ui32>(std::move(resultTestIndices));
325 }
326
327 template <class TClassId>
328 TVector<TArraySubsetIndexing<ui32>> StratifiedSplitToFolds(

Callers 2

TrainEvalSplitFunction · 0.85
PrepareTrainTestSplitFunction · 0.85

Calls 10

SplitByClassFunction · 0.85
GetClassSplitMinLenFunction · 0.85
SortFunction · 0.50
moveFunction · 0.50
GetObjectCountMethod · 0.45
sizeMethod · 0.45
push_backMethod · 0.45
emptyMethod · 0.45
beginMethod · 0.45
endMethod · 0.45

Tested by

no test coverage detected