| """PTB-XL preprocessing, following DL4mHealth/Medformer's PTB-XL_preprocessing.ipynb. |
| |
| * records500 (12-lead, 500 Hz, 10 s) -> linear resample to 250 Hz -> 2500 samples |
| * per-record StandardScaler over time, then reshape into 10 non-overlapping |
| 250-timestamp windows (the paper's Table 5: T=250, patch 10, 5 classes) |
| * label = the SCP code with the highest likelihood, mapped to the 5 diagnostic |
| superclasses NORM / MI / STTC / CD / HYP |
| * subject-wise (patient_id) split; Medformer's PTB-XL loader uses a 60/20/20 |
| patient split, which we reproduce with a fixed seed. |
| """ |
|
|
| import argparse |
| import ast |
| import os |
|
|
| import numpy as np |
| import pandas as pd |
| import wfdb |
| from scipy import interpolate |
| from sklearn.preprocessing import StandardScaler |
|
|
| SUPER = {"NORM": 0, "MI": 1, "STTC": 2, "CD": 3, "HYP": 4} |
|
|
|
|
| def build_scp_map(scp_csv): |
| """diagnostic_class column of scp_statements.csv maps each SCP code to a superclass.""" |
| df = pd.read_csv(scp_csv, index_col=0) |
| df = df[df.diagnostic == 1] |
| return {code: row.diagnostic_class for code, row in df.iterrows() |
| if isinstance(row.diagnostic_class, str) and row.diagnostic_class in SUPER} |
|
|
|
|
| def resample(arr, freq=500, target=250): |
| t = np.linspace(1, len(arr), len(arr)) |
| f = interpolate.interp1d(t, arr, kind="linear") |
| return f(np.linspace(1, len(arr), int(len(arr) / freq * target))) |
|
|
|
|
| def main(): |
| ap = argparse.ArgumentParser() |
| ap.add_argument("--root", default="data/PTBXL") |
| ap.add_argument("--out", default="data/PTBXL/processed") |
| ap.add_argument("--subject-frac", type=float, default=1.0, |
| help="fraction of patients to keep (compute budget)") |
| ap.add_argument("--seed", type=int, default=0) |
| a = ap.parse_args() |
|
|
| info = pd.read_csv(os.path.join(a.root, "ptbxl_database.csv")) |
| code2super = build_scp_map(os.path.join(a.root, "scp_statements.csv")) |
|
|
| rows = [] |
| for r in info.itertuples(): |
| codes = ast.literal_eval(r.scp_codes) |
| diag = {c: v for c, v in codes.items() if c in code2super} |
| if not diag: |
| continue |
| best = max(diag, key=diag.get) |
| rows.append((r.ecg_id, r.patient_id, SUPER[code2super[best]], r.filename_hr)) |
| df = pd.DataFrame(rows, columns=["ecg_id", "patient_id", "label", "path"]) |
| print(f"records with a diagnostic superclass: {len(df)} " |
| f"({df.patient_id.nunique()} patients)") |
| print("class counts:", df.label.value_counts().sort_index().to_dict()) |
|
|
| pats = np.array(sorted(df.patient_id.unique())) |
| rng = np.random.default_rng(a.seed) |
| rng.shuffle(pats) |
| if a.subject_frac < 1.0: |
| pats = pats[: max(2, int(len(pats) * a.subject_frac))] |
| df = df[df.patient_id.isin(pats)] |
| print(f"subsampled to {len(pats)} patients / {len(df)} records " |
| f"(subject_frac={a.subject_frac})") |
| n = len(pats) |
| splits = {"train": set(pats[: int(0.6 * n)]), |
| "val": set(pats[int(0.6 * n):int(0.8 * n)]), |
| "test": set(pats[int(0.8 * n):])} |
|
|
| os.makedirs(a.out, exist_ok=True) |
| for name, ids in splits.items(): |
| sub = df[df.patient_id.isin(ids)] |
| X, y = [], [] |
| for i, r in enumerate(sub.itertuples()): |
| sig, _ = wfdb.rdsamp(os.path.join(a.root, r.path)) |
| res = np.stack([resample(sig[:, c]) for c in range(sig.shape[1])], 1) |
| res = StandardScaler().fit_transform(res) |
| w = res[: (res.shape[0] // 250) * 250].reshape(-1, 250, res.shape[1]) |
| X.append(w.astype(np.float32)) |
| y.append(np.full(w.shape[0], r.label, dtype=np.int64)) |
| if i % 2000 == 0: |
| print(f" {name} {i}/{len(sub)}", flush=True) |
| X = np.concatenate(X) |
| y = np.concatenate(y) |
| np.save(os.path.join(a.out, f"X_{name}.npy"), X) |
| np.save(os.path.join(a.out, f"y_{name}.npy"), y) |
| print(name, X.shape, np.bincount(y, minlength=5), flush=True) |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|