"""Train RainNet on contiguous windows from an RYDL-style HDF5 file.""" import json import math import os from pathlib import Path import random import sys import h5py import numpy as np import torch from torch import nn import torch.nn.functional as F from torch.nn.parallel import DistributedDataParallel from torch.utils.data import DataLoader, Dataset, DistributedSampler import yaml ROOT = Path(__file__).resolve().parents[1] sys.path.insert(0, str(ROOT)) from model.rainnet import build_rainnet def load_config(): with (ROOT / "conf/config.yaml").open(encoding="utf-8") as handle: return yaml.safe_load(handle) def seed_everything(seed): random.seed(seed) np.random.seed(seed) torch.manual_seed(seed) if torch.cuda.is_available(): torch.cuda.manual_seed_all(seed) def setup_device(config): 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(backend="nccl" if torch.cuda.is_available() else "gloo") if torch.cuda.is_available() and config["device"] in ("auto", "cuda"): device = torch.device("cuda", local_rank) torch.cuda.set_device(device) else: device = torch.device("cpu") return device, distributed, local_rank class RainNetDataset(Dataset): def __init__(self, path, keys, input_steps=4): self.path = path self.keys = list(keys) self.input_steps = input_steps if len(self.keys) <= input_steps: raise ValueError("A split needs at least input_steps + 1 frames") def __len__(self): return len(self.keys) - self.input_steps def __getitem__(self, index): with h5py.File(self.path, "r") as handle: inputs = np.stack( [handle[key][...] for key in self.keys[index : index + self.input_steps]] ) target_key = self.keys[index + self.input_steps] target = handle[target_key][...][None] return torch.from_numpy(inputs), torch.from_numpy(target), target_key class LogCoshLoss(nn.Module): def forward(self, prediction, target): error = torch.abs(prediction - target) return (error + F.softplus(-2.0 * error) - math.log(2.0)).mean() def transform_and_pad(tensor, pad): return F.pad(torch.log(tensor + 0.01), pad, mode="reflect") def run_epoch(model, loader, criterion, device, pad, max_batches, optimizer=None): training = optimizer is not None model.train(training) losses = [] parameter_updated = False first_shapes = None context = torch.enable_grad() if training else torch.no_grad() with context: for batch_index, (inputs, targets, target_keys) in enumerate(loader): if batch_index >= max_batches: break inputs, targets = inputs.to(device), targets.to(device) padded_inputs = transform_and_pad(inputs, pad) padded_targets = transform_and_pad(targets, pad) if training: optimizer.zero_grad(set_to_none=True) output = model(padded_inputs) loss = criterion(output, padded_targets) if not torch.isfinite(loss): raise RuntimeError(f"Non-finite loss: {loss.item()}") if first_shapes is None: first_shapes = (inputs.shape, targets.shape, padded_inputs.shape, output.shape, target_keys[0]) if training: tracked = next(model.parameters()).detach().clone() loss.backward() optimizer.step() parameter_updated = parameter_updated or not torch.equal(tracked, next(model.parameters()).detach()) losses.append(loss.item()) return float(np.mean(losses)), parameter_updated, first_shapes def save_checkpoint(path, model, optimizer, epoch, val_loss, config): path.parent.mkdir(parents=True, exist_ok=True) state_model = model.module if isinstance(model, DistributedDataParallel) else model torch.save( { "model_state_dict": state_model.state_dict(), "optimizer_state_dict": optimizer.state_dict(), "epoch": epoch, "validation_loss": val_loss, "config": config, }, path, ) def main(): config = load_config() seed_everything(config["seed"]) device, distributed, local_rank = setup_device(config) is_main = local_rank == 0 data = config["data"] path = ROOT / data["path"] if not path.exists(): raise FileNotFoundError(f"Fake data not found: {path}; run scripts/fake_data.py") with h5py.File(path, "r") as handle: keys = sorted(handle.keys()) train_end = data["train_frames"] val_end = train_end + data["val_frames"] train_set = RainNetDataset(path, keys[:train_end], data["input_steps"]) val_set = RainNetDataset(path, keys[train_end:val_end], data["input_steps"]) train_sampler = DistributedSampler(train_set, shuffle=True) if distributed else None train_loader = DataLoader( train_set, batch_size=config["train"]["batch_size"], shuffle=train_sampler is None, sampler=train_sampler, num_workers=data["num_workers"], ) val_loader = DataLoader(val_set, batch_size=1, shuffle=False, num_workers=data["num_workers"]) model = build_rainnet(**config["model"]).to(device) parameter_count = sum(parameter.numel() for parameter in model.parameters()) if distributed: model = DistributedDataParallel(model, device_ids=[local_rank] if device.type == "cuda" else None) criterion = LogCoshLoss() optimizer = torch.optim.Adam(model.parameters(), lr=config["train"]["learning_rate"]) pad_h = data["padded_height"] - data["raw_height"] pad_w = data["padded_width"] - data["raw_width"] pad = (pad_w // 2, pad_w - pad_w // 2, pad_h // 2, pad_h - pad_h // 2) history = {"train_loss": [], "validation_loss": [], "learning_rate": []} best_loss = float("inf") any_update = False for epoch in range(config["train"]["epochs"]): if train_sampler: train_sampler.set_epoch(epoch) train_loss, updated, shapes = run_epoch( model, train_loader, criterion, device, pad, config["train"]["max_train_batches"], optimizer ) val_loss, _, _ = run_epoch( model, val_loader, criterion, device, pad, config["train"]["max_valid_batches"] ) any_update = any_update or updated history["train_loss"].append(train_loss) history["validation_loss"].append(val_loss) history["learning_rate"].append(optimizer.param_groups[0]["lr"]) if is_main: last_path = ROOT / config["train"]["checkpoint_last"] best_path = ROOT / config["train"]["checkpoint_best"] save_checkpoint(last_path, model, optimizer, epoch + 1, val_loss, config) if val_loss < best_loss: best_loss = val_loss save_checkpoint(best_path, model, optimizer, epoch + 1, val_loss, config) result_dir = ROOT / config["evaluation"]["result_dir"] result_dir.mkdir(parents=True, exist_ok=True) with (result_dir / "train_history.json").open("w", encoding="utf-8") as handle: json.dump(history, handle, indent=2) print(f"Device: {device}") print(f"Input shape: {tuple(shapes[0])}") print(f"Target shape: {tuple(shapes[1])}") print(f"Target key (i+4): {shapes[4]}") print(f"Padded input shape: {tuple(shapes[2])}") print(f"Model output shape: {tuple(shapes[3])}") print(f"Parameter count: {parameter_count}") print(f"Epoch: {epoch + 1}") print(f"Train loss: {train_loss:.8f}") print(f"Validation loss: {val_loss:.8f}") print(f"Learning rate: {optimizer.param_groups[0]['lr']}") print(f"parameter_update_detected: {any_update}") print(f"Checkpoint path: {best_path}") if not any_update: raise RuntimeError("No model parameter changed after optimizer.step()") if is_main: reload_model = build_rainnet(**config["model"]) checkpoint = torch.load( ROOT / config["train"]["checkpoint_best"], map_location="cpu", weights_only=False ) reload_model.load_state_dict(checkpoint["model_state_dict"]) print("checkpoint_reload_after_training: True") if distributed: torch.distributed.destroy_process_group() if __name__ == "__main__": main()