Skip to content

sample

topic_segmentation.data.mimic.sample

Sample MIMIC-IV discharge notes that contain a hospital course section.

python -m topic_segmentation.data.mimic.sample --seed 42 --n 1000 --out-dir OUT [--exclude-sample-json EARLIER.json ...]

Reservoir sampling over discharge.csv.gz; notes and patients of earlier samples are excluded. Writes OUT/discharge_random<n>.json (counts, note IDs, and metadata) and OUT/discharge_random<n>_by_id.json (the notes).

first_line(text)

Source code in src/topic_segmentation/data/mimic/sample.py
27
28
def first_line(text):
    return next((line.strip()[:160] for line in text.splitlines() if line.strip()), "")

excluded_ids(paths)

Collect note and patient IDs from earlier samples.

Source code in src/topic_segmentation/data/mimic/sample.py
31
32
33
34
35
36
37
38
def excluded_ids(paths):
    """Collect note and patient IDs from earlier samples."""
    notes, subjects = set(), set()
    for path in paths:
        for row in json.loads(Path(path).read_text(encoding="utf-8"))["rows"]:
            notes.add(row["note_id"])
            subjects.add(row["subject_id"])
    return notes, subjects

reservoir_sample(size, rng, excluded_notes, excluded_subjects)

Return a uniform sample sorted by note ID, plus filter counts.

Source code in src/topic_segmentation/data/mimic/sample.py
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
def reservoir_sample(size, rng, excluded_notes, excluded_subjects):
    """Return a uniform sample sorted by note ID, plus filter counts."""
    sample, counts = [], Counter()
    with gzip.open(NOTES, "rt", encoding="utf-8", newline="") as handle:
        for row in csv.DictReader(handle):
            counts["source_count"] += 1
            if not COURSE_HEADING_RE.search(row["text"]):
                continue
            counts["matching_count"] += 1
            if row["note_id"] in excluded_notes or row["subject_id"] in excluded_subjects:
                counts["excluded_count"] += 1
                counts["excluded_note_count" if row["note_id"] in excluded_notes else "excluded_subject_count"] += 1
                continue
            counts["eligible_count"] += 1
            slot = len(sample)
            if slot >= size:
                slot = rng.randrange(counts["eligible_count"])
                if slot >= size:
                    continue
            text = row["text"]
            item = {"table": "discharge", **{field: row[field] for field in METADATA_FIELDS}, "text_chars": len(text),
                    "line_count": text.count("\n") + 1 if text else 0, "first_line": first_line(text), "text": text}
            if slot == len(sample):
                sample.append(item)
            else:
                sample[slot] = item
    sample.sort(key=lambda note: note["note_id"])
    names = ("source_count", "matching_count", "excluded_count", "excluded_note_count", "excluded_subject_count",
             "eligible_count")
    return sample, {name: counts[name] for name in names}

main()

Source code in src/topic_segmentation/data/mimic/sample.py
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
def main():
    parser = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter)
    parser.add_argument("--seed", type=int, required=True)
    parser.add_argument("--n", type=int, required=True)
    parser.add_argument("--out-dir", type=Path, required=True)
    parser.add_argument("--exclude-sample-json", type=Path, action="append", default=[],
                        help="Earlier sample whose notes and patients are excluded; may be repeated")
    args = parser.parse_args()

    name = f"discharge_random{args.n}"
    notes, counts = reservoir_sample(args.n, random.Random(args.seed), *excluded_ids(args.exclude_sample_json))
    if len(notes) != args.n:
        raise RuntimeError(f"Requested {args.n} notes but only {len(notes)} were eligible")
    args.out_dir.mkdir(parents=True, exist_ok=True)
    summary = {"name": name, **counts, "sample_count": len(notes), "filter": FILTER,
               "note_ids": [note["note_id"] for note in notes],
               "rows": [{field: note[field] for field in METADATA_FIELDS} for note in notes]}
    (args.out_dir / f"{name}.json").write_text(json.dumps(summary, indent=2), encoding="utf-8")
    cache = {"sample_name": name, "table": "discharge", "count": len(notes),
             "order": [note["note_id"] for note in notes], "notes": {note["note_id"]: note for note in notes}}
    (args.out_dir / f"{name}_by_id.json").write_text(json.dumps(cache, ensure_ascii=False), encoding="utf-8")
    print(f"sampled {len(notes)} of {counts['eligible_count']:,} eligible notes -> {args.out_dir}")