TobiasLogic's picture
Upload ObjectModel-v1: code, README, assets (v1 checkpoint)
9b92c75 verified
Raw
History Blame Contribute Delete
10.8 kB
from __future__ import annotations
import argparse
import json
import os
import time
from pathlib import Path
import torch
import torch.distributed as dist
from torch.nn.parallel import DistributedDataParallel
from torch.optim import AdamW
from torch.optim.lr_scheduler import LambdaLR
from torch.utils.data import DataLoader, DistributedSampler
from .config import apply_overrides, load_config, save_config
from .data import build_dataset, detection_collate
from .evaluate import evaluate_coco
from .losses import ObjectModelCriterion
from .model import build_model
from .utils import (
ModelEMA,
learning_rate_factor,
move_targets,
save_checkpoint,
seed_everything,
trainable_parameter_count,
)
def distributed_context() -> tuple[int, int, int]:
world_size = int(os.environ.get("WORLD_SIZE", "1"))
rank = int(os.environ.get("RANK", "0"))
local_rank = int(os.environ.get("LOCAL_RANK", "0"))
if world_size > 1:
if not torch.cuda.is_available():
raise RuntimeError("Distributed training currently requires CUDA")
torch.cuda.set_device(local_rank)
dist.init_process_group(backend="nccl")
return rank, world_size, local_rank
def build_optimizer(model, config: dict) -> AdamW:
train = config["train"]
backbone, other = [], []
for name, parameter in model.named_parameters():
if not parameter.requires_grad:
continue
(backbone if name.startswith("backbone.") else other).append(parameter)
return AdamW(
[
{"params": other, "lr": float(train["lr"])},
{"params": backbone, "lr": float(train.get("backbone_lr", train["lr"]))},
],
weight_decay=float(train["weight_decay"]),
)
def reduce_losses(losses: dict[str, torch.Tensor], world_size: int) -> dict[str, float]:
values = torch.stack([value.detach() for value in losses.values()])
if world_size > 1:
dist.all_reduce(values)
values /= world_size
return {name: float(value) for name, value in zip(losses, values, strict=True)}
def main() -> None:
parser = argparse.ArgumentParser(description="Train ObjectModel-v1")
parser.add_argument("--config", default="configs/objectmodel_v1.yaml")
parser.add_argument("--data-root", required=True)
parser.add_argument("--output", default="outputs/objectmodel_v1")
parser.add_argument("--resume")
parser.add_argument("--device", default="cuda" if torch.cuda.is_available() else "cpu")
parser.add_argument("--set", action="append", default=[])
args = parser.parse_args()
rank, world_size, local_rank = distributed_context()
config = apply_overrides(load_config(args.config), args.set)
train_config = config["train"]
seed_everything(int(train_config["seed"]) + rank)
device = torch.device(f"cuda:{local_rank}" if world_size > 1 else args.device)
if device.type == "cuda":
torch.backends.cudnn.benchmark = True
torch.backends.cuda.matmul.allow_tf32 = True
torch.backends.cudnn.allow_tf32 = True
torch.set_float32_matmul_precision("high")
channels_last = bool(train_config.get("channels_last", False))
compile_model = bool(train_config.get("compile", False))
output_dir = Path(args.output)
if rank == 0:
output_dir.mkdir(parents=True, exist_ok=True)
save_config(config, output_dir / "config.yaml")
train_dataset = build_dataset(config, args.data_root, "train")
train_sampler = DistributedSampler(train_dataset, shuffle=True) if world_size > 1 else None
train_loader = DataLoader(
train_dataset,
batch_size=int(train_config["batch_size"]),
shuffle=train_sampler is None,
sampler=train_sampler,
num_workers=int(train_config["workers"]),
pin_memory=device.type == "cuda",
drop_last=True,
persistent_workers=int(train_config["workers"]) > 0,
prefetch_factor=int(train_config.get("prefetch_factor", 4))
if int(train_config["workers"]) > 0
else None,
collate_fn=detection_collate,
)
model = build_model(config).to(device)
if channels_last:
model = model.to(memory_format=torch.channels_last)
criterion = ObjectModelCriterion(config).to(device)
optimizer = build_optimizer(model, config)
total_steps = int(train_config["epochs"]) * len(train_loader)
scheduler = LambdaLR(
optimizer,
lambda step: learning_rate_factor(
step,
total_steps,
int(train_config["warmup_steps"]),
float(train_config["min_lr_ratio"]),
),
)
ema = ModelEMA(model, float(train_config["ema_decay"])) if rank == 0 else None
start_epoch, global_step, best_ap = 0, 0, -1.0
if args.resume:
checkpoint = torch.load(args.resume, map_location="cpu", weights_only=False)
model.load_state_dict(checkpoint["model"])
optimizer.load_state_dict(checkpoint["optimizer"])
scheduler.load_state_dict(checkpoint["scheduler"])
start_epoch = int(checkpoint["epoch"]) + 1
global_step = int(checkpoint.get("global_step", start_epoch * len(train_loader)))
best_ap = float(checkpoint.get("best_ap", -1.0))
if ema is not None and "ema" in checkpoint:
ema.model.load_state_dict(checkpoint["ema"])
if rank == 0:
print(
json.dumps(
{
"parameters": trainable_parameter_count(model),
"world_size": world_size,
"device": str(device),
"steps_per_epoch": len(train_loader),
},
indent=2,
)
)
training_model = (
DistributedDataParallel(model, device_ids=[local_rank], find_unused_parameters=False)
if world_size > 1
else model
)
if compile_model:
training_model = torch.compile(training_model, dynamic=False, mode="reduce-overhead")
use_amp = bool(train_config.get("amp", True)) and device.type == "cuda"
amp_dtype = (
torch.bfloat16 if train_config.get("amp_dtype", "float16") == "bfloat16" else torch.float16
)
scaler = torch.amp.GradScaler("cuda", enabled=use_amp and amp_dtype == torch.float16)
history_path = output_dir / "metrics.jsonl"
for epoch in range(start_epoch, int(train_config["epochs"])):
if train_sampler is not None:
train_sampler.set_epoch(epoch)
training_model.train()
epoch_start = time.perf_counter()
log_start = epoch_start
running = torch.zeros((), device=device)
for batch_index, (images, targets) in enumerate(train_loader):
images = images.to(
device,
non_blocking=True,
memory_format=torch.channels_last if channels_last else torch.preserve_format,
)
targets = move_targets(targets, device)
optimizer.zero_grad(set_to_none=True)
with torch.autocast(device_type=device.type, dtype=amp_dtype, enabled=use_amp):
outputs = training_model(images)
losses = criterion(outputs, targets)
scaler.scale(losses["loss_total"]).backward()
scaler.unscale_(optimizer)
torch.nn.utils.clip_grad_norm_(
training_model.parameters(), float(train_config["clip_grad_norm"])
)
scaler.step(optimizer)
scaler.update()
scheduler.step()
global_step += 1
if ema is not None:
ema.update(model)
running += losses["loss_total"].detach()
if rank == 0 and (batch_index + 1) % int(train_config["print_freq"]) == 0:
now = time.perf_counter()
log_steps = int(train_config["print_freq"])
avg_loss = running / (batch_index + 1)
if world_size > 1:
dist.all_reduce(avg_loss)
avg_loss = avg_loss / world_size
print(
f"epoch={epoch + 1} step={batch_index + 1}/{len(train_loader)} "
f"loss={avg_loss.item():.4f} lr={scheduler.get_last_lr()[0]:.3e} "
f"step_seconds={(now - log_start) / log_steps:.3f} "
f"images_per_second={log_steps * len(images) / (now - log_start):.2f}",
flush=True,
)
log_start = now
epoch_avg_loss = running / max(len(train_loader), 1)
if world_size > 1:
dist.all_reduce(epoch_avg_loss)
epoch_avg_loss = epoch_avg_loss / world_size
metrics: dict[str, float] = {
"epoch": epoch + 1,
"train_loss": epoch_avg_loss.item(),
"epoch_seconds": time.perf_counter() - epoch_start,
}
if world_size > 1:
dist.barrier()
should_evaluate = (epoch + 1) % int(train_config["eval_every"]) == 0
if rank == 0 and should_evaluate:
val_dataset = build_dataset(config, args.data_root, "val")
val_loader = DataLoader(
val_dataset,
batch_size=int(train_config.get("eval_batch_size", train_config["batch_size"])),
shuffle=False,
num_workers=int(train_config["workers"]),
pin_memory=device.type == "cuda",
collate_fn=detection_collate,
)
metrics.update(
evaluate_coco(
ema.model if ema is not None else model,
val_loader,
device,
output_dir / f"predictions_epoch_{epoch + 1:03d}.json",
)
)
if rank == 0:
state = {
"epoch": epoch,
"global_step": global_step,
"best_ap": max(best_ap, metrics.get("AP", -1.0)),
"model": model.state_dict(),
"ema": ema.model.state_dict() if ema is not None else model.state_dict(),
"optimizer": optimizer.state_dict(),
"scheduler": scheduler.state_dict(),
"config": config,
}
save_checkpoint(output_dir / "last.pt", **state)
if metrics.get("AP", -1.0) > best_ap:
best_ap = metrics["AP"]
state["best_ap"] = best_ap
save_checkpoint(output_dir / "best.pt", **state)
with history_path.open("a", encoding="utf-8") as handle:
handle.write(json.dumps(metrics) + "\n")
print(json.dumps(metrics))
if world_size > 1:
dist.barrier()
if world_size > 1:
dist.destroy_process_group()
if __name__ == "__main__":
main()