Download scripts/train.py from OneScience-Group/ACE2: direct link, hf CLI and curl.
- Browser
- Download file 3.09 kB
-
https://huggingface.co/OneScience-Group/ACE2/resolve/main/scripts/train.py
- Command line
-
hf download hf://OneScience-Group/ACE2/scripts/train.py
-
curl -L -o train.py https://huggingface.co/OneScience-Group/ACE2/resolve/main/scripts/train.py
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() | |