| """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() |
|
|