"""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()