File size: 3,090 Bytes
380b161 | 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 | 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()
|