File size: 20,730 Bytes
bc4c433
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
416
417
418
419
420
421
422
423
424
425
426
427
428
429
430
431
432
433
434
435
436
437
438
439
440
441
442
443
444
445
446
447
448
449
450
451
452
453
454
455
456
457
458
459
460
461
462
463
464
465
466
467
468
469
470
471
472
473
474
475
476
477
478
479
480
481
482
483
484
485
486
487
488
489
490
491
492
493
494
495
496
497
498
499
500
501
502
503
504
505
506
507
508
509
510
511
512
513
514
515
516
517
518
519
520
521
522
523
524
525
526
527
528
529
530
531
532
533
534
535
536
537
538
539
540
541
542
543
544
545
546
547
548
549
550
551
552
553
554
555
556
557
558
559
560
561
562
563
564
565
566
567
568
569
570
571
572
573
574
575
576
577
578
579
580
581
582
583
584
585
586
587
588
589
590
591
592
593
594
595
596
597
598
599
600
601
602
603
604
605
606
607
608
609
610
611
612
"""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}."
    )


@dataclass(frozen=True)
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()))