tsfp-repro-code / prep_adftd.py
riteshhf's picture
Upload folder using huggingface_hub
2188a91 verified
Raw
History Blame Contribute Delete
3.74 kB
"""ADFTD (OpenNeuro ds004504) preprocessing, following DL4mHealth/Medformer.
Steps (verbatim from Medformer's ADFTD_preprocessing.ipynb):
* read the ASR/ICA-cleaned derivatives (.set, 19 channels, 500 Hz)
* linear-interpolation resample 500 -> 256 Hz
* segment into non-overlapping 256-timestamp windows, centred on the recording
* label map {C: 0 healthy, F: 1 FTD, A: 2 AD}
Split (from Medformer's ADFTDLoader): subject-wise 60/20/20 within each class.
Normalisation: per-sample, per-channel standardisation (uea.normalize_batch_ts).
"""
import argparse
import os
import mne
import numpy as np
import pandas as pd
from scipy import interpolate
mne.set_log_level("ERROR")
LABEL_MAP = {"A": 2, "F": 1, "C": 0}
def resample(arr, freq=500, target=256):
t = np.linspace(1, len(arr), len(arr))
f = interpolate.interp1d(t, arr, kind="linear")
t_new = np.linspace(1, len(arr), int(len(arr) / freq * target))
return f(t_new)
def segment(mat, window=256):
"""mat: (T, C). Centre-anchored non-overlapping windows."""
res = []
start = mat.shape[0] // 2
left = start - (start // window) * window
right = start + ((mat.shape[0] - start) // window) * window
for i in range(left, right, window):
res.append(mat[i:i + window])
return res
def main():
ap = argparse.ArgumentParser()
ap.add_argument("--root", default="data/ADFTD")
ap.add_argument("--out", default="data/ADFTD/processed")
ap.add_argument("--window", type=int, default=256)
args = ap.parse_args()
parts = pd.read_csv(os.path.join(args.root, "participants.tsv"), sep="\t")
labels = {}
for row in parts.values:
pid = int(str(row[0])[-3:])
labels[pid] = LABEL_MAP[row[3]]
deriv = os.path.join(args.root, "derivatives")
feats, ys, subs = [], [], []
for pid in sorted(labels):
path = os.path.join(deriv, f"sub-{pid:03d}", "eeg",
f"sub-{pid:03d}_task-eyesclosed_eeg.set")
raw = mne.io.read_raw_eeglab(path, preload=True)
sig = raw.get_data() # (C, T) at 500 Hz
res = np.stack([resample(sig[c]) for c in range(sig.shape[0])], axis=1) # (T', C)
for win in segment(res, args.window):
feats.append(win.astype(np.float32))
ys.append(labels[pid])
subs.append(pid)
print(f"sub-{pid:03d} label={labels[pid]} raw={sig.shape} "
f"resampled={res.shape} windows={len(segment(res, args.window))}", flush=True)
X = np.stack(feats) # (N, 256, 19)
y = np.array(ys, dtype=np.int64)
s = np.array(subs, dtype=np.int64)
print("total", X.shape, np.bincount(y))
# subject-wise 60/20/20 within each class (Medformer ADFTDLoader, a=0.6, b=0.8)
train_ids, val_ids, test_ids = [], [], []
for cls in (0, 1, 2):
ids = sorted([p for p in labels if labels[p] == cls])
a, b = int(0.6 * len(ids)), int(0.8 * len(ids))
train_ids += ids[:a]
val_ids += ids[a:b]
test_ids += ids[b:]
print("subjects train/val/test:", len(train_ids), len(val_ids), len(test_ids))
os.makedirs(args.out, exist_ok=True)
for name, ids in (("train", train_ids), ("val", val_ids), ("test", test_ids)):
m = np.isin(s, ids)
Xi = X[m]
mu = Xi.mean(axis=1, keepdims=True)
sd = Xi.std(axis=1, keepdims=True)
sd[sd == 0] = 1.0
Xi = ((Xi - mu) / sd).astype(np.float32)
np.save(os.path.join(args.out, f"X_{name}.npy"), Xi)
np.save(os.path.join(args.out, f"y_{name}.npy"), y[m])
np.save(os.path.join(args.out, f"subj_{name}.npy"), s[m])
print(name, Xi.shape, np.bincount(y[m]))
if __name__ == "__main__":
main()