"""Train the ConvLSTM radar encoder-forecaster with full-sequence BPTT.""" 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.convlstm import ConvLSTM class RadarDataset(Dataset): def __init__(self, path, config): self.data = np.load(path) data = config["data"] if str(self.data["format_version"]) != data["format_version"]: raise ValueError("incompatible radar data format") expected_input = (int(data["input_frames"]), int(data["channels"]), int(data["height"]), int(data["width"])) expected_target = (int(data["output_frames"]), int(data["channels"]), int(data["height"]), int(data["width"])) if self.data["inputs"].shape[1:] != expected_input or self.data["targets"].shape[1:] != expected_target: raise ValueError("radar tensors do not preserve the paper dimensions") def __len__(self): return len(self.data["inputs"]) def __getitem__(self, index): return torch.from_numpy(self.data["inputs"][index]).float(), torch.from_numpy(self.data["targets"][index]).float() def device_from_config(config, rank=0): if config["runtime"]["device"] == "auto": return torch.device("cuda", rank) if torch.cuda.is_available() else torch.device("cpu") return torch.device(config["runtime"]["device"]) def main(): config = yaml.safe_load((ROOT / "conf/config.yaml").read_text()) 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) dataset = RadarDataset(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 = ConvLSTM(config["model"]).to(device) if distributed: model = DistributedDataParallel(model, device_ids=[local_rank] if device.type == "cuda" else None) optimizer = torch.optim.RMSprop(model.parameters(), lr=float(config["train"]["learning_rate"]), alpha=float(config["train"]["rmsprop_alpha"]), weight_decay=float(config["train"]["weight_decay"])) history = [] for epoch in range(int(config["train"]["epochs"])): model.train() total, steps = 0.0, 0 for inputs, targets in loader: _, logits = model(inputs.to(device)) patched_target = torch.nn.functional.pixel_unshuffle(targets.to(device).flatten(0, 1), int(config["model"]["patch_size"])).unflatten(0, targets.shape[:2]) loss = torch.nn.functional.binary_cross_entropy_with_logits(logits, patched_target) optimizer.zero_grad(set_to_none=True) loss.backward() torch.nn.utils.clip_grad_norm_(model.parameters(), float(config["train"]["gradient_clip_norm"])) optimizer.step() total += float(loss.detach()) steps += 1 metrics = {"epoch": epoch + 1, "binary_cross_entropy": total / max(steps, 1)} history.append(metrics) if rank == 0: print(f"epoch={epoch + 1} binary_cross_entropy={metrics['binary_cross_entropy']:.6f}") if rank == 0: checkpoint, metrics_path = ROOT / config["paths"]["checkpoint"], ROOT / config["paths"]["training_metrics"] checkpoint.parent.mkdir(parents=True, exist_ok=True) metrics_path.parent.mkdir(parents=True, exist_ok=True) state = model.module.state_dict() if distributed else model.state_dict() torch.save({"model": state, "model_config": config["model"], "format_version": config["data"]["format_version"]}, checkpoint) metrics_path.write_text(json.dumps({"history": history}, indent=2) + "\n") if distributed: torch.distributed.destroy_process_group() if __name__ == "__main__": main()