Detector / scripts /train_detector.py
Benxelua's picture
Do not back up smoke-test checkpoints
80cbeb3 verified
Raw History Blame Contribute Delete
9.16 kB
#!/usr/bin/env python3
"""Train one locked-protocol detector experiment and select only by validation AP."""
from __future__ import annotations
import argparse, csv, json, math, os, random
from pathlib import Path
import sys
import numpy as np
import torch
import yaml
from torch.utils.data import DataLoader
sys.path.insert(0, str(Path(__file__).resolve().parent))
from detector_lib import (CanonicalMAP, LockedLODDataset, build_model, collate, hf_batch, image_tensors)
def seed_all(seed):
random.seed(seed); np.random.seed(seed); torch.manual_seed(seed); torch.cuda.manual_seed_all(seed)
@torch.no_grad()
def evaluate(model, processor, kind, loader, device):
model.eval(); metric = CanonicalMAP()
for images, targets in loader:
if kind == "hf":
encoded = processor(images=list(images), return_tensors="pt")
output = model(**{k: v.to(device) for k, v in encoded.items()})
sizes = torch.stack([target["orig_size"] for target in targets]).to(device)
predictions = processor.post_process_object_detection(output, target_sizes=sizes, threshold=0.0)
else:
predictions = model(image_tensors(images, device))
normalized = [{"boxes": item["boxes"].detach().cpu(), "scores": item["scores"].detach().cpu(),
"labels": item["labels"].detach().cpu() - (1 if kind == "torchvision" else 0)} for item in predictions]
metric.update(normalized, targets)
return metric.compute()
def save_checkpoint(path, model, optimizer, scheduler, epoch, best, cfg):
torch.save({"model": model.state_dict(), "optimizer": optimizer.state_dict(),
"scheduler": scheduler.state_dict() if scheduler else None,
"epoch": epoch, "best_val_map50_95": best, "config": cfg}, path)
def verified_hf_backup(output_dir: Path, epoch: int, final: bool = False) -> None:
"""Backup recoverable artifacts when an explicit runtime credential exists.
A backup failure never discards the local checkpoint or aborts training;
verification is recorded locally only after the remote file listing agrees.
"""
repo = os.environ.get("HF_BACKUP_REPO")
prefix = os.environ.get("HF_BACKUP_PREFIX")
token = os.environ.get("HF_TOKEN") or os.environ.get("HUGGINGFACE_HUB_TOKEN")
interval = int(os.environ.get("HF_BACKUP_EVERY", "5"))
if not (repo and prefix and token) or (not final and epoch % interval):
return
try:
from huggingface_hub import HfApi
api = HfApi()
names = ["config.yaml", "train_log.csv", "best_checkpoint.pt", "last_checkpoint.pt"]
api.upload_folder(repo_id=repo, repo_type="model", folder_path=str(output_dir),
path_in_repo=prefix, allow_patterns=names,
commit_message=f"checkpoint backup epoch {epoch}", token=token)
remote = set(api.list_repo_files(repo_id=repo, repo_type="model", token=token))
required = {f"{prefix}/{name}" for name in names if (output_dir / name).is_file()}
missing = sorted(required - remote)
if missing: raise RuntimeError(f"HF verification missing {missing}")
(output_dir / "hf_backup.json").write_text(json.dumps({"repo": repo, "prefix": prefix, "epoch": epoch, "verified": True}, indent=2) + "\n")
except Exception as exc:
(output_dir / "hf_backup_warning.txt").write_text(f"epoch={epoch}\n{type(exc).__name__}: {exc}\n")
def main():
parser = argparse.ArgumentParser()
parser.add_argument("--config", type=Path, required=True)
parser.add_argument("--dataset-root", type=Path, required=True)
parser.add_argument("--labels-root", type=Path, default=Path(__file__).resolve().parents[1] / "labels")
parser.add_argument("--output-dir", type=Path, required=True)
parser.add_argument("--smoke", action="store_true")
args = parser.parse_args()
cfg = yaml.safe_load(args.config.read_text()); seed_all(cfg["seed"])
out = args.output_dir; out.mkdir(parents=True, exist_ok=True)
(out / "config.yaml").write_text(yaml.safe_dump(cfg, sort_keys=False))
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
model, processor, kind = build_model(cfg); model.to(device)
params = sum(p.numel() for p in model.parameters())
limit = 2 if args.smoke else None
train_set = LockedLODDataset(args.config.parent.parent / cfg["train_manifest"], args.dataset_root, args.labels_root, limit)
val_set = LockedLODDataset(args.config.parent.parent / cfg["val_manifest"], args.dataset_root, args.labels_root, limit)
# A smoke test demonstrates an end-to-end optimization step rather than
# reproducing production throughput. Faster R-CNN at native 1200x800 can
# exceed a 4-GB development GPU at batch 2, while the Kaggle production
# configuration remains batch 2. Keep SSDLite at two samples because its
# final BatchNorm feature map needs batch statistics during training.
smoke_batch = 1 if cfg["detector"] == "fasterrcnn_r50_fpn" else cfg["batch_size"]
train_loader = DataLoader(train_set, batch_size=min(smoke_batch if args.smoke else cfg["batch_size"], len(train_set)), shuffle=True, num_workers=0, collate_fn=collate)
val_loader = DataLoader(val_set, batch_size=1, num_workers=0, collate_fn=collate)
if cfg.get("optimizer", "adamw").lower() == "sgd":
optimizer = torch.optim.SGD(model.parameters(), lr=cfg["lr"], momentum=cfg.get("momentum", 0.9),
weight_decay=cfg["weight_decay"], nesterov=cfg.get("nesterov", False))
else:
optimizer = torch.optim.AdamW(model.parameters(), lr=cfg["lr"], weight_decay=cfg["weight_decay"])
scheduler = torch.optim.lr_scheduler.MultiStepLR(optimizer, milestones=cfg["lr_milestones"], gamma=cfg["lr_gamma"])
start, best = 1, -1.0
last = out / "last_checkpoint.pt"
if last.exists() and not args.smoke:
state = torch.load(last, map_location="cpu"); model.load_state_dict(state["model"]); optimizer.load_state_dict(state["optimizer"])
if state.get("scheduler"): scheduler.load_state_dict(state["scheduler"])
start, best = state["epoch"] + 1, state["best_val_map50_95"]
epochs = 1 if args.smoke else cfg["epochs"]
log = out / "train_log.csv"
for epoch in range(start, epochs + 1):
model.train(); total_loss = 0.0
for images, targets in train_loader:
optimizer.zero_grad(set_to_none=True)
if kind == "hf":
loss = model(**hf_batch(processor, images, targets, device)).loss
else:
model_targets = [{"boxes": t["boxes"].to(device), "labels": t["labels"].to(device) + 1} for t in targets]
loss = sum(model(image_tensors(images, device), model_targets).values())
if not torch.isfinite(loss):
raise RuntimeError(f"non-finite loss at epoch {epoch}")
loss.backward(); optimizer.step(); total_loss += float(loss.detach())
if args.smoke: break
val = evaluate(model, processor, kind, val_loader, device)
score = val["map50_95"]
if not math.isfinite(score) or score < 0: raise RuntimeError(f"invalid validation metric: {score}")
scheduler.step(); save_checkpoint(last, model, optimizer, scheduler, epoch, max(best, score), cfg)
if score > best:
best = score; save_checkpoint(out / "best_checkpoint.pt", model, optimizer, scheduler, epoch, best, cfg)
new = not log.exists()
with log.open("a", newline="") as handle:
writer = csv.DictWriter(handle, fieldnames=["epoch", "loss", "val_map50_95", "val_map50"])
if new: writer.writeheader()
writer.writerow({"epoch": epoch, "loss": total_loss, "val_map50_95": score, "val_map50": val["map50"]})
if not args.smoke:
verified_hf_backup(out, epoch, final=(epoch == epochs))
state = torch.load(out / "best_checkpoint.pt", map_location="cpu")
model.load_state_dict(state["model"])
val = evaluate(model, processor, kind, val_loader, device)
(out / "val_metrics.json").write_text(json.dumps({"best_epoch": state["epoch"], **val}, indent=2) + "\n")
if args.smoke:
(out / "smoke.json").write_text(json.dumps({"status": "passed", "params": params, "device": str(device), "validation": val}, indent=2) + "\n")
return
tests = {}
for manifest in cfg["test_manifests"]:
dataset = LockedLODDataset(args.config.parent.parent / manifest, args.dataset_root, args.labels_root)
tests[Path(manifest).stem] = evaluate(model, processor, kind, DataLoader(dataset, batch_size=1, num_workers=0, collate_fn=collate), device)
(out / "test_metrics.json").write_text(json.dumps(tests, indent=2) + "\n")
(out / "run_metadata.json").write_text(json.dumps({"experiment_id": cfg["experiment_id"], "params": params,
"torch": torch.__version__, "cuda": torch.version.cuda, "device": str(device), "pretrained": cfg["pretrained"]}, indent=2) + "\n")
verified_hf_backup(out, state["epoch"], final=True)
if __name__ == "__main__": main()