File size: 3,785 Bytes
8bfc737
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
from __future__ import annotations

import os
from pathlib import Path
from typing import Any

import torch
import yaml


ROOT = Path(__file__).resolve().parents[1]
DEFAULT_CONFIG = ROOT / "conf/config.yaml"


def load_config(path: str | Path = DEFAULT_CONFIG) -> dict[str, Any]:
    config_path = Path(path).expanduser().resolve()
    with config_path.open("r", encoding="utf-8") as handle:
        config = yaml.safe_load(handle)
    if not isinstance(config, dict):
        raise ValueError("configuration must be a YAML mapping")
    for section in ("data", "model", "train", "dataloader", "distributed"):
        if section not in config:
            raise ValueError(f"missing configuration section: {section}")
    data, model = config["data"], config["model"]
    for key in ("input_length", "output_length", "height", "width", "channels"):
        if int(data.get(key, 0)) <= 0:
            raise ValueError(f"data.{key} must be positive")
    if int(data["height"]) % 4 or int(data["width"]) % 4:
        raise ValueError("data height and width must be divisible by four")
    dims, depths = model.get("dims"), model.get("depths")
    if not isinstance(dims, list) or len(dims) != 2 or int(dims[1]) != 2 * int(dims[0]):
        raise ValueError("model.dims must be [D, 2*D]")
    if not isinstance(depths, list) or len(depths) != 2 or min(int(x) for x in depths) < 1:
        raise ValueError("model.depths must contain two positive integers")
    heads = int(model.get("heads", 0))
    if heads < 1 or any(int(dim) % heads for dim in dims):
        raise ValueError("model.heads must divide both hidden dimensions")
    normalization = data.get("normalization", "unit")
    if normalization not in ("unit", "uint8_255"):
        raise ValueError("data.normalization must be 'unit' or 'uint8_255'")
    for key in ("data_dir", "train_npz", "val_npz", "test_npz"):
        if data.get(key):
            value = Path(data[key]).expanduser()
            data[key] = str(value if value.is_absolute() else ROOT / value)
    output_dir = Path(config["train"]["output_dir"]).expanduser()
    config["train"]["output_dir"] = str(output_dir if output_dir.is_absolute() else ROOT / output_dir)
    return config


def resolve_device(requested: str, local_rank: int = 0) -> torch.device:
    if requested not in ("auto", "cpu", "cuda"):
        raise ValueError("device must be auto, cpu, or cuda")
    use_accelerator = requested == "cuda" or (requested == "auto" and torch.cuda.is_available())
    if use_accelerator:
        if not torch.cuda.is_available():
            raise RuntimeError("CUDA/ROCm device requested but torch.cuda.is_available() is false")
        torch.cuda.set_device(local_rank)
        return torch.device("cuda", local_rank)
    return torch.device("cpu")


def resolve_cli_path(value: str | None) -> str | None:
    if not value:
        return value
    path = Path(value).expanduser()
    return str(path if path.is_absolute() else ROOT / path)


def load_checkpoint_payload(path: str | Path, device: torch.device) -> dict[str, Any]:
    payload = torch.load(path, map_location=device, weights_only=False)
    if not isinstance(payload, dict) or "model" not in payload or "config" not in payload:
        raise ValueError("checkpoint must contain model and config")
    return payload


def clean_state_dict(state_dict: dict[str, torch.Tensor]) -> dict[str, torch.Tensor]:
    return {key.removeprefix("module."): value for key, value in state_dict.items()}


def atomic_torch_save(payload: dict[str, Any], path: str | Path) -> Path:
    path = Path(path)
    path.parent.mkdir(parents=True, exist_ok=True)
    temporary = path.with_name(f".{path.name}.{os.getpid()}.tmp")
    torch.save(payload, temporary)
    os.replace(temporary, path)
    return path