EchoCast-3D / scripts /train.py
zhangrenchao's picture
Publish EchoCast-3D reproduction
e0a6aa0 verified
Raw
History Blame Contribute Delete
4.08 kB
"""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()