File size: 4,081 Bytes
e0a6aa0
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
"""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()