Skip to content

splits

topic_segmentation.data.mimic.splits

Create nested, patient-disjoint MIMIC training and validation sets.

python -m topic_segmentation.data.mimic.splits POOL --train-sizes 50 100 300 1000 4000 --prefix-sizes 8 16 32
python -m topic_segmentation.data.mimic.splits LARGE_POOL --extends POOL --train-sizes 10000 30000 100000

Read note caches and segmentation_labels/ from POOL; write to POOL/splits_valid1000_seed42. Reserve 1,000 notes for validation and keep patients separate across train, validation, and test. --prefix-sizes selects whole-patient prefixes of the smallest training set. --extends retains the earlier pool's validation set and largest training subset.

read_ids(path)

Source code in src/topic_segmentation/data/mimic/splits.py
25
26
def read_ids(path):
    return [line.strip() for line in path.read_text(encoding="utf-8").splitlines() if line.strip()]

note_ids(patients, notes_of)

Source code in src/topic_segmentation/data/mimic/splits.py
29
30
def note_ids(patients, notes_of):
    return {note_id for patient in patients for note_id in notes_of[patient]}

take_patients(candidates, notes_of, size)

Select whole patients in order to reach size notes, skipping groups that would exceed it.

Source code in src/topic_segmentation/data/mimic/splits.py
33
34
35
36
37
38
39
40
41
42
43
44
def take_patients(candidates, notes_of, size):
    """Select whole patients in order to reach size notes, skipping groups that would exceed it."""
    taken, rest, count = [], [], 0
    for patient in candidates:
        if count + len(notes_of[patient]) <= size:
            taken.append(patient)
            count += len(notes_of[patient])
        else:
            rest.append(patient)
    if count != size:
        raise ValueError(f"no patient-disjoint split of exactly {size} notes; reached {count}")
    return taken, rest

whole_patients(ids, subject, notes_of)

Return patient IDs, checking that ids contains each patient's full set of notes.

Source code in src/topic_segmentation/data/mimic/splits.py
47
48
49
50
51
52
def whole_patients(ids, subject, notes_of):
    """Return patient IDs, checking that ids contains each patient's full set of notes."""
    patients = sorted({subject[note_id] for note_id in ids})
    if note_ids(patients, notes_of) != set(ids):
        raise ValueError("note ids do not cover whole patients")
    return patients

patient_prefix(ids, size, subject)

Select the first size notes grouped by patient, requiring a whole-patient prefix.

Source code in src/topic_segmentation/data/mimic/splits.py
55
56
57
58
59
60
61
62
63
64
65
66
67
68
def patient_prefix(ids, size, subject):
    """Select the first size notes grouped by patient, requiring a whole-patient prefix."""
    by_patient = {}
    for note_id in ids:
        by_patient.setdefault(subject[note_id], []).append(note_id)
    chosen = []
    for notes in by_patient.values():  # patients in order of their first note
        if len(chosen) >= size:
            break
        chosen.extend(notes)
    if len(chosen) != size:
        raise ValueError(f"no patient-complete prefix of {size} notes")
    chosen = set(chosen)
    return [note_id for note_id in ids if note_id in chosen]

write_split(directory, name, ids, rows)

Write ids/<name>.txt and <level>/<name>.jsonl; return the split statistics.

Source code in src/topic_segmentation/data/mimic/splits.py
71
72
73
74
75
76
77
78
79
80
81
82
def write_split(directory, name, ids, rows):
    """Write `ids/<name>.txt` and `<level>/<name>.jsonl`; return the split statistics."""
    (directory / "ids").mkdir(parents=True, exist_ok=True)
    (directory / "ids" / f"{name}.txt").write_text("".join(f"{note_id}\n" for note_id in ids), encoding="utf-8")
    statistics = {}
    for level, by_id in rows.items():
        (directory / level).mkdir(parents=True, exist_ok=True)
        with (directory / level / f"{name}.jsonl").open("w", encoding="utf-8") as stream:
            stream.writelines(json.dumps(by_id[note_id], ensure_ascii=False) + "\n" for note_id in ids)
        statistics[level] = {"sentences": sum(int(by_id[note_id]["n_sentences"]) for note_id in ids),
                             "internal_boundaries": sum(int(by_id[note_id]["n_internal_boundaries"]) for note_id in ids)}
    return statistics

main()

Source code in src/topic_segmentation/data/mimic/splits.py
 85
 86
 87
 88
 89
 90
 91
 92
 93
 94
 95
 96
 97
 98
 99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
