Download scripts/train.py from OneScience-Group/ClimODE: direct link, hf CLI and curl.
- Browser
- Download file 17.2 kB
-
https://huggingface.co/OneScience-Group/ClimODE/resolve/main/scripts/train.py
- Command line
-
hf download hf://OneScience-Group/ClimODE/scripts/train.py
-
curl -L -o train.py https://huggingface.co/OneScience-Group/ClimODE/resolve/main/scripts/train.py
17.2 kB
| """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 | |
| # Allow ``python scripts/train.py`` to resolve project-local packages. | |
| 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) | |
| # Relative paths follow the project working directory, matching the | |
| # config/checkpoint conventions used by the reference earth projects. | |
| 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() | |