| """Pre-train + fine-tune harness for the TS-Fingerprint reproduction. |
| |
| Reproduces the protocol of Sec 4.1 / 4.2: |
| * subject-wise 60/20/20 split (Medformer protocol) |
| * Adam, lr 1e-3 (pre-train) / 1e-4 (downstream), <=100 epochs, |
| early stopping on validation F1 with patience 10 |
| * macro Accuracy / Precision / Recall / F1 / AUROC on the test split |
| """ |
|
|
| import argparse |
| import json |
| import os |
| import sys |
| import time |
|
|
| import numpy as np |
| import torch |
| import torch.nn as nn |
| from sklearn.metrics import (accuracy_score, f1_score, precision_score, |
| recall_score, roc_auc_score) |
| from torch.utils.data import DataLoader, TensorDataset |
|
|
| sys.path.insert(0, os.path.dirname(os.path.abspath(__file__))) |
| from model import build |
|
|
|
|
| def set_seed(s): |
| np.random.seed(s) |
| torch.manual_seed(s) |
| torch.cuda.manual_seed_all(s) |
|
|
|
|
| def load_split(data_dir, split, stride=1): |
| x = np.load(os.path.join(data_dir, f"X_{split}.npy"), mmap_mode="r") |
| y = np.load(os.path.join(data_dir, f"y_{split}.npy")) |
| if stride > 1: |
| |
| |
| |
| x, y = x[::stride], y[::stride] |
| x = np.ascontiguousarray(x) |
| return torch.from_numpy(x).float(), torch.from_numpy(y).long() |
|
|
|
|
| def metrics(y_true, y_prob): |
| y_pred = y_prob.argmax(1) |
| n_cls = y_prob.shape[1] |
| out = { |
| "accuracy": accuracy_score(y_true, y_pred) * 100, |
| "precision": precision_score(y_true, y_pred, average="macro", zero_division=0) * 100, |
| "recall": recall_score(y_true, y_pred, average="macro", zero_division=0) * 100, |
| "f1": f1_score(y_true, y_pred, average="macro", zero_division=0) * 100, |
| } |
| try: |
| if n_cls == 2: |
| out["auroc"] = roc_auc_score(y_true, y_prob[:, 1]) * 100 |
| else: |
| out["auroc"] = roc_auc_score(y_true, y_prob, multi_class="ovr", |
| average="macro") * 100 |
| except ValueError: |
| out["auroc"] = float("nan") |
| return out |
|
|
|
|
| AMP = {"enabled": False} |
|
|
|
|
| def autocast(device): |
| return torch.autocast("cuda", dtype=torch.float16, |
| enabled=AMP["enabled"] and device == "cuda") |
|
|
|
|
| @torch.no_grad() |
| def evaluate(model, loader, device): |
| model.eval() |
| probs, ys = [], [] |
| for xb, yb in loader: |
| with autocast(device): |
| logits = model(xb.to(device)) |
| probs.append(torch.softmax(logits.float(), -1).cpu().numpy()) |
| ys.append(yb.numpy()) |
| return metrics(np.concatenate(ys), np.concatenate(probs)) |
|
|
|
|
| def pretrain(model, loader, device, args, log=print): |
| opt = torch.optim.Adam(model.parameters(), lr=args.pre_lr) |
| scaler = torch.amp.GradScaler("cuda", enabled=AMP["enabled"] and device == "cuda") |
| hist = [] |
| for ep in range(args.pre_epochs): |
| model.train() |
| tot = rec = div = 0.0 |
| n = 0 |
| for xb, _ in loader: |
| xb = xb.to(device, non_blocking=True) |
| with autocast(device): |
| loss, lr_, ld_ = model.pretrain_step( |
| xb, mask_ratio=args.mask_ratio, lam=args.lam, use_div=args.use_div) |
| opt.zero_grad(set_to_none=True) |
| scaler.scale(loss).backward() |
| scaler.unscale_(opt) |
| torch.nn.utils.clip_grad_norm_(model.parameters(), 4.0) |
| scaler.step(opt) |
| scaler.update() |
| bs = xb.shape[0] |
| tot += loss.item() * bs |
| rec += lr_.item() * bs |
| div += ld_.item() * bs |
| n += bs |
| hist.append({"epoch": ep, "loss": tot / n, "l_rec": rec / n, "l_div": div / n}) |
| log(f"[pretrain] ep {ep:3d} loss {tot/n:.5f} rec {rec/n:.5f} div {div/n:.4f}") |
| return hist |
|
|
|
|
| def finetune(model, tr, va, te, device, args, log=print): |
| opt = torch.optim.Adam(model.parameters(), lr=args.ft_lr) |
| scaler = torch.amp.GradScaler("cuda", enabled=AMP["enabled"] and device == "cuda") |
| crit = nn.CrossEntropyLoss() |
| best_f1, best_state, patience, hist = -1.0, None, 0, [] |
| for ep in range(args.ft_epochs): |
| model.train() |
| tot, n = 0.0, 0 |
| for xb, yb in tr: |
| xb, yb = xb.to(device, non_blocking=True), yb.to(device, non_blocking=True) |
| with autocast(device): |
| loss = crit(model(xb), yb) |
| opt.zero_grad(set_to_none=True) |
| scaler.scale(loss).backward() |
| scaler.unscale_(opt) |
| torch.nn.utils.clip_grad_norm_(model.parameters(), 4.0) |
| scaler.step(opt) |
| scaler.update() |
| tot += loss.item() * xb.shape[0] |
| n += xb.shape[0] |
| vm = evaluate(model, va, device) |
| hist.append({"epoch": ep, "train_loss": tot / n, "val_f1": vm["f1"], |
| "val_acc": vm["accuracy"]}) |
| log(f"[finetune] ep {ep:3d} loss {tot/n:.4f} val_f1 {vm['f1']:.2f} " |
| f"val_acc {vm['accuracy']:.2f}") |
| if vm["f1"] > best_f1: |
| best_f1 = vm["f1"] |
| best_state = {k: v.detach().cpu().clone() for k, v in model.state_dict().items()} |
| patience = 0 |
| else: |
| patience += 1 |
| if patience >= args.patience: |
| log(f"[finetune] early stop at epoch {ep}") |
| break |
| if best_state is not None: |
| model.load_state_dict(best_state) |
| return evaluate(model, te, device), hist, best_f1 |
|
|
|
|
| def main(): |
| p = argparse.ArgumentParser() |
| p.add_argument("--data-dir", required=True) |
| p.add_argument("--model", default="tsfp", choices=["tsfp", "timae", "simmtm"]) |
| p.add_argument("--mode", default="pretrain_ft", |
| choices=["scratch", "pretrain_ft"]) |
| p.add_argument("--use-div", type=int, default=1) |
| p.add_argument("--k", type=int, default=8) |
| p.add_argument("--d-model", type=int, default=128) |
| p.add_argument("--n-heads", type=int, default=8) |
| p.add_argument("--enc-layers", type=int, default=6) |
| p.add_argument("--dec-layers", type=int, default=2) |
| p.add_argument("--patch-size", type=int, default=8) |
| p.add_argument("--mask-ratio", type=float, default=0.6) |
| p.add_argument("--lam", type=float, default=1e-4) |
| p.add_argument("--pre-lr", type=float, default=1e-3) |
| p.add_argument("--ft-lr", type=float, default=1e-4) |
| p.add_argument("--pre-epochs", type=int, default=100) |
| p.add_argument("--ft-epochs", type=int, default=100) |
| p.add_argument("--patience", type=int, default=10) |
| p.add_argument("--batch-size", type=int, default=128) |
| p.add_argument("--seeds", type=int, nargs="+", default=[41, 42, 43, 44, 45]) |
| p.add_argument("--out", default="results.json") |
| p.add_argument("--tag", default="") |
| p.add_argument("--amp", type=int, default=1) |
| p.add_argument("--stride", type=int, default=1, |
| help="keep every Nth window (compute-budget subsampling)") |
| args = p.parse_args() |
| args.use_div = bool(args.use_div) |
| AMP["enabled"] = bool(args.amp) |
|
|
| device = "cuda" if torch.cuda.is_available() else "cpu" |
| xtr, ytr = load_split(args.data_dir, "train", args.stride) |
| xva, yva = load_split(args.data_dir, "val", args.stride) |
| xte, yte = load_split(args.data_dir, "test", args.stride) |
| n_cls = int(max(ytr.max(), yva.max(), yte.max())) + 1 |
| c_in, t_len = xtr.shape[2], xtr.shape[1] |
| print(f"data {tuple(xtr.shape)} / {tuple(xva.shape)} / {tuple(xte.shape)} " |
| f"classes={n_cls} channels={c_in} T={t_len} device={device}", flush=True) |
|
|
| runs = [] |
| for seed in args.seeds: |
| t0 = time.time() |
| set_seed(seed) |
| g = torch.Generator().manual_seed(seed) |
| tr = DataLoader(TensorDataset(xtr, ytr), batch_size=args.batch_size, |
| shuffle=True, generator=g, drop_last=True, num_workers=2, |
| pin_memory=True) |
| va = DataLoader(TensorDataset(xva, yva), batch_size=512, num_workers=2) |
| te = DataLoader(TensorDataset(xte, yte), batch_size=512, num_workers=2) |
| model = build(args.model, c_in=c_in, patch_size=args.patch_size, |
| n_classes=n_cls, d_model=args.d_model, n_heads=args.n_heads, |
| enc_layers=args.enc_layers, dec_layers=args.dec_layers, |
| k=args.k, max_patches=max(512, t_len // args.patch_size + 8)) |
| model.to(device) |
| n_par = sum(p.numel() for p in model.parameters()) |
| pre_hist = [] |
| if args.mode == "pretrain_ft": |
| pre_hist = pretrain(model, tr, device, args) |
| test_m, ft_hist, best_val = finetune(model, tr, va, te, device, args) |
| dt = time.time() - t0 |
| print(f"== seed {seed} [{args.model}/{args.mode}/div={args.use_div}] " |
| f"{json.dumps({k: round(v, 2) for k, v in test_m.items()})} " |
| f"({dt:.0f}s, {n_par/1e6:.2f}M params)", flush=True) |
| runs.append({"seed": seed, "test": test_m, "best_val_f1": best_val, |
| "seconds": dt, "params": n_par, |
| "pretrain_hist": pre_hist, "finetune_hist": ft_hist}) |
|
|
| agg = {m: {"mean": float(np.mean([r["test"][m] for r in runs])), |
| "std": float(np.std([r["test"][m] for r in runs]))} |
| for m in runs[0]["test"]} |
| out = {"config": vars(args), "runs": runs, "aggregate": agg} |
| with open(args.out, "w") as f: |
| json.dump(out, f, indent=2) |
| print("AGGREGATE " + json.dumps({k: f"{v['mean']:.2f}+-{v['std']:.2f}" |
| for k, v in agg.items()}), flush=True) |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|