"""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 # noqa: E402 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: # keep every `stride`-th window. Windows are stored in recording order, so # this thins each recording uniformly and leaves the subject-wise split, # the subject count and the class balance untouched. 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()