ACE2 / scripts /train.py
zhangrenchao's picture
Publish ACE2 reproduction
380b161 verified
Raw
History Blame Contribute Delete
3.09 kB
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()