Download code/training/src/training_validation/common.py from lsh9034/ci-net: direct link, hf CLI and curl.
- Browser
- Download file 15.5 kB
-
https://huggingface.co/lsh9034/ci-net/resolve/main/code/training/src/training_validation/common.py
- Command line
-
hf download hf://lsh9034/ci-net/code/training/src/training_validation/common.py
-
curl -L -o common.py https://huggingface.co/lsh9034/ci-net/resolve/main/code/training/src/training_validation/common.py
15.5 kB
| from __future__ import annotations | |
| import csv | |
| import json | |
| import os | |
| import random | |
| import shutil | |
| import sys | |
| from pathlib import Path | |
| from typing import Any | |
| import torch | |
| from torch.utils.data import DataLoader, Dataset | |
| CODE_ROOT = Path(__file__).resolve().parents[2] | |
| if str(CODE_ROOT) not in sys.path: | |
| sys.path.insert(0, str(CODE_ROOT)) | |
| from src.data_pipeline.basic_dataset import BasicDataset, skip_missing_collate # noqa: E402 | |
| from src.data_pipeline.fast_dataset import FastDataset # noqa: E402 | |
| from src.config import load_config as load_release_config # noqa: E402 | |
| from src.model import SimVPCI, SimVPBTAux, SwinLSTMCI # noqa: E402 | |
| from src.training_validation.ram_chunk import FastRAMChunkedDataset, FastUniqueRAMCachedDataset # noqa: E402 | |
| def set_seed(seed: int, deterministic: bool = True) -> None: | |
| random.seed(int(seed)) | |
| try: | |
| import numpy as np | |
| np.random.seed(int(seed)) | |
| except ImportError: | |
| pass | |
| os.environ["PYTHONHASHSEED"] = str(int(seed)) | |
| torch.manual_seed(int(seed)) | |
| torch.cuda.manual_seed(int(seed)) | |
| torch.cuda.manual_seed_all(int(seed)) | |
| if deterministic: | |
| torch.backends.cudnn.deterministic = True | |
| torch.backends.cudnn.benchmark = False | |
| class CachedDataset(Dataset): | |
| def __init__(self, dataset: Dataset): | |
| self.samples = [dataset[i] for i in range(len(dataset))] | |
| def __len__(self) -> int: | |
| return len(self.samples) | |
| def __getitem__(self, idx: int) -> dict[str, Any]: | |
| return self.samples[int(idx)] | |
| def load_config(path: str | Path) -> dict[str, Any]: | |
| return load_release_config(path) | |
| def config_snapshot(config: dict[str, Any]) -> dict[str, Any]: | |
| return {k: v for k, v in config.items() if not k.startswith("_")} | |
| def _evenly_spaced(items: list[Any], count: int) -> list[Any]: | |
| count = int(count) | |
| if count <= 0 or len(items) <= count: | |
| return items | |
| if count == 1: | |
| return [items[0]] | |
| last = len(items) - 1 | |
| indices = [round(i * last / (count - 1)) for i in range(count)] | |
| return [items[int(i)] for i in indices] | |
| def _apply_sample_filter(dataset: Dataset, split_cfg: dict[str, Any]) -> None: | |
| if not hasattr(dataset, "samples"): | |
| return | |
| samples = list(getattr(dataset, "samples")) | |
| sample_stride = int(split_cfg.get("sample_stride", 1) or 1) | |
| if sample_stride > 1: | |
| samples = samples[::sample_stride] | |
| max_samples = split_cfg.get("max_samples") | |
| if max_samples is not None: | |
| samples = _evenly_spaced(samples, int(max_samples)) | |
| setattr(dataset, "samples", samples) | |
| def build_dataset(config: dict[str, Any], split: str, mode: str) -> Dataset: | |
| ds_cfg = dict(config.get("dataset", {})) | |
| dataset_type = str(ds_cfg.get("type", "fast")).lower() | |
| split_cfg = dict(config.get(mode, {})) | |
| ram_chunk_cfg = _ram_chunk_config(config, split_cfg) | |
| use_ram_chunk = bool(ram_chunk_cfg.get("enabled", False)) | |
| required_inputs = list(split_cfg.get("input_sources", config.get("input_sources", config.get("required_inputs", ["concat"])))) | |
| required_labels = list(split_cfg.get("required_labels", config.get("required_labels", ["ci"]))) | |
| dataset_config = ds_cfg.get("config", config.get("dataset_config", config)) | |
| if not isinstance(dataset_config, dict): | |
| dataset_config = load_release_config(dataset_config) | |
| else: | |
| dataset_config = dict(dataset_config) | |
| if split_cfg.get("time_ranges") is not None: | |
| dataset_config.setdefault("splits", {}) | |
| dataset_config["splits"][split] = list(split_cfg["time_ranges"]) | |
| if dataset_type == "fast": | |
| dataset: Dataset = FastDataset(dataset_config, split=split, required_inputs=required_inputs, required_labels=required_labels) | |
| elif dataset_type == "basic": | |
| if use_ram_chunk: | |
| raise ValueError("ram_chunk.enabled=true only supports dataset.type=fast") | |
| dataset = BasicDataset(dataset_config, split=split, inputs=required_inputs, labels=required_labels) | |
| else: | |
| raise ValueError(f"unsupported dataset.type: {dataset_type}") | |
| _apply_sample_filter(dataset, split_cfg) | |
| cache_cfg = config.get("ram_cache", {}) | |
| use_cache = bool(split_cfg.get("use_ram_cache", cache_cfg.get(f"{mode}_use_ram_cache", cache_cfg.get("use_ram_cache", False)))) | |
| if use_ram_chunk and use_cache: | |
| raise ValueError("ram_chunk.enabled and use_ram_cache cannot be true at the same time") | |
| if use_ram_chunk: | |
| chunk_order = str(ram_chunk_cfg.get("chunk_order", "sequential")) | |
| swap_policy = str(ram_chunk_cfg.get("swap_policy", "repeat_current")) | |
| if chunk_order not in {"sequential", "circular"}: | |
| raise ValueError("ram_chunk.chunk_order currently supports only 'sequential' and 'circular'") | |
| if swap_policy != "repeat_current": | |
| raise ValueError("ram_chunk.swap_policy currently supports only 'repeat_current'") | |
| dataset = FastRAMChunkedDataset( | |
| dataset, # type: ignore[arg-type] | |
| chunk_ram_gb=float(ram_chunk_cfg.get("chunk_ram_gb", 40.0)), | |
| cache_dtype=str(ram_chunk_cfg.get("cache_dtype", "float16")), | |
| read_block_rows=int(ram_chunk_cfg.get("read_block_rows", 64)), | |
| async_prefetch=bool(ram_chunk_cfg.get("async_prefetch", True)), | |
| chunk_order=chunk_order, | |
| verbose=bool(ram_chunk_cfg.get("verbose", True)), | |
| ) | |
| if not bool(config.get("_defer_ram_chunk_initial_load", False)): | |
| initial_chunk_id = int(ram_chunk_cfg.get("initial_chunk_id", 0)) | |
| dataset.load_chunk_sync(initial_chunk_id) # type: ignore[attr-defined] | |
| if use_cache: | |
| if dataset_type == "fast": | |
| cache_dtype = str(cache_cfg.get("cache_dtype", config.get("ram_chunk", {}).get("cache_dtype", "float16"))) | |
| read_block_rows = int(cache_cfg.get("read_block_rows", config.get("ram_chunk", {}).get("read_block_rows", 64))) | |
| verbose = bool(cache_cfg.get("verbose", config.get("ram_chunk", {}).get("verbose", True))) | |
| dataset = FastUniqueRAMCachedDataset( | |
| dataset, # type: ignore[arg-type] | |
| cache_dtype=cache_dtype, | |
| read_block_rows=read_block_rows, | |
| verbose=verbose, | |
| ) | |
| else: | |
| dataset = CachedDataset(dataset) | |
| return dataset | |
| def is_ram_chunk_dataset(dataset: Dataset) -> bool: | |
| return isinstance(dataset, FastRAMChunkedDataset) | |
| def _ram_chunk_config(config: dict[str, Any], split_cfg: dict[str, Any]) -> dict[str, Any]: | |
| merged = dict(config.get("ram_chunk", {})) | |
| split_ram_chunk = split_cfg.get("ram_chunk") | |
| if isinstance(split_ram_chunk, dict): | |
| merged.update(split_ram_chunk) | |
| if "use_ram_chunk" in split_cfg: | |
| merged["enabled"] = bool(split_cfg["use_ram_chunk"]) | |
| return merged | |
| def build_dataloader(config: dict[str, Any], dataset: Dataset, mode: str) -> DataLoader: | |
| loader_cfg = dict(config.get("dataloader", {})) | |
| split_cfg = dict(config.get(mode, {})) | |
| batch_size = int(split_cfg.get("batch_size", loader_cfg.get("batch_size", 1))) | |
| num_workers = int(split_cfg.get("num_workers", loader_cfg.get("num_workers", 0))) | |
| shuffle_default = mode == "train" | |
| shuffle = bool(split_cfg.get("shuffle", loader_cfg.get(f"{mode}_shuffle", shuffle_default))) | |
| ds_cfg = dict(config.get("dataset", {})) | |
| collate_fn = None | |
| if str(ds_cfg.get("type", "fast")).lower() == "basic": | |
| required_inputs = list(split_cfg.get("input_sources", config.get("input_sources", config.get("required_inputs", ["concat"])))) | |
| required_labels = list(split_cfg.get("required_labels", config.get("required_labels", ["ci"]))) | |
| collate_fn = skip_missing_collate(required_inputs, required_labels) | |
| loader_kwargs = { | |
| "batch_size": batch_size, | |
| "shuffle": shuffle, | |
| "num_workers": num_workers, | |
| "pin_memory": bool(loader_cfg.get("pin_memory", True)), | |
| "drop_last": bool(split_cfg.get("drop_last", mode == "train")), | |
| "collate_fn": collate_fn, | |
| } | |
| if num_workers > 0: | |
| loader_kwargs["persistent_workers"] = bool( | |
| split_cfg.get("persistent_workers", loader_cfg.get("persistent_workers", False)) | |
| ) | |
| loader_kwargs["prefetch_factor"] = int( | |
| split_cfg.get("prefetch_factor", loader_cfg.get("prefetch_factor", 2)) | |
| ) | |
| return DataLoader(dataset, **loader_kwargs) | |
| def shutdown_dataloader(loader: DataLoader | None) -> None: | |
| if loader is None: | |
| return | |
| iterator = getattr(loader, "_iterator", None) | |
| if iterator is not None: | |
| shutdown = getattr(iterator, "_shutdown_workers", None) | |
| if shutdown is not None: | |
| try: | |
| shutdown() | |
| except Exception: | |
| pass | |
| del loader | |
| def build_model(config: dict[str, Any]) -> torch.nn.Module: | |
| model_cfg = dict(config.get("model", {})) | |
| name = str(model_cfg.get("name", "simvp_ci")).lower() | |
| params = dict(model_cfg.get("params", {})) | |
| if name in {"simvp_ci", "simvp"}: | |
| return SimVPCI(**params) | |
| if name in {"simvp_bt_aux", "simvp_ci_bt", "cinet_bt_predict"}: | |
| return SimVPBTAux(**params) | |
| if name in {"swinlstm_ci", "swinlstm", "swinlstm_b", "swinlstm_d"}: | |
| if name == "swinlstm_b": | |
| params.setdefault("variant", "b") | |
| elif name == "swinlstm_d": | |
| params.setdefault("variant", "d") | |
| return SwinLSTMCI(**params) | |
| raise ValueError(f"unsupported model.name: {name}") | |
| def build_optimizer(config: dict[str, Any], model: torch.nn.Module) -> torch.optim.Optimizer: | |
| opt_cfg = dict(config.get("optimizer", {})) | |
| name = str(opt_cfg.get("name", "adamw")).lower() | |
| params = dict(opt_cfg.get("params", {})) | |
| if name == "adam": | |
| return torch.optim.Adam(model.parameters(), **params) | |
| if name == "adamw": | |
| return torch.optim.AdamW(model.parameters(), **params) | |
| if name == "sgd": | |
| return torch.optim.SGD(model.parameters(), **params) | |
| raise ValueError(f"unsupported optimizer.name: {name}") | |
| def build_scheduler(config: dict[str, Any], optimizer: torch.optim.Optimizer): | |
| sched_cfg = dict(config.get("scheduler", {})) | |
| name = str(sched_cfg.get("name", "none")).lower() | |
| params = dict(sched_cfg.get("params", {})) | |
| if name in {"none", "null", ""}: | |
| return None | |
| if name == "cosine": | |
| return torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, **params) | |
| if name == "step": | |
| return torch.optim.lr_scheduler.StepLR(optimizer, **params) | |
| if name == "plateau": | |
| return torch.optim.lr_scheduler.ReduceLROnPlateau(optimizer, **params) | |
| raise ValueError(f"unsupported scheduler.name: {name}") | |
| def pack_inputs(batch: dict[str, Any], input_sources: list[str], device: torch.device) -> torch.Tensor: | |
| arrays = [] | |
| for name in input_sources: | |
| value = batch["inputs"].get(name) | |
| if value is None: | |
| raise ValueError(f"batch input {name!r} is None") | |
| arrays.append(value.to(device=device, dtype=torch.float32, non_blocking=True)) | |
| return torch.cat(arrays, dim=2) | |
| def get_label(batch: dict[str, Any], label_key: str, device: torch.device) -> torch.Tensor: | |
| value = batch["labels"].get(label_key) | |
| if value is None: | |
| raise ValueError(f"batch label {label_key!r} is None") | |
| value = value.to(device=device, dtype=torch.float32, non_blocking=True) | |
| if value.ndim == 4 and value.shape[1] == 1: | |
| value = value[:, 0] | |
| return value | |
| def atomic_save(payload: dict[str, Any], path: str | Path) -> None: | |
| path = Path(path) | |
| path.parent.mkdir(parents=True, exist_ok=True) | |
| tmp = path.with_suffix(path.suffix + ".tmp") | |
| torch.save(payload, tmp) | |
| os.replace(tmp, path) | |
| def save_epoch_checkpoints( | |
| config: dict[str, Any], | |
| model: torch.nn.Module, | |
| optimizer: torch.optim.Optimizer, | |
| scheduler: Any, | |
| epoch: int, | |
| train_summary: dict[str, Any], | |
| criterion: torch.nn.Module | None = None, | |
| ) -> None: | |
| out_dir = Path(config.get("output_dir", config.get("checkpoint_dir", "runs/default"))) | |
| ckpt_dir = out_dir / "checkpoints" | |
| snapshot = config_snapshot(config) | |
| model_payload = { | |
| "epoch": int(epoch), | |
| "model_state_dict": model.state_dict(), | |
| "config": snapshot, | |
| "train_summary": train_summary, | |
| } | |
| atomic_save(model_payload, ckpt_dir / f"epoch_{int(epoch):04d}_model.pt") | |
| full_payload = dict(model_payload) | |
| full_payload.update( | |
| { | |
| "optimizer_state_dict": optimizer.state_dict(), | |
| "scheduler_state_dict": scheduler.state_dict() if scheduler is not None else None, | |
| } | |
| ) | |
| if criterion is not None: | |
| full_payload["criterion_state_dict"] = criterion.state_dict() | |
| atomic_save(full_payload, ckpt_dir / "latest_full.pt") | |
| def load_model_checkpoint(model: torch.nn.Module, path: str | Path, device: torch.device) -> dict[str, Any]: | |
| path = Path(path) | |
| if path.suffix == ".safetensors": | |
| from safetensors.torch import load_file | |
| state = load_file(str(path), device=str(device)) | |
| model.load_state_dict(state, strict=True) | |
| return {"model_state_dict": state} | |
| payload = torch.load(path, map_location=device, weights_only=True) | |
| state = payload.get("model_state_dict", payload) | |
| model.load_state_dict(state) | |
| return payload if isinstance(payload, dict) else {"model_state_dict": payload} | |
| def resume_full_checkpoint( | |
| path: str | Path, | |
| model: torch.nn.Module, | |
| optimizer: torch.optim.Optimizer, | |
| scheduler: Any, | |
| device: torch.device, | |
| criterion: torch.nn.Module | None = None, | |
| ) -> int: | |
| path = Path(path) | |
| if not path.exists(): | |
| return 0 | |
| payload = torch.load(path, map_location=device, weights_only=True) | |
| model.load_state_dict(payload["model_state_dict"], strict=True) | |
| optimizer.load_state_dict(payload["optimizer_state_dict"]) | |
| for state in optimizer.state.values(): | |
| for key, value in state.items(): | |
| if torch.is_tensor(value): | |
| state[key] = value.to(device) | |
| if scheduler is not None and payload.get("scheduler_state_dict") is not None: | |
| scheduler.load_state_dict(payload["scheduler_state_dict"]) | |
| if criterion is not None and payload.get("criterion_state_dict") is not None: | |
| criterion.load_state_dict(payload["criterion_state_dict"], strict=False) | |
| return int(payload.get("epoch", 0)) | |
| def append_csv_row(path: str | Path, row: dict[str, Any]) -> None: | |
| path = Path(path) | |
| path.parent.mkdir(parents=True, exist_ok=True) | |
| write_header = not path.exists() | |
| with path.open("a", newline="", encoding="utf-8") as f: | |
| writer = csv.DictWriter(f, fieldnames=list(row.keys())) | |
| if write_header: | |
| writer.writeheader() | |
| writer.writerow(row) | |
| def write_json(path: str | Path, payload: dict[str, Any]) -> None: | |
| path = Path(path) | |
| path.parent.mkdir(parents=True, exist_ok=True) | |
| with path.open("w", encoding="utf-8") as f: | |
| json.dump(payload, f, indent=2, ensure_ascii=False) | |
| def maybe_copy_best(src: Path, dst: Path, enabled: bool) -> None: | |
| if enabled: | |
| dst.parent.mkdir(parents=True, exist_ok=True) | |
| # Some mounted filesystems allow writing file contents but reject chmod/copystat. | |
| # copyfile keeps best_model.pt useful without copying metadata. | |
| shutil.copyfile(src, dst) | |