PET / scripts /train_pet_foundation_epoch_ckpt.py
DesonDai's picture
Add files using upload-large-folder tool
212e9d7 verified
Raw
History Blame Contribute Delete
5.07 kB
"""
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()