File size: 5,209 Bytes
2e913c2 | 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 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 | """Train four independent paper CNN-LSTM branches with RMSprop and MSE."""
import json
import os
import sys
from pathlib import Path
import numpy as np
import torch
import yaml
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.climatebench import ClimateBench, ClimateBenchBranch, TARGETS
class ClimateDataset(Dataset):
def __init__(self, path: Path, config: dict):
self.data = np.load(path)
expected = config["data"]
shape = (int(expected["time_steps"]), 4, int(expected["height"]), int(expected["width"]))
target_shape = (4, int(expected["height"]), int(expected["width"]))
if str(self.data["format_version"]) != expected["format_version"] or str(self.data["storage_layout"]) != "NTCHW":
raise ValueError("incompatible ClimateBench NPZ format or storage layout")
if self.data["inputs"].shape[1:] != shape or self.data["targets"].shape[1:] != target_shape:
raise ValueError(f"expected inputs [N,{shape}] and targets [N,{target_shape}]")
if self.data["inputs"].dtype != np.float32 or self.data["targets"].dtype != np.float32:
raise TypeError("inputs and targets must be float32")
if tuple(self.data["channel_names"].tolist()) != tuple(expected["channels"]):
raise ValueError("forcing channels must be [co2_cumulative,ch4,so2,bc]")
def __len__(self) -> int:
return len(self.data["inputs"])
def __getitem__(self, index: int):
return torch.from_numpy(self.data["inputs"][index]), torch.from_numpy(self.data["targets"][index])
def device_from_config(config: dict, rank: int = 0) -> torch.device:
requested = config["runtime"]["device"]
if requested == "auto":
return torch.device("cuda", rank) if torch.cuda.is_available() else torch.device("cpu")
return torch.device(requested)
def main() -> None:
config = yaml.safe_load((ROOT / "conf/config.yaml").read_text())
np.random.seed(int(config["seed"]))
torch.manual_seed(int(config["seed"]))
distributed = int(os.environ.get("WORLD_SIZE", "1")) > 1
local_rank = int(os.environ.get("LOCAL_RANK", "0"))
if distributed:
torch.distributed.init_process_group("nccl" if torch.cuda.is_available() else "gloo")
rank = torch.distributed.get_rank() if distributed else 0
device = device_from_config(config, local_rank)
if device.type == "cuda":
torch.cuda.set_device(device)
dataset = ClimateDataset(ROOT / config["data"]["root"] / "train.npz", config)
sampler = DistributedSampler(dataset, shuffle=True) if distributed else None
loader = DataLoader(dataset, batch_size=int(config["train"]["batch_size"]), sampler=sampler,
shuffle=sampler is None, num_workers=int(config["train"]["num_workers"]))
model = ClimateBench(int(config["data"]["height"]), int(config["data"]["width"])).to(device)
counts = model.parameter_counts()
expected_count = int(config["model"]["parameters_per_target"])
if any(count != expected_count for count in counts.values()):
raise RuntimeError(f"parameter count mismatch: {counts}")
wrapped = DistributedDataParallel(model, device_ids=[local_rank] if device.type == "cuda" else None) if distributed else model
optimizer = torch.optim.RMSprop(wrapped.parameters(), lr=float(config["train"]["learning_rate"]))
history = []
for epoch in range(int(config["train"]["epochs"])):
if sampler is not None:
sampler.set_epoch(epoch)
total, batches = 0.0, 0
for inputs, targets in loader:
prediction = wrapped(inputs.to(device))
loss = torch.nn.functional.mse_loss(prediction, targets.to(device))
optimizer.zero_grad(set_to_none=True)
loss.backward()
torch.nn.utils.clip_grad_norm_(wrapped.parameters(), float(config["train"]["gradient_clip_norm"]))
optimizer.step()
total += float(loss.detach())
batches += 1
history.append({"epoch": epoch + 1, "mse": total / max(batches, 1), "batches": batches})
model = wrapped.module if distributed else wrapped
if rank == 0:
checkpoint_path = ROOT / config["paths"]["checkpoint"]
metrics_path = ROOT / config["paths"]["training_metrics"]
checkpoint_path.parent.mkdir(parents=True, exist_ok=True)
metrics_path.parent.mkdir(parents=True, exist_ok=True)
torch.save({"model": model.state_dict(), "parameter_counts": counts, "targets": TARGETS,
"format_version": config["data"]["format_version"], "storage_layout": "NTCHW",
"model_config": {"height": model.height, "width": model.width}}, checkpoint_path)
metrics_path.write_text(json.dumps({"history": history, "parameter_counts": counts}, indent=2) + "\n")
print(f"checkpoint={checkpoint_path.relative_to(ROOT)} batches={history[-1]['batches']} parameters={counts}")
if distributed:
torch.distributed.destroy_process_group()
if __name__ == "__main__":
main()
|