from pathlib import Path import sys import numpy as np import torch from torch.nn.parallel import DistributedDataParallel from torch.utils.data import DataLoader, Dataset, DistributedSampler ROOT = Path(__file__).resolve().parents[1] sys.path.insert(0, str(ROOT)) from model.ace2 import build_model, hard_correct, init_distributed, load_config, seed_all class Windows(Dataset): def __init__(self, path): data = np.load(path) self.state = data["state"] self.forcing = data["forcing"] self.windows = [(n, t) for n in range(len(self.state)) for t in range(self.state.shape[1] - 2)] def __len__(self): return len(self.windows) def __getitem__(self, index): n, t = self.windows[index] return tuple(torch.from_numpy(x.astype(np.float32)) for x in ( self.state[n, t], self.state[n, t + 1], self.state[n, t + 2], self.forcing[n, t + 1], self.forcing[n, t + 2])) def main(): cfg = load_config(ROOT) seed_all(cfg["seed"]) distributed, rank, device = init_distributed() dataset = Windows(ROOT / cfg["data"]["path"]) sampler = DistributedSampler(dataset, shuffle=True) if distributed else None loader = DataLoader(dataset, batch_size=cfg["train"]["batch_size"], sampler=sampler, shuffle=sampler is None, num_workers=0) model = build_model(cfg).to(device) if distributed: model = DistributedDataParallel(model, device_ids=[device.index] if device.type == "cuda" else None) optimizer = torch.optim.Adam(model.parameters(), lr=cfg["train"]["learning_rate"]) for epoch in range(cfg["train"]["epochs"]): if sampler: sampler.set_epoch(epoch) total = 0.0 for x0, y1, y2, f1, f2 in loader: x0, y1, y2, f1, f2 = (x.to(device) for x in (x0, y1, y2, f1, f2)) p1 = hard_correct(x0, model(x0, f1)) p2 = hard_correct(p1, model(p1, f2)) loss = torch.mean((p1 - y1) ** 2) + cfg["train"]["two_step_weight"] * torch.mean((p2 - y2) ** 2) optimizer.zero_grad(set_to_none=True) loss.backward() optimizer.step() total += loss.item() if rank == 0: print(f"epoch={epoch + 1} two_step_loss={total / len(loader):.7f}") if rank == 0: path = ROOT / cfg["train"]["checkpoint"] path.parent.mkdir(parents=True, exist_ok=True) raw_model = model.module if distributed else model torch.save({"model": raw_model.state_dict(), "model_config": cfg["model"], "format_version": cfg["data"]["format_version"]}, path) metrics = ROOT / "result/training/metrics.json" metrics.parent.mkdir(parents=True, exist_ok=True) metrics.write_text(__import__("json").dumps({"history": [{"epoch": cfg["train"]["epochs"], "loss": total / len(loader)}], "world_size": int(__import__("os").environ.get("WORLD_SIZE", "1"))}, indent=2)) print(f"saved {path}") if distributed: torch.distributed.destroy_process_group() if __name__ == "__main__": main()