| """Train or fine-tune ClimODE with OneScience ERA5 data.""" |
|
|
| from __future__ import annotations |
|
|
| import argparse |
| import json |
| import os |
| import random |
| import sys |
| from pathlib import Path |
|
|
| import numpy as np |
| import torch |
| import torch.distributed as dist |
| import torch.nn as nn |
| import yaml |
| from torch.nn.parallel import DistributedDataParallel |
| from torch.utils.data import DataLoader |
| from torch.utils.data.distributed import DistributedSampler |
|
|
| |
| PROJECT_ROOT = Path(__file__).resolve().parents[1] |
| if str(PROJECT_ROOT) not in sys.path: |
| sys.path.insert(0, str(PROJECT_ROOT)) |
|
|
| from model.climode import ClimODE, load_checkpoint |
| from scripts.data_loader import ClimODESeriesDataset, load_constants |
| from scripts.velocity import fit_velocity_cache, load_velocity_cache |
|
|
|
|
| def set_seed(seed: int) -> None: |
| random.seed(seed) |
| np.random.seed(seed) |
| torch.manual_seed(seed) |
| if torch.cuda.is_available(): |
| torch.cuda.manual_seed_all(seed) |
| torch.backends.cudnn.deterministic = True |
| torch.backends.cudnn.benchmark = False |
|
|
|
|
| def _device(value: str | None) -> torch.device: |
| if value: |
| return torch.device(value) |
| return torch.device("cuda" if torch.cuda.is_available() else "cpu") |
|
|
|
|
| def _init_distributed(backend: str) -> tuple[bool, int, int]: |
| world_size = int(os.environ.get("WORLD_SIZE", "1")) |
| if world_size == 1: |
| return False, 0, 1 |
| if not dist.is_initialized(): |
| dist.init_process_group(backend=backend) |
| return True, dist.get_rank(), world_size |
|
|
|
|
| def _nll(mean: torch.Tensor, std: torch.Tensor, truth: torch.Tensor, var_coeff: float) -> torch.Tensor: |
| distribution = torch.distributions.Normal(mean, 1.0e-3 + std) |
| return (-distribution.log_prob(truth)).mean() + var_coeff * (std.square()).sum() |
|
|
|
|
| def _load_yaml(path: Path) -> dict: |
| with path.open("r", encoding="utf-8") as handle: |
| return yaml.safe_load(handle) |
|
|
|
|
| def _resolve(path: str | Path) -> Path: |
| value = Path(path) |
| return value if value.is_absolute() else PROJECT_ROOT / value |
|
|
|
|
| def _parse_years(value: str | None, fallback: list[int]) -> list[int]: |
| if value is None: |
| return list(fallback) |
| years = [int(item.strip()) for item in value.split(",") if item.strip()] |
| if not years: |
| raise ValueError("year override must contain at least one integer") |
| return years |
|
|
|
|
| def _model_from_args(config: dict, args: argparse.Namespace, device: torch.device) -> nn.Module: |
| model_cfg = config["model"] |
| use_pretrained = bool(getattr(args, "use_pretrained", False)) |
| pretrained_checkpoint = getattr(args, "pretrained_checkpoint", None) |
| if args.mode == "resume": |
| checkpoint = args.checkpoint or _resolve(model_cfg["default_checkpoint"]) |
| model = load_checkpoint(checkpoint, map_location="cpu") |
| elif args.mode == "finetune" or use_pretrained: |
| checkpoint = args.checkpoint |
| if checkpoint is None and (use_pretrained or pretrained_checkpoint is not None): |
| checkpoint = pretrained_checkpoint or model_cfg.get("pretrained_checkpoint") |
| if checkpoint is None: |
| raise ValueError( |
| "finetune requires --checkpoint, or explicitly pass " |
| "--use-pretrained [--pretrained-checkpoint PATH]" |
| ) |
| checkpoint = _resolve(checkpoint) |
| if not checkpoint.is_file(): |
| raise FileNotFoundError(f"Checkpoint not found: {checkpoint}") |
| model = load_checkpoint(checkpoint, map_location="cpu") |
| else: |
| model = ClimODE( |
| num_channels=5, |
| const_channels=2, |
| out_types=5, |
| method=args.solver or model_cfg.get("solver", "euler"), |
| use_attention=model_cfg.get("use_attention", True), |
| use_uncertainty=model_cfg.get("use_uncertainty", True), |
| use_positional_encoder=model_cfg.get("use_positional_encoder", False), |
| ) |
| return model.to(device) |
|
|
|
|
| def _run_epoch( |
| model, |
| loader, |
| velocity, |
| constants, |
| lat, |
| lon, |
| device, |
| optimizer, |
| var_coeff, |
| max_batches, |
| atol, |
| rtol, |
| ): |
| training = optimizer is not None |
| model.train(training) |
| total = 0.0 |
| count = 0 |
| for batch_index, batch in enumerate(loader): |
| if max_batches is not None and batch_index >= max_batches: |
| break |
| observations = batch["observations"].squeeze(0).to(device) |
| time_steps = batch["time_steps"].squeeze(0).to(device) |
| sequence_index = int(batch["sequence_index"].item()) |
| past_velocity = velocity[sequence_index].to(device) |
| target = observations |
| initial = observations[0].unsqueeze(1) |
| model_core = model.module if isinstance(model, DistributedDataParallel) else model |
| model_core.update_param([past_velocity, constants, lat, lon]) |
| if training: |
| optimizer.zero_grad(set_to_none=True) |
| with torch.set_grad_enabled(training): |
| mean, std, _ = model(time_steps, initial, atol=atol, rtol=rtol) |
| loss = _nll(mean, std, target, var_coeff) |
| loss = loss + 0.001 * sum(parameter.square().sum() for parameter in model.parameters()) |
| if training: |
| loss.backward() |
| optimizer.step() |
| total += float(loss.detach()) |
| count += 1 |
| if dist.is_initialized(): |
| totals = torch.tensor([total, float(count)], dtype=torch.float64, device=device) |
| dist.all_reduce(totals, op=dist.ReduceOp.SUM) |
| total, count = float(totals[0].item()), int(totals[1].item()) |
| return total / max(count, 1), count |
|
|
|
|
| def _prepare_velocity( |
| dataset, |
| constants: torch.Tensor, |
| lat: torch.Tensor, |
| lon: torch.Tensor, |
| path: Path, |
| epochs: int, |
| learning_rate: float, |
| smoothing_alpha: float, |
| kernel_sigma: float, |
| distributed: bool, |
| rank: int, |
| ) -> torch.Tensor: |
| """Build a split cache once, then let every DDP rank read the same result.""" |
|
|
| if path.is_file(): |
| return load_velocity_cache(path, len(dataset)) |
| if distributed: |
| if rank == 0: |
| fit_velocity_cache( |
| dataset, |
| constants, |
| lat, |
| lon, |
| path, |
| epochs=epochs, |
| learning_rate=learning_rate, |
| smoothing_alpha=smoothing_alpha, |
| kernel_sigma=kernel_sigma, |
| ) |
| dist.barrier() |
| return load_velocity_cache(path, len(dataset)) |
| return fit_velocity_cache( |
| dataset, |
| constants, |
| lat, |
| lon, |
| path, |
| epochs=epochs, |
| learning_rate=learning_rate, |
| smoothing_alpha=smoothing_alpha, |
| kernel_sigma=kernel_sigma, |
| ) |
|
|
|
|
| def main() -> None: |
| parser = argparse.ArgumentParser(description=__doc__) |
| parser.add_argument( |
| "--config", type=Path, default=PROJECT_ROOT / "conf/config.yaml" |
| ) |
| parser.add_argument("--mode", choices=["scratch", "finetune", "resume"], default=None) |
| parser.add_argument("--checkpoint", type=Path, default=None) |
| parser.add_argument( |
| "--use-pretrained", |
| action="store_true", |
| help="Explicitly initialize from the official pretrained checkpoint", |
| ) |
| parser.add_argument( |
| "--pretrained-checkpoint", |
| type=Path, |
| default=None, |
| help="Override model.pretrained_checkpoint when --use-pretrained is set", |
| ) |
| parser.add_argument("--solver", choices=["euler", "rk4", "dopri5", "dopri8", "midpoint"], default=None) |
| parser.add_argument("--epochs", type=int, default=None) |
| parser.add_argument("--sequence-length", type=int, default=None) |
| parser.add_argument("--velocity-epochs", type=int, default=None) |
| parser.add_argument("--velocity-cache", type=Path, default=None) |
| parser.add_argument("--data-dir", type=Path, default=None, help="Override data.data_dir") |
| parser.add_argument("--stats-dir", type=Path, default=None, help="Override data.stats_dir") |
| parser.add_argument("--static-file", type=Path, default=None, help="Override data.static_file") |
| parser.add_argument("--checkpoint-dir", type=Path, default=None) |
| parser.add_argument("--log-file", type=Path, default=None) |
| parser.add_argument("--device", type=str, default=None) |
| parser.add_argument("--max-batches", type=int, default=None) |
| parser.add_argument("--seed", type=int, default=None) |
| parser.add_argument("--train-years", type=str, default=None, help="Comma-separated year override") |
| parser.add_argument("--val-years", type=str, default=None, help="Comma-separated year override") |
| args = parser.parse_args() |
| args.config = _resolve(args.config) |
| args.checkpoint = _resolve(args.checkpoint) if args.checkpoint is not None else None |
| args.pretrained_checkpoint = ( |
| _resolve(args.pretrained_checkpoint) |
| if args.pretrained_checkpoint is not None |
| else None |
| ) |
| args.data_dir = _resolve(args.data_dir) if args.data_dir is not None else None |
| args.stats_dir = _resolve(args.stats_dir) if args.stats_dir is not None else None |
| args.static_file = _resolve(args.static_file) if args.static_file is not None else None |
| args.velocity_cache = ( |
| _resolve(args.velocity_cache) if args.velocity_cache is not None else None |
| ) |
| args.checkpoint_dir = ( |
| _resolve(args.checkpoint_dir) if args.checkpoint_dir is not None else None |
| ) |
| args.log_file = _resolve(args.log_file) if args.log_file is not None else None |
| config = _load_yaml(args.config) |
| model_cfg, data_cfg, vel_cfg, train_cfg = config["model"], config["data"], config["velocity"], config["training"] |
| args.mode = args.mode or train_cfg.get("mode", "scratch") |
| if args.mode == "resume" and args.use_pretrained: |
| raise ValueError("--use-pretrained cannot be combined with --mode resume") |
| if ( |
| args.pretrained_checkpoint is not None |
| and args.mode not in {"finetune"} |
| and not args.use_pretrained |
| ): |
| raise ValueError( |
| "--pretrained-checkpoint requires --use-pretrained or " |
| "--mode finetune" |
| ) |
| if args.mode == "scratch" and args.checkpoint is not None: |
| raise ValueError( |
| "--checkpoint is ignored in scratch mode; use --mode resume or " |
| "--mode finetune explicitly" |
| ) |
| if args.use_pretrained and args.mode == "scratch": |
| args.mode = "finetune" |
| args.solver = args.solver or model_cfg.get("solver", "euler") |
| args.sequence_length = args.sequence_length or data_cfg.get("sequence_length", 8) |
| args.velocity_epochs = args.velocity_epochs if args.velocity_epochs is not None else vel_cfg.get("epochs", 200) |
| args.max_batches = args.max_batches if args.max_batches is not None else train_cfg.get("max_batches") |
| set_seed(args.seed if args.seed is not None else train_cfg.get("seed", 42)) |
| distributed, rank, world_size = _init_distributed(train_cfg.get("ddp_backend", "nccl")) |
| device = _device(args.device) |
| if distributed and device.type == "cuda": |
| device = torch.device("cuda", int(os.environ.get("LOCAL_RANK", "0"))) |
| if device.type == "cuda": |
| if device.index is None: |
| device = torch.device("cuda", 0) |
| torch.cuda.set_device(device) |
|
|
| root = _resolve(args.data_dir or data_cfg["data_dir"]) |
| stats_dir = _resolve(args.stats_dir or data_cfg.get("stats_dir", root / "static")) |
| train_set = ClimODESeriesDataset( |
| root, |
| _parse_years(args.train_years, data_cfg["train_years"]), |
| stats_dir=stats_dir, |
| model_size=(data_cfg["model_height"], data_cfg["model_width"]), |
| sequence_length=args.sequence_length, |
| normalize=data_cfg.get("normalize", True), |
| ) |
| val_set = ClimODESeriesDataset( |
| root, |
| _parse_years(args.val_years, data_cfg["val_years"]), |
| stats_dir=stats_dir, |
| model_size=(data_cfg["model_height"], data_cfg["model_width"]), |
| sequence_length=args.sequence_length, |
| normalize=data_cfg.get("normalize", True), |
| ) |
| train_sampler = DistributedSampler(train_set, shuffle=True) if distributed else None |
| val_sampler = DistributedSampler(val_set, shuffle=False) if distributed else None |
| train_loader = DataLoader(train_set, batch_size=1, sampler=train_sampler, shuffle=train_sampler is None, num_workers=data_cfg["dataloader"]["num_workers"]) |
| val_loader = DataLoader(val_set, batch_size=1, sampler=val_sampler, shuffle=False, num_workers=data_cfg["dataloader"]["num_workers"]) |
| static_file = _resolve(args.static_file or data_cfg["static_file"]) |
| constants, lat, lon = load_constants(static_file, (data_cfg["model_height"], data_cfg["model_width"])) |
| constants, lat, lon = constants.to(device), lat.unsqueeze(0).to(device), lon.unsqueeze(0).to(device) |
|
|
| |
| |
| velocity_root = _resolve(vel_cfg["cache_dir"]) |
| velocity_path = args.velocity_cache or (velocity_root / "train.pt") |
| train_velocity = _prepare_velocity( |
| train_set, |
| constants, |
| lat.squeeze(0).cpu(), |
| lon.squeeze(0).cpu(), |
| velocity_path, |
| epochs=args.velocity_epochs, |
| learning_rate=vel_cfg["learning_rate"], |
| smoothing_alpha=vel_cfg["smoothing_alpha"], |
| kernel_sigma=vel_cfg["kernel_sigma"], |
| distributed=distributed, |
| rank=rank, |
| ) |
| val_velocity_path = velocity_path.with_name("val.pt") |
| val_velocity = _prepare_velocity( |
| val_set, |
| constants, |
| lat.squeeze(0).cpu(), |
| lon.squeeze(0).cpu(), |
| val_velocity_path, |
| epochs=args.velocity_epochs, |
| learning_rate=vel_cfg["learning_rate"], |
| smoothing_alpha=vel_cfg["smoothing_alpha"], |
| kernel_sigma=vel_cfg["kernel_sigma"], |
| distributed=distributed, |
| rank=rank, |
| ) |
|
|
| model = _model_from_args(config, args, device) |
| if distributed: |
| model = DistributedDataParallel(model, device_ids=[device.index] if device.type == "cuda" else None) |
| lr = train_cfg.get("finetune_learning_rate", 5.0e-5) if args.mode == "finetune" else model_cfg.get("learning_rate", 5.0e-4) |
| optimizer = torch.optim.AdamW(model.parameters(), lr=lr, weight_decay=model_cfg.get("weight_decay", 1.0e-5)) |
| epochs = args.epochs or (train_cfg.get("finetune_epochs", 40) if args.mode == "finetune" else train_cfg.get("epochs", 300)) |
| scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, epochs) |
| start_epoch = 0 |
| if args.mode == "resume": |
| resume_path = args.checkpoint or _resolve(model_cfg["default_checkpoint"]) |
| try: |
| resume = torch.load(resume_path, map_location="cpu", weights_only=True) |
| except TypeError: |
| resume = torch.load(resume_path, map_location="cpu") |
| state = resume.get("model", resume.get("state_dict")) |
| (model.module if isinstance(model, DistributedDataParallel) else model).load_state_dict(state) |
| if "optimizer" in resume: |
| optimizer.load_state_dict(resume["optimizer"]) |
| if "scheduler" in resume: |
| scheduler.load_state_dict(resume["scheduler"]) |
| start_epoch = int(resume.get("epoch", -1)) + 1 |
|
|
| checkpoint_dir = _resolve(args.checkpoint_dir or model_cfg["checkpoint_dir"]) |
| checkpoint_dir.mkdir(parents=True, exist_ok=True) |
| log_path = _resolve(args.log_file or train_cfg.get("log_file", "./result/train.jsonl")) |
| log_path.parent.mkdir(parents=True, exist_ok=True) |
| best_val = float("inf") |
| for epoch in range(start_epoch, epochs): |
| if train_sampler is not None: |
| train_sampler.set_epoch(epoch) |
| var_coeff = 1.0e-3 if epoch == 0 else 2.0 * scheduler.get_last_lr()[0] |
| train_loss, train_count = _run_epoch( |
| model, train_loader, train_velocity, constants, lat, lon, device, |
| optimizer, var_coeff, args.max_batches, model_cfg["atol"], model_cfg["rtol"] |
| ) |
| with torch.no_grad(): |
| val_loss, val_count = _run_epoch( |
| model, val_loader, val_velocity, constants, lat, lon, device, |
| None, var_coeff, args.max_batches, model_cfg["atol"], model_cfg["rtol"] |
| ) |
| scheduler.step() |
| record = {"epoch": epoch, "train_loss": train_loss, "val_loss": val_loss, "train_batches": train_count, "val_batches": val_count, "lr": scheduler.get_last_lr()[0]} |
| if rank == 0: |
| with log_path.open("a", encoding="utf-8") as handle: |
| handle.write(json.dumps(record) + "\n") |
| if val_loss < best_val: |
| best_val = val_loss |
| torch.save({"model": (model.module if isinstance(model, DistributedDataParallel) else model).state_dict(), "optimizer": optimizer.state_dict(), "scheduler": scheduler.state_dict(), "epoch": epoch}, checkpoint_dir / "model_bak.pth") |
| print(json.dumps(record)) |
| if distributed: |
| dist.destroy_process_group() |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|