| 9 | |
| 10 | |
| 11 | def train_val_test_split(dset_len, train_size, val_size, test_size, seed): |
| 12 | |
| 13 | assert (train_size is None) + (val_size is None) + (test_size is None) <= 1, "Only one of train_size, val_size, test_size is allowed to be None." |
| 14 | |
| 15 | is_float = (isinstance(train_size, float), isinstance(val_size, float), isinstance(test_size, float)) |
| 16 | |
| 17 | train_size = round(dset_len * train_size) if is_float[0] else train_size |
| 18 | val_size = round(dset_len * val_size) if is_float[1] else val_size |
| 19 | test_size = round(dset_len * test_size) if is_float[2] else test_size |
| 20 | |
| 21 | if train_size is None: |
| 22 | train_size = dset_len - val_size - test_size |
| 23 | elif val_size is None: |
| 24 | val_size = dset_len - train_size - test_size |
| 25 | elif test_size is None: |
| 26 | test_size = dset_len - train_size - val_size |
| 27 | |
| 28 | if train_size + val_size + test_size > dset_len: |
| 29 | if is_float[2]: |
| 30 | test_size -= 1 |
| 31 | elif is_float[1]: |
| 32 | val_size -= 1 |
| 33 | elif is_float[0]: |
| 34 | train_size -= 1 |
| 35 | |
| 36 | assert train_size >= 0 and val_size >= 0 and test_size >= 0, ( |
| 37 | f"One of training ({train_size}), validation ({val_size}) or " |
| 38 | f"testing ({test_size}) splits ended up with a negative size." |
| 39 | ) |
| 40 | |
| 41 | total = train_size + val_size + test_size |
| 42 | assert dset_len >= total, f"The dataset ({dset_len}) is smaller than the combined split sizes ({total})." |
| 43 | |
| 44 | if total < dset_len: |
| 45 | rank_zero_warn(f"{dset_len - total} samples were excluded from the dataset") |
| 46 | |
| 47 | idxs = np.arange(dset_len, dtype=np.int64) |
| 48 | idxs = np.random.default_rng(seed).permutation(idxs) |
| 49 | |
| 50 | idx_train = idxs[:train_size] |
| 51 | idx_val = idxs[train_size: train_size + val_size] |
| 52 | idx_test = idxs[train_size + val_size: total] |
| 53 | |
| 54 | return np.array(idx_train), np.array(idx_val), np.array(idx_test) |
| 55 | |
| 56 | |
| 57 | def make_splits(dataset_len, train_size, val_size, test_size, seed, filename=None, splits=None): |