diff --git a/giant/cli.py b/giant/cli.py index 0cd80c6..4c4fd88 100644 --- a/giant/cli.py +++ b/giant/cli.py @@ -41,6 +41,7 @@ from giant.data.transforms import ( ) from giant.checkpoint_io import CheckpointCompatibilityError, load_for_inference from giant.geometry import GeometryOracle +from giant.materials import MATERIAL_PROPERTIES from giant.pipeline import run_train_job from giant.rollout import ( L1DistCollector, @@ -854,6 +855,52 @@ def new_run( typer.echo(f" giant train {data_arg} --config {config_path} --out {run_dir}") +model_app = typer.Typer( + no_args_is_help=True, + help="Inspect a resolved model architecture without training.", +) +app.add_typer(model_app, name="model") + + +@model_app.command("summary") +def model_summary( + config: Annotated[ + Optional[Path], + typer.Option("--config", "-c", help="TOML config file (default: built-in defaults)"), + ] = None, + pdg_vocab: Annotated[ + int, + typer.Option( + "--pdg-vocab", + help="Placeholder PDG vocab size for conditioning.particle.type='embedding' " + "or a pdg/process router (no dataset attached to derive the real training vocab)", + ), + ] = 300, + mat_vocab: Annotated[ + int, + typer.Option( + "--mat-vocab", + help="Placeholder material vocab size for conditioning.material.type='embedding' " + "or a process router (default: the number of known materials in giant.materials)", + ), + ] = len(MATERIAL_PROPERTIES), +) -> None: + """Build the resolved model graph from a config with no dataset attached, and + print per-module parameter counts, trunk widths, which heads exist, and which + conditioning/stage1_model/stage2_model config keys actually shaped it.""" + from giant.model.summary import render_summary, summarize_model + + cfg = gconfig.merge_cli_overrides(gconfig.DEFAULT_CONFIG, config, {}) + try: + gconfig.validate_config(cfg) + except ValueError as exc: + typer.echo(f"error: {exc}", err=True) + raise typer.Exit(1) + + summary = summarize_model(cfg, pdg_vocab=pdg_vocab, mat_vocab=mat_vocab) + typer.echo(render_summary(summary)) + + @app.command() def predict( data: Annotated[Path, typer.Argument(help="Parquet file or directory of parquet files")], diff --git a/giant/config.py b/giant/config.py index 2c83541..3873e86 100644 --- a/giant/config.py +++ b/giant/config.py @@ -843,6 +843,25 @@ class GiantConfig: DEFAULT_CONFIG: dict = GiantConfig().to_dict() +def leaf_paths(node: dict, prefix: str = "") -> list[str]: + """Every dotted leaf path in a DEFAULT_CONFIG-shaped dict, e.g. + "stage1_model.router.n_experts". `[meta]` (run provenance, no schema + counterpart) is skipped at the top level, matching `validate_config_keys`. + Shared by `tests/test_config_consumed_keys.py` (the static per-identifier + audit) and `giant.model.summary` (the runtime per-config audit, gitea + #46) so both walk the exact same tree.""" + paths = [] + for key, value in node.items(): + if prefix == "" and key == "meta": + continue + path = f"{prefix}.{key}" if prefix else key + if isinstance(value, dict): + paths.extend(leaf_paths(value, path)) + else: + paths.append(path) + return paths + + def git_hash() -> str: try: return subprocess.check_output(["git", "rev-parse", "HEAD"], stderr=subprocess.DEVNULL).decode().strip() diff --git a/giant/model/summary.py b/giant/model/summary.py new file mode 100644 index 0000000..ff8fdc6 --- /dev/null +++ b/giant/model/summary.py @@ -0,0 +1,310 @@ +"""Build-only model introspection (gitea #46): construct the resolved +Stage1/Stage2/critic graph from a config with no dataset attached, and report +per-module parameter counts, trunk widths, which heads exist, and — via +differential probing — which `conditioning`/`stage1_model`/`stage2_model` +config keys actually shape the built model. This is the runtime counterpart +to `tests/test_config_consumed_keys.py`'s static per-identifier audit: that +test asks "does any code reference this key's name at all", this module asks +"given *this* resolved config, does the key change what `build_models`/ +`build_critics` (`giant/model/builders.py`) actually produces". + +Differential probing, not identifier matching: build the model once from the +resolved config and take a structural fingerprint (`_fingerprint` — which +submodules exist, every parameter's/buffer's shape+dtype, every plain scalar +attribute stored on any module). Then, for each in-scope leaf key, perturb +just that one value (`_perturb`), rebuild, and re-fingerprint. A changed +fingerprint — or a rebuild that raises — means the key was consumed; an +identical fingerprint means construction never looked at it under this +particular config. A key can be genuinely inert under one config and live +under another (e.g. any `stage1_model.router.*` key when `router.enabled = +false`) — that config-dependence is exactly the "silently degenerate +combination" issue #46 is after, so it is reported per-run rather than +baked into a static table. + +Keys legitimately owned by the trainer/sampler/rollout rather than by +`build_models`/`build_critics` (loss weights, WGAN-GP training +hyperparameters, teacher-forcing schedules, ...) are cataloged in +`_NOT_BUILD_TIME` below so the report doesn't flag them as suspicious. A +couple of leaves are inert under every config today — +`stage2_model.autoregressive.order`, `stage2_model.stage1_context` — matching +`tests/test_config_consumed_keys.py`'s own `_KNOWN_UNUSED` entries; they are +deliberately *not* in `_NOT_BUILD_TIME`, since "always inert" is itself the +finding those two tests independently converge on. +""" + +import copy +from dataclasses import dataclass, field + +import torch.nn as nn + +from giant.config import _get_path, _set_path, leaf_paths +from giant.model.builders import build_critics, build_models +from giant.model.trunks import RoutedTrunk + +_IN_SCOPE_ROOTS = ("conditioning", "stage1_model", "stage2_model") + +_PROBE_STR = "__giant_model_summary_probe__" + +# A handful of string leaves branch on equality against one specific literal +# (e.g. `builders.py`: `stop_token = s2_spec.n_sec.mode == "stop_token"`), +# where every value other than that literal behaves identically. A single +# generic sentinel probe would then falsely read as inert whenever the +# config's *current* value is already one of those identically-behaving +# "other" values (e.g. mode="head") — it never crosses the one boundary that +# actually matters. Named here so probing tries the real alternative(s) too; +# every other string leaf is registry-validated (raises on garbage, still +# correctly detected as consumed) or genuinely value-independent, so doesn't +# need an entry. +_STRING_ALTERNATIVES: dict[str, tuple[str, ...]] = { + "stage2_model.n_sec.owner": ("stage1", "stage2"), + "stage2_model.n_sec.mode": ("stop_token", "head", "truth"), + "stage2_model.particle_type.target": ("physical", "onehot", "embedding"), +} + +# Verified by reading giant/training/trainers.py, giant/training/stage2_inputs.py +# and giant/rollout.py while implementing gitea #46 — not auto-derived, so a +# future reader touching these fields should re-check this table still holds. +_NOT_BUILD_TIME: dict[str, str] = { + "stage1_model.lambda": "trainers.py: StageSpec.lambda_weight, the total-loss mix weight", + "stage2_model.lambda": "trainers.py: StageSpec.lambda_weight, the total-loss mix weight", + "stage2_model.n_sec.lambda": "trainers.py: StageSpec.n_sec_lambda, the n_sec-head loss weight", + "stage2_model.particle_type.lambda": "trainers.py: Stage2Trainer.particle_type_lambda, the type-head loss weight", + "stage2_model.particle_type.other_policy": "giant/rollout.py: resolves an 'other'-bucket secondary's PDG code at inference", + "stage2_model.autoregressive.teacher_forcing": "giant/training/stage2_inputs.py's training-time input assembly", + "stage2_model.autoregressive.tf_p_start": "trainers.py's teacher-forcing schedule", + "stage2_model.autoregressive.tf_p_end": "trainers.py's teacher-forcing schedule", + "stage1_model.router.lambda_balance": "trainers.py's load-balancing auxiliary loss weight", + "stage1_model.router.lambda_entropy": "trainers.py's entropy-regularization auxiliary loss weight", + "stage1_model.router.lambda_proc": "trainers.py's supervised process-classification auxiliary loss weight", + "stage1_model.router.gumbel_tau_start": "trainers.py's expert-combination Gumbel-softmax temperature anneal", + "stage1_model.router.gumbel_tau_end": "trainers.py's expert-combination Gumbel-softmax temperature anneal", + "stage2_model.router.lambda_balance": "trainers.py's load-balancing auxiliary loss weight", + "stage2_model.router.lambda_entropy": "trainers.py's entropy-regularization auxiliary loss weight", + "stage2_model.router.lambda_proc": "trainers.py's supervised process-classification auxiliary loss weight", + "stage2_model.router.gumbel_tau_start": "trainers.py's expert-combination Gumbel-softmax temperature anneal", + "stage2_model.router.gumbel_tau_end": "trainers.py's expert-combination Gumbel-softmax temperature anneal", + "stage1_model.wgan.n_critic": "trainers.py's WGAN-GP critic-update cadence", + "stage1_model.wgan.gp_weight": "trainers.py's WGAN-GP gradient-penalty coefficient", + "stage1_model.wgan.critic_lr": "trainers.py's critic optimizer learning rate", + "stage2_model.wgan.n_critic": "trainers.py's WGAN-GP critic-update cadence", + "stage2_model.wgan.gp_weight": "trainers.py's WGAN-GP gradient-penalty coefficient", + "stage2_model.wgan.critic_lr": "trainers.py's critic optimizer learning rate", + "stage2_model.wgan.gumbel_tau_start": "trainers.py's type-slice Gumbel-softmax temperature anneal (type_gumbel_tau_start)", + "stage2_model.wgan.gumbel_tau_end": "trainers.py's type-slice Gumbel-softmax temperature anneal (type_gumbel_tau_end)", +} + + +@dataclass +class ModelSummary: + modules: dict[str, nn.Module] + consumed: list[str] + inert: list[str] + elsewhere: list[str] + pdg_vocab: int + mat_vocab: int + vocab_caveats: list[str] = field(default_factory=list) + + +def _build_model_config(cfg: dict, pdg_vocab: int, mat_vocab: int) -> dict: + return { + "pdg_vocab": pdg_vocab, + "mat_vocab": mat_vocab, + "conditioning": cfg["conditioning"], + "stage1_model": cfg["stage1_model"], + "stage2_model": cfg["stage2_model"], + } + + +def _built_modules(cfg: dict, pdg_vocab: int, mat_vocab: int) -> dict[str, nn.Module]: + model_config = _build_model_config(cfg, pdg_vocab, mat_vocab) + modules: dict[str, nn.Module] = {} + for name, m in build_models(model_config).items(): + if m is not None: + modules[name] = m + for name, m in build_critics(model_config).items(): + if m is not None: + modules[f"{name}_critic"] = m + return modules + + +def _fingerprint(modules: dict[str, nn.Module]) -> list: + """A config-shape fingerprint of the built graph: which submodules + exist, every parameter's/buffer's shape+dtype (never values — those are + randomly initialized and irrelevant to *structure*), and every plain + scalar attribute any module stores on itself (e.g. `Stage2Autoregressive + .stop_sampling`, `EnergyRouter.temperature`) — this is what makes a + non-parametric key's effect on construction observable.""" + sig = [] + for stage_name, module in modules.items(): + for mod_name, m in module.named_modules(): + full = f"{stage_name}.{mod_name}" if mod_name else stage_name + for k, v in vars(m).items(): + if k.startswith("_"): + continue + if v is None or isinstance(v, (bool, int, float, str)): + sig.append((full, k, v)) + for pname, p in module.named_parameters(): + sig.append((stage_name, "param", pname, tuple(p.shape), str(p.dtype))) + for bname, b in module.named_buffers(): + sig.append((stage_name, "buffer", bname, tuple(b.shape), str(b.dtype))) + return sorted(sig, key=repr) + + +def _perturb_candidates(path: str, value) -> list: + """Values to try perturbing `path`'s current `value` to, in order — + probing stops at the first one that changes the fingerprint or raises. + Almost always a single candidate; see `_STRING_ALTERNATIVES`.""" + if isinstance(value, bool): + return [not value] + if isinstance(value, int): + return [value + 1] + if isinstance(value, float): + return [value + 1.0] + if isinstance(value, str): + alternatives = [v for v in _STRING_ALTERNATIVES.get(path, ()) if v != value] + return [*alternatives, _PROBE_STR] + raise TypeError(f"gitea #46 probing: unsupported leaf value type {type(value)!r} ({value!r})") + + +def _vocab_caveats(cfg: dict) -> list[str]: + caveats = [] + if _get_path(cfg, "conditioning.particle.type") == "embedding": + caveats.append( + "conditioning.particle.type = 'embedding' -- pdg_vocab below is a " + "placeholder (no dataset attached to derive the real training vocab size)" + ) + if _get_path(cfg, "conditioning.material.type") == "embedding": + caveats.append( + "conditioning.material.type = 'embedding' -- mat_vocab below is a " + "placeholder (no dataset attached to derive the real training vocab size)" + ) + for stage in ("stage1_model", "stage2_model"): + router_type = _get_path(cfg, f"{stage}.router.type") + if _get_path(cfg, f"{stage}.router.enabled") and router_type in ("pdg", "process"): + caveats.append( + f"{stage}.router.type = {router_type!r} builds its own pdg_vocab-sized " + "embedding -- the count above is a placeholder" + ) + return caveats + + +def summarize_model(cfg: dict, pdg_vocab: int, mat_vocab: int) -> ModelSummary: + """Build `cfg`'s model with no dataset attached and report its resolved + graph, plus which `conditioning`/`stage1_model`/`stage2_model` config + keys actually shaped it (differential probing — see module docstring). + `cfg` must already be a fully-merged v0.3 config (`merge_cli_overrides` + output) — this does not migrate or validate it.""" + modules = _built_modules(cfg, pdg_vocab, mat_vocab) + baseline_fp = _fingerprint(modules) + + in_scope = [p for p in leaf_paths(cfg) if p.split(".", 1)[0] in _IN_SCOPE_ROOTS] + consumed: list[str] = [] + inert: list[str] = [] + elsewhere: list[str] = [] + for path in in_scope: + original = _get_path(cfg, path) + changed = False + for candidate in _perturb_candidates(path, original): + probe_cfg = copy.deepcopy( + { + "conditioning": cfg["conditioning"], + "stage1_model": cfg["stage1_model"], + "stage2_model": cfg["stage2_model"], + } + ) + _set_path(probe_cfg, path, candidate) + try: + changed = _fingerprint(_built_modules(probe_cfg, pdg_vocab, mat_vocab)) != baseline_fp + except Exception: + changed = True + if changed: + break + if changed: + consumed.append(path) + elif path in _NOT_BUILD_TIME: + elsewhere.append(path) + else: + inert.append(path) + + return ModelSummary( + modules=modules, + consumed=sorted(consumed), + inert=sorted(inert), + elsewhere=sorted(elsewhere), + pdg_vocab=pdg_vocab, + mat_vocab=mat_vocab, + vocab_caveats=_vocab_caveats(cfg), + ) + + +def _tree_lines(module: nn.Module, name: str, indent: int = 0) -> list[str]: + total = sum(p.numel() for p in module.parameters()) + in_dim = getattr(module, "in_dim", None) + out_dim = getattr(module, "out_dim", None) + widths = f" [in={in_dim}, out={out_dim}]" if in_dim is not None and out_dim is not None else "" + lines = [f"{' ' * indent}{name} ({type(module).__name__}): {total:,}{widths}"] + for child_name, child in module.named_children(): + lines.extend(_tree_lines(child, child_name, indent + 1)) + return lines + + +_HEAD_NAMES = ("n_sec_head", "type_head", "stop_head") + + +def _stage_header(name: str, module: nn.Module) -> list[str]: + total = sum(p.numel() for p in module.parameters()) + lines = [f"{name}: {type(module).__name__} -- {total:,} parameters"] + generator = getattr(module, "generator_kind", None) + if generator is not None: + lines.append(f" generator: {generator}") + trunk = getattr(module, "trunk", None) + if trunk is not None: + in_dim = getattr(trunk, "in_dim", "?") + out_dim = getattr(trunk, "out_dim", "?") + if isinstance(trunk, RoutedTrunk): + detail = f"routed, n_experts={trunk.router.n_experts}, expert type={type(trunk.experts[0]).__name__}" + else: + detail = f"unrouted, {type(trunk).__name__}" + lines.append(f" trunk: {detail}, in={in_dim}, out={out_dim}") + history_kind = getattr(module, "history_kind", None) + if history_kind is not None: + lines.append(f" autoregressive history: {history_kind}") + present = [h for h in _HEAD_NAMES if getattr(module, h, None) is not None] + absent = [h for h in _HEAD_NAMES if hasattr(module, h) and getattr(module, h) is None] + if present or absent: + lines.append(f" heads present: {', '.join(present) if present else 'none'}") + if absent: + lines.append(f" heads absent: {', '.join(absent)}") + return lines + + +def render_summary(summary: ModelSummary) -> str: + lines: list[str] = [] + for name, module in summary.modules.items(): + lines.extend(_stage_header(name, module)) + lines.extend(_tree_lines(module, name, indent=1)) + lines.append("") + + lines.append( + f"config keys read during construction: {len(summary.consumed)} / " + f"read elsewhere (trainer/sampler/rollout): {len(summary.elsewhere)} / " + f"inert under this config: {len(summary.inert)}" + ) + if summary.elsewhere: + lines.append("read elsewhere, not by construction:") + for path in summary.elsewhere: + lines.append(f" {path} ({_NOT_BUILD_TIME[path]})") + lines.append("inert under this config (declared, parsed, but doing nothing here):") + if summary.inert: + for path in summary.inert: + lines.append(f" {path}") + else: + lines.append(" (none)") + + if summary.vocab_caveats: + lines.append("") + lines.append("vocab placeholder caveats:") + for caveat in summary.vocab_caveats: + lines.append(f" {caveat}") + + return "\n".join(lines) diff --git a/giant/model/trunks.py b/giant/model/trunks.py index 415d64d..398523e 100644 --- a/giant/model/trunks.py +++ b/giant/model/trunks.py @@ -73,6 +73,7 @@ class ExpertTrunk(nn.Module): block_conditioning: str = "add", ) -> None: super().__init__() + self.in_dim = in_dim self.out_dim = out_dim self.input_proj = nn.Linear(in_dim, hidden_dim) self.blocks = nn.ModuleList( @@ -130,7 +131,10 @@ class Trunk(nn.Module): """Interface implemented by a standalone trunk body (any `TRUNK_REGISTRY` entry, e.g. `ExpertTrunk`) and by `RoutedTrunk`: everything downstream of the fused conditioning vector, i.e. the actual generative trunk of a - stage.""" + stage. Implementations are expected to expose `in_dim`/`out_dim` + attributes (as `ExpertTrunk`/`RoutedTrunk` do) — `giant.model.summary` + (gitea #46) reads them to report trunk widths without needing to know the + body architecture.""" def forward( self, @@ -157,6 +161,8 @@ class RoutedTrunk(Trunk): ) -> None: super().__init__() self.router = router + self.in_dim = in_dim + self.out_dim = out_dim self.experts = nn.ModuleList( [ build_expert_body( diff --git a/tests/test_config_consumed_keys.py b/tests/test_config_consumed_keys.py index 095f9d2..9672a8c 100644 --- a/tests/test_config_consumed_keys.py +++ b/tests/test_config_consumed_keys.py @@ -28,6 +28,7 @@ import ast from pathlib import Path from giant.config import DEFAULT_CONFIG +from giant.config import leaf_paths as _leaf_paths _REPO_ROOT = Path(__file__).resolve().parents[1] @@ -72,19 +73,6 @@ _KNOWN_UNUSED = { _FIELD_NAME_OVERRIDES = {"lambda": "lambda_weight"} -def _leaf_paths(node: dict, prefix: str = "") -> list[str]: - paths = [] - for key, value in node.items(): - if prefix == "" and key == "meta": - continue - path = f"{prefix}.{key}" if prefix else key - if isinstance(value, dict): - paths.extend(_leaf_paths(value, path)) - else: - paths.append(path) - return paths - - def _field_name(leaf_path: str) -> str: name = leaf_path.rsplit(".", 1)[-1] return _FIELD_NAME_OVERRIDES.get(name, name) diff --git a/tests/test_model_summary.py b/tests/test_model_summary.py new file mode 100644 index 0000000..5685858 --- /dev/null +++ b/tests/test_model_summary.py @@ -0,0 +1,137 @@ +"""Tests for `giant model summary` (gitea #46).""" + +from __future__ import annotations + +from pathlib import Path + +import pytest +from typer.testing import CliRunner + +from giant import config as gconfig +from giant.cli import app +from giant.materials import MATERIAL_PROPERTIES +from giant.model.summary import _NOT_BUILD_TIME, _built_modules, _vocab_caveats, summarize_model + +runner = CliRunner() + +_PDG_VOCAB = 300 +_MAT_VOCAB = len(MATERIAL_PROPERTIES) + + +def _cfg(overrides: dict | None = None) -> dict: + return gconfig.merge_cli_overrides(gconfig.DEFAULT_CONFIG, None, overrides or {}) + + +@pytest.fixture(scope="module") +def default_summary(): + return summarize_model(_cfg(), pdg_vocab=_PDG_VOCAB, mat_vocab=_MAT_VOCAB) + + +def test_default_config_builds_both_stages_with_a_real_tree(default_summary): + assert set(default_summary.modules) >= {"stage1", "stage2"} + for module in default_summary.modules.values(): + assert sum(p.numel() for p in module.parameters()) > 0 + stage1 = default_summary.modules["stage1"] + assert hasattr(stage1, "cond_enc") + assert hasattr(stage1, "trunk") + assert {"input_proj", "blocks", "out_proj"} <= {n for n, _ in stage1.trunk.named_children()} + + +def test_every_in_scope_leaf_is_classified(default_summary): + in_scope = { + p + for p in gconfig.leaf_paths(gconfig.DEFAULT_CONFIG) + if p.split(".", 1)[0] in ("conditioning", "stage1_model", "stage2_model") + } + classified = set(default_summary.consumed) | set(default_summary.inert) | set(default_summary.elsewhere) + assert classified == in_scope + + +def test_not_build_time_allow_list_has_no_stale_entries(): + in_scope = set(gconfig.leaf_paths(gconfig.DEFAULT_CONFIG)) + stale = set(_NOT_BUILD_TIME) - in_scope + assert not stale, f"_NOT_BUILD_TIME entries no longer in DEFAULT_CONFIG: {sorted(stale)}" + + +def test_router_disabled_by_default_so_its_fields_are_inert(default_summary): + assert "stage1_model.router.n_experts" in default_summary.inert + assert "stage1_model.router.temperature" in default_summary.inert + + +def test_markov_history_leaves_attention_dims_inert_but_history_itself_consumed(default_summary): + assert "stage2_model.autoregressive.attn_n_heads" in default_summary.inert + assert "stage2_model.autoregressive.attn_n_layers" in default_summary.inert + assert "stage2_model.autoregressive.history" in default_summary.consumed + + +def test_single_literal_branch_fields_are_correctly_seen_as_consumed(default_summary): + """Regression guard: n_sec.owner ("stage2"), n_sec.mode ("head") and + particle_type.target ("onehot") each branch as `== "one specific other + literal"` in giant/model/builders.py|models.py. A naive single generic + sentinel probe lands in the same "not that literal" bucket as the + current value and never crosses the boundary that actually matters -- + this is exactly what _STRING_ALTERNATIVES exists to fix.""" + assert "stage2_model.n_sec.owner" in default_summary.consumed + assert "stage2_model.n_sec.mode" in default_summary.consumed + assert "stage2_model.particle_type.target" in default_summary.consumed + + +def test_stage1_wgan_generator_swaps_flow_time_dim_for_critic_dims(): + summary = summarize_model(_cfg({"stage1_model": {"generator": "wgan"}}), pdg_vocab=_PDG_VOCAB, mat_vocab=_MAT_VOCAB) + assert "stage1_model.flow.time_dim" in summary.inert + assert "stage1_model.wgan.noise_dim" in summary.consumed + assert "stage1_model.wgan.critic_hidden_dim" in summary.consumed + + +def test_stage2_one_shot_decoder_makes_autoregressive_block_inert(): + summary = summarize_model( + _cfg({"stage2_model": {"decoder": "one_shot"}}), pdg_vocab=_PDG_VOCAB, mat_vocab=_MAT_VOCAB + ) + assert "stage2_model.autoregressive.history" in summary.inert + assert "history_encoder" not in {n for n, _ in summary.modules["stage2"].named_children()} + + +def test_energy_router_enabled_consumes_core_fields_but_not_process_only_fields(): + summary = summarize_model( + _cfg({"stage1_model": {"router": {"enabled": True, "type": "energy", "n_experts": 4}}}), + pdg_vocab=_PDG_VOCAB, + mat_vocab=_MAT_VOCAB, + ) + assert "stage1_model.router.n_experts" in summary.consumed + assert "stage1_model.router.temperature" in summary.consumed + # emb_dim/hidden_dim are pdg/process-router-only kwargs -- build_router's + # signature filter drops them for an energy router. + assert "stage1_model.router.hidden_dim" in summary.inert + assert "stage1_model.router.emb_dim" in summary.inert + + +def test_vocab_caveat_text_for_embedding_particle_conditioning(): + cfg = _cfg({"conditioning": {"particle": {"type": "embedding"}}}) + caveats = _vocab_caveats(cfg) + assert any("pdg_vocab" in c and "embedding" in c for c in caveats) + assert not any("mat_vocab" in c for c in caveats) + + +def test_pdg_vocab_flag_changes_embedding_table_size(): + cfg = _cfg({"conditioning": {"particle": {"type": "embedding"}}}) + small = _built_modules(cfg, pdg_vocab=10, mat_vocab=_MAT_VOCAB) + big = _built_modules(cfg, pdg_vocab=1000, mat_vocab=_MAT_VOCAB) + assert big["stage1"].cond_enc.pdg_emb.weight.numel() > small["stage1"].cond_enc.pdg_emb.weight.numel() + + +def test_invalid_combo_exits_nonzero_with_validate_config_message(tmp_path: Path): + config_path = tmp_path / "bad.toml" + config_path.write_text('[meta]\nconfig_version = 3\n\n[stage2_model.particle_type]\ntarget = "embedding"\n') + result = runner.invoke(app, ["model", "summary", "--config", str(config_path)]) + assert result.exit_code == 1 + assert "requires conditioning.particle.type = 'embedding'" in result.output + + +def test_cli_default_smoke(): + result = runner.invoke(app, ["model", "summary"]) + assert result.exit_code == 0, result.output + assert "stage1" in result.output + assert "stage2" in result.output + assert "parameters" in result.output + assert "trunk" in result.output + assert "inert under this config" in result.output