| 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( |
no test coverage detected