Download src/objectmodel_v1/train.py from bench-labs/objectmodel-v1: direct link, hf CLI and curl.
- Browser
- Download file 10.8 kB
-
https://huggingface.co/bench-labs/objectmodel-v1/resolve/main/src/objectmodel_v1/train.py
- Command line
-
hf download hf://bench-labs/objectmodel-v1/src/objectmodel_v1/train.py
-
curl -L -o train.py https://huggingface.co/bench-labs/objectmodel-v1/resolve/main/src/objectmodel_v1/train.py
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() | |