""" 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) # Apply MedicalNet layer4 partial tuning 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) # Save every epoch 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()