Problem 232846 · medium · Level 02 Linear Data Structures

Keep the Rare Class in Both Sets

train/test split · stratification · seeded shuffling · grouping

A plain shuffled split can be unlucky with a rare class: with 5 fraudulent payments among 200, the test set might receive none of them, and then nobody can measure how well fraud is caught. A stratified split splits every class separately.

Write stratified_split(y, test_pct, seed) that returns a tuple (train_idx, test_idx) of example indices (positions in y), each list sorted in increasing order. Follow this procedure exactly:

  1. Create one generator, rng = random.Random(seed).
  2. Group the indices by label. Within a group the indices are in increasing order.
  3. Go through the labels in sorted order. For each label, shuffle its list of indices with rng.shuffle(...); the first count * test_pct // 100 indices of the shuffled list (where count is the size of the group) go to the test set, the rest to the training set.

Examples

Input:  y = ["spam", "ham", "ham", "spam", "ham", "ham", "ham", "ham", "spam", "ham", "ham", "ham"]
        test_pct = 34, seed = 5
Output: ([0, 1, 6, 7, 8, 9, 10, 11], [2, 3, 4, 5])
Explanation: "ham" comes first. Its 9 indices [1, 2, 4, 5, 6, 7, 9, 10, 11] shuffle to
[4, 5, 2, 1, 11, 10, 9, 7, 6], and 9 * 34 // 100 = 3 of them go to the test set: 4, 5, 2.
Then the "spam" indices [0, 3, 8] shuffle to [3, 0, 8], and 3 * 34 // 100 = 1 goes to the test set: 3.

Input:  y = [0, 0, 1, 1, 1, 0, 1], test_pct = 50, seed = 7
Output: ([0, 1, 4, 6], [2, 3, 5])

Constraints

  • 0 <= len(y) <= 10**5; the labels are all strings or all whole numbers
  • 0 <= test_pct <= 100
  • use the one generator for all shuffles, in the order described; the result must not depend on anything else

Goals

  • Split a labelled dataset so that every class keeps its share in the test set
  • Shuffle indices, not examples, with a seeded random generator
  • Follow a precisely stated random procedure so the result can be reproduced
Starting Python…