| """DDP-capable EchoCast-3D training entry point.""" |
|
|
| from __future__ import annotations |
|
|
| import argparse |
| import json |
| import os |
| from pathlib import Path |
|
|
| import torch |
| import torch.distributed as dist |
| from torch.nn.parallel import DistributedDataParallel |
| from torch.utils.data import DataLoader, DistributedSampler |
| import sys |
| ROOT=Path(__file__).resolve().parents[1];sys.path.insert(0,str(ROOT)) |
|
|
| from model.echocast_3d import PackedRadarDataset, load_config, DiffusionSchedule, EchoCast3D, masked_objective |
|
|
|
|
| def main() -> None: |
| parser = argparse.ArgumentParser() |
| parser.add_argument("--config", default="conf/config.yaml") |
| parser.add_argument("--data", default="data") |
| parser.add_argument("--output", default="result/checkpoints/echocast_3d.pt") |
| args = parser.parse_args() |
| config = load_config(ROOT / args.config) |
| distributed = int(os.environ.get("WORLD_SIZE", "1")) > 1 |
| local_rank = int(os.environ.get("LOCAL_RANK", "0")) |
| if distributed: |
| dist.init_process_group("nccl" if torch.cuda.is_available() else "gloo") |
| use_cuda=torch.cuda.is_available() and torch.cuda.device_count()>=int(os.environ.get("WORLD_SIZE","1")) and config["runtime"]["device"]!="cpu" |
| device = torch.device(f"cuda:{local_rank}" if use_cuda else "cpu") |
| if device.type == "cuda": |
| torch.cuda.set_device(device) |
| torch.manual_seed(config["seed"] + local_rank) |
| dataset = PackedRadarDataset(ROOT / args.data, "train") |
| sampler = DistributedSampler(dataset, shuffle=True) if distributed else None |
| loader = DataLoader(dataset, batch_size=config["training"]["batch_size"], sampler=sampler, shuffle=sampler is None) |
| raw_model = EchoCast3D(config).to(device) |
| model = DistributedDataParallel(raw_model, device_ids=[local_rank]) if distributed else raw_model |
| optimizer = torch.optim.AdamW( |
| model.parameters(), lr=config["training"]["learning_rate"], weight_decay=config["training"]["weight_decay"] |
| ) |
| schedule = DiffusionSchedule(device=device, **{k: config["diffusion"][k] for k in ("steps", "beta_start", "beta_end")}) |
| for epoch in range(config["training"]["epochs"]): |
| if sampler: |
| sampler.set_epoch(epoch) |
| for batch in loader: |
| truth = batch["truth"].float().to(device) |
| values = batch["values"].float().to(device) |
| validity = batch["validity"].bool().to(device) |
| observed = batch["observed"].bool().to(device) |
| clean_target = truth[:, 3:] |
| timestep = torch.randint(schedule.steps, (truth.shape[0],), device=device) |
| noisy_target, _ = schedule.add_noise(clean_target, timestep) |
| token_mask = torch.rand(truth.shape[0], raw_model.token_count, device=device) < config["model"]["mask_ratio"] |
| prediction = model(values[:, :3], noisy_target, observed[:, :3], timestep, token_mask) |
| loss, parts = masked_objective( |
| raw_model, prediction, clean_target, validity, token_mask, config["training"]["reconstruction_weight"] |
| ) |
| optimizer.zero_grad(set_to_none=True) |
| loss.backward() |
| optimizer.step() |
| if local_rank == 0: |
| print(f"epoch={epoch + 1} loss={loss.item():.5f} score={parts['score']:.5f} recon={parts['reconstruction']:.5f}") |
| if local_rank == 0: |
| output = ROOT / args.output |
| output.parent.mkdir(parents=True, exist_ok=True) |
| torch.save({"model": raw_model.state_dict(), "model_config": config["model"], "format_version": config["data"]["format_version"], "config": config}, output) |
| metrics=ROOT/config["paths"]["training_metrics"];metrics.parent.mkdir(parents=True,exist_ok=True);metrics.write_text(json.dumps({"epochs":config["training"]["epochs"],"loss":float(loss),"world_size":int(os.environ.get("WORLD_SIZE","1"))},indent=2)) |
| print(f"saved {output}; geometry tokens={raw_model.token_count}, paper states {config['model']['paper_token_count']}") |
| if distributed: |
| dist.destroy_process_group() |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|