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