def main():
    parser = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter)
    parser.add_argument("pool", type=Path, help="Folder with the note cache and segmentation_labels/")
    parser.add_argument("--train-sizes", type=int, nargs="+", required=True)
    parser.add_argument("--prefix-sizes", type=int, nargs="+", default=[])
    parser.add_argument("--extends", type=Path, help="Smaller pool whose validation set and largest training subset are kept")
    args = parser.parse_args()

    order, subject = [], {}
    for pool in ([args.extends] if args.extends else []) + [args.pool]:
        [cache_path] = pool.glob("*_by_id.json")
        cache = json.loads(cache_path.read_text(encoding="utf-8"))
        for note_id in cache["order"]:
            if note_id in subject:
                raise ValueError(f"note {note_id} is in two caches")
            order.append(note_id)
            subject[note_id] = str(cache["notes"][note_id]["subject_id"])
    notes_of = {}
    for note_id in order:
        notes_of.setdefault(subject[note_id], []).append(note_id)
    train_sizes = sorted(set(args.train_sizes))
    if VALID_SIZE + train_sizes[-1] != len(order):
        raise ValueError(f"{VALID_SIZE} validation and {train_sizes[-1]} training notes must make up the {len(order)} notes")
    test = json.loads(TEST.read_text(encoding="utf-8"))["notes"]
    if set(order) & set(test) or set(notes_of) & {str(note["subject_id"]) for note in test.values()}:
        raise AssertionError("the pool overlaps the test sample")

    if args.extends:
        kept = max((args.extends / SPLITS / "ids").glob("train_*.txt"), key=lambda path: int(path.stem.split("_")[1]))
        valid_patients = whole_patients(read_ids(args.extends / SPLITS / "ids/valid_1000.txt"), subject, notes_of)
        excluded = set(valid_patients)
        train_patients = [patient for patient in notes_of if patient not in excluded]
        nested = whole_patients(read_ids(kept), subject, notes_of)
    else:
        shuffled = list(notes_of)
        random.Random(SEED).shuffle(shuffled)
        valid_patients, train_patients = take_patients(shuffled, notes_of, VALID_SIZE)
        nested = []
    excluded = set(nested)
    available = [patient for patient in train_patients if patient not in excluded]
    random.Random(SEED + 1).shuffle(available)
    subsets, size = {}, len(note_ids(nested, notes_of))
    for train_size in train_sizes:
        added, available = take_patients(available, notes_of, train_size - size)
        nested += added
        selected = note_ids(nested, notes_of)
        subsets[train_size], size = [note_id for note_id in order if note_id in selected], train_size
    for prefix_size in args.prefix_sizes:
        subsets[prefix_size] = patient_prefix(subsets[train_sizes[0]], prefix_size, subject)
    subsets = dict(sorted(subsets.items()))
    sizes = list(subsets)
    if any(not set(subsets[small]) < set(subsets[large]) for small, large in zip(sizes, sizes[1:])):
        raise AssertionError("training subsets are not nested")
    selected = note_ids(valid_patients, notes_of)
    valid = [note_id for note_id in order if note_id in selected]

    rows = {}
    for level in LEVELS:
        with (args.pool / f"segmentation_labels/mimic_discharge_{level}.jsonl").open(encoding="utf-8") as stream:
            rows[level] = {row["note_id"]: row for row in map(json.loads, stream)}
        if set(rows[level]) != set(order):
            raise ValueError(f"{level} labels do not match the pool notes")
    out = args.pool / SPLITS
    manifest = {"seed": SEED, "pool_documents": len(order), "pool_subjects": len(notes_of),
                "valid": {"documents": len(valid), "subjects": len(valid_patients),
                          "labels": write_split(out, "valid_1000", valid, rows)},
                "train": {}}
    for train_size, ids in subsets.items():
        manifest["train"][str(train_size)] = {"documents": len(ids), "subjects": len({subject[note_id] for note_id in ids}),
                                              "labels": write_split(out, f"train_{train_size:04d}", ids, rows)}
        print(f"train_{train_size:04d}: {len(ids)} notes of {manifest['train'][str(train_size)]['subjects']} patients")
    (out / "manifest.json").write_text(json.dumps(manifest, indent=2), encoding="utf-8")
    print(f"valid_1000: {len(valid)} notes of {len(valid_patients)} patients -> {out}")