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