| """ |
| ReMAP-PET training with per-epoch checkpointing. |
| Saves a checkpoint every epoch for later clinical probe selection. |
| """ |
| from __future__ import annotations |
|
|
| import argparse, inspect, importlib.util, sys, types |
| from pathlib import Path |
|
|
| import torch |
| from torch import nn |
| from torch.utils.data import DataLoader |
|
|
| from pet_vlm_dataset import PETSUVRDataset, collate_pet_suvr |
| from train_pet_foundation import ( |
| PETSUVRFoundationModel, build_encoder, alignment_loss, |
| ) |
|
|
|
|
| def main(): |
| parser = argparse.ArgumentParser() |
| parser.add_argument("--backbone", default="medicalnet") |
| parser.add_argument("--encoder-train-scope", default="layer4") |
| parser.add_argument("--epochs", type=int, default=50) |
| parser.add_argument("--batch-size", type=int, default=4) |
| parser.add_argument("--lr", type=float, default=1e-5) |
| parser.add_argument("--num-workers", type=int, default=2) |
| parser.add_argument("--output-size", type=int, nargs=3, default=(96, 96, 96)) |
| parser.add_argument("--embed-dim", type=int, default=256) |
| parser.add_argument("--contrastive-weight", type=float, default=0.2) |
| parser.add_argument("--regression-weight", type=float, default=1.0) |
| parser.add_argument("--temperature", type=float, default=0.07) |
| parser.add_argument("--medicalnet-weights", type=Path, default=Path("pretrained/medicalnet/resnet_50_23dataset.pth")) |
| parser.add_argument("--train-csv", type=Path, default=Path("metadata/splits/train.csv")) |
| parser.add_argument("--val-csv", type=Path, default=Path("metadata/splits/val.csv")) |
| parser.add_argument("--out-dir", type=Path, default=Path("runs/foundation/remap_epochs")) |
| args = parser.parse_args() |
|
|
| args.out_dir.mkdir(parents=True, exist_ok=True) |
| device = torch.device("cuda" if torch.cuda.is_available() else "cpu") |
|
|
| train_set = PETSUVRDataset(args.train_csv, output_size=tuple(args.output_size)) |
| val_set = PETSUVRDataset(args.val_csv, output_size=tuple(args.output_size)) |
| train_loader = DataLoader(train_set, batch_size=args.batch_size, shuffle=True, |
| num_workers=args.num_workers, collate_fn=collate_pet_suvr) |
| val_loader = DataLoader(val_set, batch_size=args.batch_size, shuffle=False, |
| num_workers=args.num_workers, collate_fn=collate_pet_suvr) |
| n_regions = int(train_set[0]["suvr"].numel()) |
|
|
| encoder = build_encoder(args) |
| model = PETSUVRFoundationModel(encoder, n_regions, args.embed_dim, False).to(device) |
|
|
| |
| for name, param in model.pet_encoder.named_parameters(): |
| param.requires_grad = ("layer4" in name) |
|
|
| trainable = sum(p.numel() for p in model.parameters() if p.requires_grad) |
| total = sum(p.numel() for p in model.parameters()) |
| print(f"Trainable params: {trainable:,} / {total:,} ({100*trainable/total:.1f}%)") |
|
|
| optimizer = torch.optim.AdamW( |
| [p for p in model.parameters() if p.requires_grad], |
| lr=args.lr, weight_decay=1e-4, |
| ) |
|
|
| model.temperature.data.fill_(args.temperature) |
| best_val_loss = float("inf") |
|
|
| for epoch in range(1, args.epochs + 1): |
| model.train() |
| model.pet_encoder.train() |
| train_loss = 0.0 |
| for batch in train_loader: |
| image = batch["image"].to(device, non_blocking=True) |
| suvr = batch["suvr"].to(device, non_blocking=True) |
| outputs = model(image, suvr) |
| loss, _ = alignment_loss(outputs, suvr, args.contrastive_weight, args.regression_weight) |
| optimizer.zero_grad(set_to_none=True) |
| loss.backward() |
| optimizer.step() |
| train_loss += float(loss) * image.shape[0] |
|
|
| train_loss /= len(train_set) |
|
|
| model.eval() |
| val_loss = 0.0 |
| with torch.no_grad(): |
| for batch in val_loader: |
| image = batch["image"].to(device, non_blocking=True) |
| suvr = batch["suvr"].to(device, non_blocking=True) |
| outputs = model(image, suvr) |
| loss, _ = alignment_loss(outputs, suvr, args.contrastive_weight, args.regression_weight) |
| val_loss += float(loss) * image.shape[0] |
| val_loss /= len(val_set) |
|
|
| print(f"epoch={epoch} train_loss={train_loss:.6f} val_loss={val_loss:.6f}", flush=True) |
|
|
| |
| ckpt_path = args.out_dir / f"epoch_{epoch:02d}.pt" |
| torch.save({ |
| "model": model.state_dict(), |
| "args": vars(args), |
| "epoch": epoch, |
| "val_loss": val_loss, |
| }, ckpt_path) |
|
|
| if val_loss < best_val_loss: |
| best_val_loss = val_loss |
| best_path = args.out_dir / "best.pt" |
| torch.save({ |
| "model": model.state_dict(), |
| "args": vars(args), |
| "epoch": epoch, |
| "val_loss": val_loss, |
| }, best_path) |
| print(f" -> best (val_loss={val_loss:.6f})", flush=True) |
|
|
| print(f"Done. Saved {args.epochs} checkpoints to {args.out_dir}") |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|