Spaces:
Running on Zero
Running on Zero
Download src/config.py from guyPerry/scatteringnet: direct link, hf CLI and curl.
- Browser
- Download file 20.7 kB
-
https://huggingface.co/spaces/guyPerry/scatteringnet/resolve/main/src/config.py
- Command line
-
hf download hf://spaces/guyPerry/scatteringnet/src/config.py
-
curl -L -o config.py https://huggingface.co/spaces/guyPerry/scatteringnet/resolve/main/src/config.py
20.7 kB
| """Hybrid config loader for catalog occupancy training. | |
| Static experiment knobs live in ``config.yaml``. ``device`` is resolved | |
| here from CUDA availability. Training knobs (epochs, lr, batch_size, | |
| optimizer, catalog, val split) are YAML-owned so ``train_multi_npz`` | |
| does not hardcode them. | |
| This module does not read ``.env`` and does not open NPZ files. | |
| """ | |
| from __future__ import annotations | |
| import os | |
| from dataclasses import dataclass | |
| from pathlib import Path | |
| from typing import Any, Mapping, TypedDict | |
| import torch | |
| import yaml | |
| # Repo root: src/config.py → parents[1]. | |
| _REPO_ROOT = Path(__file__).resolve().parents[1] | |
| _DEFAULT_YAML = _REPO_ROOT / "config.yaml" | |
| _REQUIRED_YAML_KEYS = ( | |
| "hidden", | |
| "depth", | |
| "seed", | |
| "data_dir", | |
| "epochs", | |
| "lr", | |
| ) | |
| class YamlKnobs(TypedDict): | |
| """Subset of OccupancyConfig that is stored in YAML (paths as strings).""" | |
| hidden: int | |
| depth: int | |
| seed: int | |
| data_dir: str | |
| epochs: int | |
| lr: float | |
| val_fraction: float | |
| latent_dim: int | None | |
| npz_glob: str | |
| npz_paths: tuple[str, ...] | |
| npz_catalog: tuple[tuple[str, int | None], ...] | |
| max_files_per_shape: int | None | |
| run_name: str | |
| checkpoint_metric: str | |
| batch_size: int | |
| optimizer: str | |
| n_surface: int | |
| knn_k: int | |
| knn_local_dim: int | None | |
| shape_encoder: str | |
| # Explicit BCE pos_weight; None when omitted or when auto is set. | |
| pos_weight: float | None | |
| pos_weight_auto: bool | |
| def get_device() -> torch.device: | |
| """ | |
| CUDA when a GPU is visible; otherwise CPU. | |
| Hugging Face CPU Spaces have no CUDA: infer still runs (slower). | |
| ``SCATTERINGNET_DEVICE=cpu`` forces CPU even if a GPU exists. | |
| ``SCATTERINGNET_DEVICE=cuda`` uses CUDA only when ``is_available()``; | |
| otherwise it falls back to CPU (no crash). | |
| Training scripts should still warn on a long catalog run on CPU. | |
| """ | |
| forced = os.environ.get("SCATTERINGNET_DEVICE", "").strip().lower() | |
| if forced == "cpu": | |
| return torch.device("cpu") | |
| if forced == "cuda": | |
| if torch.cuda.is_available(): | |
| return torch.device("cuda") | |
| return torch.device("cpu") | |
| if torch.cuda.is_available(): | |
| return torch.device("cuda") | |
| return torch.device("cpu") | |
| def repo_root() -> Path: | |
| """Git / project root (folder that contains ``src/`` and ``config.yaml``).""" | |
| return _REPO_ROOT | |
| def gpu_name(device: torch.device | None = None) -> str | None: | |
| """ | |
| Human GPU name for the run snapshot (``None`` on CPU). | |
| Uses ``cfg.device`` when given so a forced-CPU train does not stamp a | |
| card that was not used. | |
| """ | |
| dev = device if device is not None else get_device() | |
| if dev.type != "cuda" or not torch.cuda.is_available(): | |
| return None | |
| index = 0 if dev.index is None else int(dev.index) | |
| if index < 0 or index >= torch.cuda.device_count(): | |
| return None | |
| name = str(torch.cuda.get_device_name(index)).strip() | |
| return name or None | |
| def _as_positive_int(name: str, value: Any) -> int: | |
| """YAML may yield int or (rarely) str; occupancy dims must be int >= 1.""" | |
| try: | |
| parsed = int(value) | |
| except (TypeError, ValueError) as exc: | |
| raise ValueError(f"{name} must be an integer, got {value!r}") from exc | |
| if parsed < 1: | |
| raise ValueError(f"{name} must be >= 1, got {parsed}") | |
| return parsed | |
| def _as_int_in_range(name: str, value: Any, lo: int, hi: int) -> int: | |
| """Inclusive integer range (YAML ``knn_k`` is 0–4096).""" | |
| try: | |
| parsed = int(value) | |
| except (TypeError, ValueError) as exc: | |
| raise ValueError(f"{name} must be an integer, got {value!r}") from exc | |
| if parsed < lo or parsed > hi: | |
| raise ValueError(f"{name} must be in [{lo}, {hi}], got {parsed}") | |
| return parsed | |
| def _as_positive_float(name: str, value: Any) -> float: | |
| """Learning-rate style knobs must be a finite float > 0.""" | |
| try: | |
| parsed = float(value) | |
| except (TypeError, ValueError) as exc: | |
| raise ValueError(f"{name} must be a float, got {value!r}") from exc | |
| if parsed <= 0.0 or parsed != parsed: | |
| raise ValueError(f"{name} must be > 0, got {parsed}") | |
| return parsed | |
| def _as_pos_weight_pair(raw: Mapping[str, Any]) -> tuple[float | None, bool]: | |
| """ | |
| YAML ``pos_weight``: omit / null → unweighted BCE. | |
| ``auto`` → compute n_outside / n_inside on the train split at train time. | |
| A finite float > 0 is used as-is (1.0 is unweighted). | |
| """ | |
| if "pos_weight" not in raw: | |
| return None, False | |
| value = raw["pos_weight"] | |
| if value is None or value is False: | |
| return None, False | |
| if isinstance(value, str): | |
| text = value.strip().lower() | |
| if text in ("", "none", "off", "false"): | |
| return None, False | |
| if text == "auto": | |
| return None, True | |
| parsed = _as_positive_float("pos_weight", value) | |
| return parsed, False | |
| def _as_open_unit_interval(name: str, value: Any) -> float: | |
| """Hold-out fractions must be in (0, 1) so both splits are non-empty.""" | |
| try: | |
| parsed = float(value) | |
| except (TypeError, ValueError) as exc: | |
| raise ValueError(f"{name} must be a float, got {value!r}") from exc | |
| if parsed != parsed or parsed <= 0.0 or parsed >= 1.0: | |
| raise ValueError(f"{name} must be in (0, 1), got {parsed}") | |
| return parsed | |
| def _as_nonempty_path_string(name: str, value: Any) -> str: | |
| if value is None or (isinstance(value, str) and not value.strip()): | |
| raise ValueError(f"{name} must be a non-empty path string in config.yaml") | |
| return str(value).strip() | |
| def _as_data_dir_string(value: Any) -> str: | |
| return _as_nonempty_path_string("data_dir", value) | |
| def _as_optional_positive_int(name: str, value: Any) -> int | None: | |
| """YAML null → unlimited catalog cap; otherwise int >= 1.""" | |
| if value is None: | |
| return None | |
| return _as_positive_int(name, value) | |
| def _as_run_name(value: Any) -> str: | |
| """Optional YAML suffix for ``runs/<timestamp>_<name>/``; empty → ``run``.""" | |
| if value is None: | |
| return "run" | |
| text = str(value).strip() | |
| return text if text else "run" | |
| _CHECKPOINT_METRIC_ALIASES = { | |
| "test_acc": "val_acc", | |
| "test_iou": "val_iou", | |
| } | |
| def _as_checkpoint_metric(value: Any) -> str: | |
| """Name of the scalar used to decide ``best.pt`` (strict improve).""" | |
| if value is None or (isinstance(value, str) and not value.strip()): | |
| raise ValueError("checkpoint_metric must be a non-empty string") | |
| name = str(value).strip() | |
| return _CHECKPOINT_METRIC_ALIASES.get(name, name) | |
| def _as_val_fraction(raw: Mapping[str, Any]) -> float: | |
| """Prefer ``val_fraction``; accept legacy ``test_fraction``.""" | |
| if "val_fraction" in raw: | |
| return _as_open_unit_interval("val_fraction", raw["val_fraction"]) | |
| if "test_fraction" in raw: | |
| return _as_open_unit_interval("test_fraction", raw["test_fraction"]) | |
| raise ValueError("Config YAML missing keys: val_fraction") | |
| def _as_optional_latent_dim(raw: Mapping[str, Any]) -> int | None: | |
| """YAML omit / null → use ``hidden`` at train time.""" | |
| if "latent_dim" not in raw or raw["latent_dim"] is None: | |
| return None | |
| return _as_positive_int("latent_dim", raw["latent_dim"]) | |
| # Names accepted in config.yaml ``optimizer``. Used by train_multi_npz. | |
| _ALLOWED_OPTIMIZERS = ("adam", "adamw", "sgd") | |
| # ``none`` = OccupancyMLP (xyz only). ``surface`` = envelope encoder. | |
| _ALLOWED_SHAPE_ENCODERS = ("none", "surface") | |
| def _as_optimizer(value: Any) -> str: | |
| """Optimizer family for multi-NPZ train; default Adam.""" | |
| if value is None or (isinstance(value, str) and not str(value).strip()): | |
| return "adam" | |
| name = str(value).strip().lower() | |
| if name not in _ALLOWED_OPTIMIZERS: | |
| allowed = ", ".join(_ALLOWED_OPTIMIZERS) | |
| raise ValueError(f"optimizer must be one of {allowed}, got {value!r}") | |
| return name | |
| def _as_shape_encoder(value: Any) -> str: | |
| """Occupancy head family; default xyz-only so older YAML still loads.""" | |
| if value is None or (isinstance(value, str) and not str(value).strip()): | |
| return "none" | |
| name = str(value).strip().lower() | |
| if name not in _ALLOWED_SHAPE_ENCODERS: | |
| allowed = ", ".join(_ALLOWED_SHAPE_ENCODERS) | |
| raise ValueError(f"shape_encoder must be one of {allowed}, got {value!r}") | |
| return name | |
| def _as_npz_paths(value: Any) -> tuple[str, ...]: | |
| """Explicit NPZ list relative to data_dir (empty → use glob).""" | |
| if value is None: | |
| return () | |
| if isinstance(value, str): | |
| item = value.strip() | |
| return (item,) if item else () | |
| if not isinstance(value, list): | |
| raise ValueError(f"npz_paths must be a list of strings or null, got {type(value).__name__}") | |
| out: list[str] = [] | |
| for i, raw in enumerate(value): | |
| text = _as_nonempty_path_string(f"npz_paths[{i}]", raw) | |
| out.append(text) | |
| return tuple(out) | |
| def _as_npz_catalog(value: Any) -> tuple[tuple[str, int | None], ...]: | |
| """Union of globs; optional ``max_shapes`` is unique meshes per glob.""" | |
| if value is None: | |
| return () | |
| if not isinstance(value, list): | |
| raise ValueError( | |
| f"npz_catalog must be a list or null, got {type(value).__name__}" | |
| ) | |
| out: list[tuple[str, int | None]] = [] | |
| for i, raw in enumerate(value): | |
| if isinstance(raw, str): | |
| glob_s = _as_nonempty_path_string(f"npz_catalog[{i}]", raw) | |
| out.append((glob_s, None)) | |
| continue | |
| if not isinstance(raw, Mapping): | |
| raise ValueError( | |
| f"npz_catalog[{i}] must be a glob string or mapping, " | |
| f"got {type(raw).__name__}" | |
| ) | |
| if "glob" not in raw: | |
| raise ValueError(f"npz_catalog[{i}] missing glob") | |
| glob_s = _as_nonempty_path_string(f"npz_catalog[{i}].glob", raw["glob"]) | |
| max_shapes: int | None = None | |
| if "max_shapes" in raw and raw["max_shapes"] is not None: | |
| max_shapes = _as_positive_int( | |
| f"npz_catalog[{i}].max_shapes", raw["max_shapes"] | |
| ) | |
| out.append((glob_s, max_shapes)) | |
| return tuple(out) | |
| def as_repo_relative(path: Path | str, *, root: Path | None = None) -> str: | |
| """ | |
| POSIX string relative to the git repo when ``path`` is inside it. | |
| Already-relative inputs are returned as POSIX. Absolute paths on another | |
| drive (the dataset disk) cannot be repo-relative and stay absolute POSIX. | |
| """ | |
| text = str(path).strip() | |
| if not text: | |
| return text | |
| parsed = Path(text) | |
| if not parsed.is_absolute(): | |
| return parsed.as_posix() | |
| base = (root or _REPO_ROOT).resolve() | |
| try: | |
| return parsed.resolve().relative_to(base).as_posix() | |
| except ValueError: | |
| return parsed.resolve().as_posix() | |
| def as_data_relative(path: Path | str, data_dir: Path | str) -> str: | |
| """ | |
| POSIX string relative to ``data_dir`` (``exports/...``, not ``E:/...``). | |
| Already-relative inputs are returned as POSIX. Paths outside ``data_dir`` | |
| (unit-test temp trees) fall back to absolute POSIX. | |
| """ | |
| text = str(path).strip() | |
| if not text: | |
| return text | |
| parsed = Path(text) | |
| if not parsed.is_absolute(): | |
| return parsed.as_posix() | |
| root = Path(data_dir).expanduser().resolve() | |
| resolved = parsed.expanduser().resolve() | |
| try: | |
| return resolved.relative_to(root).as_posix() | |
| except ValueError: | |
| return resolved.as_posix() | |
| def load_yaml_knobs(path: Path) -> YamlKnobs: | |
| """ | |
| Read YAML settings. Does not check that data_dir exists on disk. | |
| Parameters | |
| ---------- | |
| path: | |
| Path to ``config.yaml``. | |
| Returns | |
| ------- | |
| YamlKnobs | |
| Typed dict of experiment knobs (paths still strings). | |
| """ | |
| if not path.is_file(): | |
| raise FileNotFoundError(f"Config YAML not found: {path}") | |
| raw = yaml.safe_load(path.read_text(encoding="utf-8")) | |
| if not isinstance(raw, Mapping): | |
| raise ValueError(f"Config YAML must be a mapping, got {type(raw).__name__}") | |
| missing = [k for k in _REQUIRED_YAML_KEYS if k not in raw] | |
| if missing: | |
| raise ValueError(f"Config YAML missing keys: {', '.join(missing)}") | |
| pos_weight, pos_weight_auto = _as_pos_weight_pair(raw) | |
| return { | |
| "hidden": _as_positive_int("hidden", raw["hidden"]), | |
| "depth": _as_positive_int("depth", raw["depth"]), | |
| "seed": _as_positive_int("seed", raw["seed"]), | |
| "data_dir": _as_data_dir_string(raw["data_dir"]), | |
| "epochs": _as_positive_int("epochs", raw["epochs"]), | |
| "lr": _as_positive_float("lr", raw["lr"]), | |
| "val_fraction": _as_val_fraction(raw), | |
| "latent_dim": _as_optional_latent_dim(raw), | |
| "npz_glob": ( | |
| _as_nonempty_path_string("npz_glob", raw["npz_glob"]) | |
| if "npz_glob" in raw | |
| else "exports/dataset/*.npz" | |
| ), | |
| "npz_paths": _as_npz_paths(raw.get("npz_paths")), | |
| "npz_catalog": _as_npz_catalog(raw.get("npz_catalog")), | |
| "max_files_per_shape": ( | |
| _as_optional_positive_int("max_files_per_shape", raw["max_files_per_shape"]) | |
| if "max_files_per_shape" in raw | |
| else 2 | |
| ), | |
| "run_name": ( | |
| _as_run_name(raw["run_name"]) if "run_name" in raw else "run" | |
| ), | |
| "checkpoint_metric": ( | |
| _as_checkpoint_metric(raw["checkpoint_metric"]) | |
| if "checkpoint_metric" in raw | |
| else "val_acc" | |
| ), | |
| "batch_size": ( | |
| _as_positive_int("batch_size", raw["batch_size"]) | |
| if "batch_size" in raw | |
| else 1024 | |
| ), | |
| "optimizer": ( | |
| _as_optimizer(raw["optimizer"]) if "optimizer" in raw else "adam" | |
| ), | |
| "n_surface": ( | |
| _as_positive_int("n_surface", raw["n_surface"]) | |
| if "n_surface" in raw | |
| else 1024 | |
| ), | |
| "knn_k": ( | |
| _as_int_in_range("knn_k", raw["knn_k"], 0, 4096) | |
| if "knn_k" in raw | |
| else 0 | |
| ), | |
| "knn_local_dim": ( | |
| _as_positive_int("knn_local_dim", raw["knn_local_dim"]) | |
| if "knn_local_dim" in raw | |
| else None | |
| ), | |
| "shape_encoder": ( | |
| _as_shape_encoder(raw["shape_encoder"]) | |
| if "shape_encoder" in raw | |
| else "none" | |
| ), | |
| "pos_weight": pos_weight, | |
| "pos_weight_auto": pos_weight_auto, | |
| } | |
| def _warn_missing_data_dir(data_dir: Path, yaml_path: Path) -> None: | |
| """Print a terminal hint so the user can fix config.yaml (no .env involved).""" | |
| print( | |
| "\n" | |
| "Dataset folder not found.\n" | |
| f" Looked for: {data_dir}\n" | |
| "\n" | |
| "Update `data_dir` in config.yaml to the folder that contains " | |
| "`exports/` and `meshes/`.\n" | |
| f" Config file: {yaml_path}\n" | |
| ) | |
| def require_data_dir(data_dir: Path, *, yaml_path: Path) -> None: | |
| """Validate the dataset root before training or NPZ loading.""" | |
| if data_dir.is_dir(): | |
| return | |
| _warn_missing_data_dir(data_dir, yaml_path) | |
| raise FileNotFoundError( | |
| f"Dataset directory does not exist: {data_dir}. " | |
| f"Set data_dir in {yaml_path}." | |
| ) | |
| class OccupancyConfig: | |
| """Resolved experiment settings from YAML plus detected device.""" | |
| data_dir: Path | |
| device: torch.device | |
| hidden: int | |
| depth: int | |
| seed: int | |
| epochs: int | |
| lr: float | |
| # Fraction of catalog **meshes** held out as val (selection split, not a locked test). | |
| val_fraction: float | |
| # Catalog knobs (optional in YAML; omitted keys keep these defaults). | |
| npz_glob: str = "exports/dataset/*.npz" | |
| npz_paths: tuple[str, ...] = () | |
| # ``(glob, max_shapes)`` rows. Empty → use ``npz_glob``. ``max_shapes`` | |
| # None keeps every mesh that glob hits (after ``max_files_per_shape``). | |
| npz_catalog: tuple[tuple[str, int | None], ...] = () | |
| max_files_per_shape: int | None = 2 | |
| # Suffix for runs/<timestamp>_<name>/ (device stays runtime-only). | |
| run_name: str = "run" | |
| # Which logged scalar selects best.pt (strict improve). | |
| checkpoint_metric: str = "val_acc" | |
| # Encoder latent width; None → use ``hidden`` at train / infer time. | |
| latent_dim: int | None = None | |
| # Mini-batch size and optimizer family (train_multi_npz). | |
| batch_size: int = 1024 | |
| optimizer: str = "adam" | |
| # Envelope sample count (YAML). Used when ``shape_encoder`` is ``surface``. | |
| n_surface: int = 1024 | |
| # 0 = global envelope z only. >0 = that many nearest envelope dots per query. | |
| knn_k: int = 0 | |
| # Width of z_local; None → same as occupancy latent_dim / hidden. | |
| knn_local_dim: int | None = None | |
| # ``none`` keeps OccupancyMLP; ``surface`` uses the envelope PointNet. | |
| shape_encoder: str = "none" | |
| # BCE inside-class weight. None + auto=False = unweighted (legacy). | |
| pos_weight: float | None = None | |
| pos_weight_auto: bool = False | |
| def load_config( | |
| yaml_path: Path | None = None, | |
| *, | |
| require_existing_data_dir: bool = True, | |
| ) -> OccupancyConfig: | |
| """ | |
| Compose OccupancyConfig from ``config.yaml`` (not from ``.env``). | |
| When ``require_existing_data_dir`` is True (default), a missing folder | |
| prints a short instruction and then raises FileNotFoundError — used for | |
| training and data loading. | |
| Parameters | |
| ---------- | |
| yaml_path: | |
| Config file; default is repo-root ``config.yaml``. | |
| require_existing_data_dir: | |
| If True, refuse to return a config whose ``data_dir`` is missing. | |
| Returns | |
| ------- | |
| OccupancyConfig | |
| YAML knobs plus detected ``device``. | |
| """ | |
| cfg_path = yaml_path or _DEFAULT_YAML | |
| knobs = load_yaml_knobs(cfg_path) | |
| data_dir = Path(knobs["data_dir"]) | |
| if require_existing_data_dir: | |
| require_data_dir(data_dir, yaml_path=cfg_path) | |
| return OccupancyConfig( | |
| data_dir=data_dir, | |
| device=get_device(), | |
| hidden=knobs["hidden"], | |
| depth=knobs["depth"], | |
| seed=knobs["seed"], | |
| epochs=knobs["epochs"], | |
| lr=knobs["lr"], | |
| val_fraction=knobs["val_fraction"], | |
| latent_dim=knobs["latent_dim"], | |
| npz_glob=knobs["npz_glob"], | |
| npz_paths=knobs["npz_paths"], | |
| npz_catalog=knobs["npz_catalog"], | |
| max_files_per_shape=knobs["max_files_per_shape"], | |
| run_name=knobs["run_name"], | |
| checkpoint_metric=knobs["checkpoint_metric"], | |
| batch_size=knobs["batch_size"], | |
| optimizer=knobs["optimizer"], | |
| n_surface=knobs["n_surface"], | |
| knn_k=knobs["knn_k"], | |
| knn_local_dim=knobs["knn_local_dim"], | |
| shape_encoder=knobs["shape_encoder"], | |
| pos_weight=knobs["pos_weight"], | |
| pos_weight_auto=knobs["pos_weight_auto"], | |
| ) | |
| def encoder_latent_dim(cfg: OccupancyConfig) -> int: | |
| """OccupancyEncoder ``z`` width: YAML ``latent_dim`` or ``hidden``.""" | |
| if cfg.latent_dim is None: | |
| return int(cfg.hidden) | |
| return int(cfg.latent_dim) | |
| def encoder_knn_local_dim(cfg: OccupancyConfig) -> int: | |
| """Local envelope code width: YAML ``knn_local_dim`` or global latent.""" | |
| if cfg.knn_local_dim is None: | |
| return encoder_latent_dim(cfg) | |
| return int(cfg.knn_local_dim) | |
| def format_config(cfg: OccupancyConfig) -> str: | |
| """ | |
| Pretty-print for CLI smoke checks. | |
| Parameters | |
| ---------- | |
| cfg: | |
| Resolved config. | |
| Returns | |
| ------- | |
| str | |
| Multi-line ``OccupancyConfig(...)`` dump. | |
| """ | |
| return ( | |
| f"OccupancyConfig(\n" | |
| f" data_dir={cfg.data_dir}\n" | |
| f" device={cfg.device}\n" | |
| f" gpu={gpu_name(cfg.device)}\n" | |
| f" hidden={cfg.hidden}\n" | |
| f" depth={cfg.depth}\n" | |
| f" seed={cfg.seed}\n" | |
| f" epochs={cfg.epochs}\n" | |
| f" lr={cfg.lr}\n" | |
| f" val_fraction={cfg.val_fraction}\n" | |
| f" latent_dim={cfg.latent_dim}\n" | |
| f" npz_glob={cfg.npz_glob}\n" | |
| f" npz_paths={list(cfg.npz_paths)}\n" | |
| f" npz_catalog={list(cfg.npz_catalog)}\n" | |
| f" max_files_per_shape={cfg.max_files_per_shape}\n" | |
| f" run_name={cfg.run_name}\n" | |
| f" checkpoint_metric={cfg.checkpoint_metric}\n" | |
| f" batch_size={cfg.batch_size}\n" | |
| f" optimizer={cfg.optimizer}\n" | |
| f" n_surface={cfg.n_surface}\n" | |
| f" knn_k={cfg.knn_k}\n" | |
| f" knn_local_dim={cfg.knn_local_dim}\n" | |
| f" shape_encoder={cfg.shape_encoder}\n" | |
| f" pos_weight={cfg.pos_weight}\n" | |
| f" pos_weight_auto={cfg.pos_weight_auto}\n" | |
| f")" | |
| ) | |
| if __name__ == "__main__": | |
| print(format_config(load_config())) | |