tsfp-repro-code / train.py
riteshhf's picture
Upload folder using huggingface_hub
2188a91 verified
Raw
History Blame Contribute Delete
9.68 kB
"""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()