9 Commits

Author SHA1 Message Date
gitea-actions f80fc90758 chore: update changelog for v0.3.13 [skip ci] 2026-08-28 09:50:33 +00:00
gitea-actions 1cf16526c9 chore: bump version 0.3.12 -> 0.3.13 [skip ci] 2026-08-28 09:50:33 +00:00
lars 5c93457081 Merge pull request 'Fix/issue 87' (#89) from fix/issue-87 into master
CI / Sync project version with tag (push) Has been skipped
CI / Lint (ruff check) (push) Successful in 25s
CI / Format (ruff format) (push) Successful in 1m8s
CI / Type check (ty) (push) Successful in 1m29s
CI / Tests (push) Successful in 3m21s
CI / Bump version, tag, and update changelog on merge to master (push) Successful in 1m32s
Reviewed-on: #89
2026-08-28 11:45:00 +02:00
lars 1ec333ff6d Merge remote-tracking branch 'origin/master' into fix/issue-87
CI / Sync project version with tag (push) Has been skipped
CI / Lint (ruff check) (push) Successful in 1m36s
CI / Format (ruff format) (push) Successful in 1m41s
CI / Sync project version with tag (pull_request) Has been skipped
CI / Type check (ty) (push) Successful in 2m17s
CI / Lint (ruff check) (pull_request) Successful in 1m19s
CI / Format (ruff format) (pull_request) Successful in 2m15s
CI / Type check (ty) (pull_request) Successful in 4m13s
CI / Tests (pull_request) Successful in 7m18s
CI / Bump version, tag, and update changelog on merge to master (pull_request) Has been skipped
CI / Tests (push) Successful in 11m5s
CI / Bump version, tag, and update changelog on merge to master (push) Has been skipped
# Conflicts:
#	giant/config.py
2026-08-28 11:31:05 +02:00
gitea-actions 0654fa3f12 chore: update changelog for v0.3.12 [skip ci] 2026-08-28 09:25:57 +00:00
gitea-actions 36fe9bd66d chore: bump version 0.3.11 -> 0.3.12 [skip ci] 2026-08-28 09:25:56 +00:00
lars bd255419e1 Add inference-time model_config overrides with a sampling-key allowlist (gitea #87)
CI / Sync project version with tag (push) Has been skipped
CI / Lint (ruff check) (push) Successful in 1m45s
CI / Format (ruff format) (push) Successful in 3m0s
CI / Type check (ty) (push) Successful in 3m16s
CI / Tests (push) Successful in 3m30s
CI / Bump version, tag, and update changelog on merge to master (push) Has been skipped
giant predict/rollout rebuilt models straight from ckpt["model_config"] with
no way to change sampling-only keys (e.g. stage2_model.n_sec.stop_sampling)
without retraining. Adds config_overrides to load_for_inference, validated
against giant.config.INFERENCE_OVERRIDES so a typo or shape-bearing key
raises CheckpointCompatibilityError up front instead of an opaque
load_state_dict mismatch. Wired as a repeatable --set dotted.path=value on
both CLI commands, recorded in the rollout YAML sidecar, and surfaced in
`giant model summary`'s output.

Co-Authored-By: Claude Sonnet 5 <noreply@anthropic.com>
2026-08-28 11:21:36 +02:00
lars 73975a4587 Merge pull request 'Add sampled n_sec under n_sec.mode = 'head' (gitea #86)' (#88) from fix/issue-86 into master
CI / Sync project version with tag (push) Has been skipped
CI / Lint (ruff check) (push) Successful in 45s
CI / Format (ruff format) (push) Successful in 54s
CI / Type check (ty) (push) Successful in 59s
CI / Tests (push) Successful in 5m24s
CI / Bump version, tag, and update changelog on merge to master (push) Successful in 2m31s
Reviewed-on: #88
2026-08-28 11:18:06 +02:00
lars fcd77c2f4b Add sampled n_sec under n_sec.mode = 'head' (gitea #86)
CI / Sync project version with tag (push) Has been skipped
CI / Lint (ruff check) (push) Successful in 1m29s
CI / Type check (ty) (push) Successful in 1m26s
CI / Format (ruff format) (push) Successful in 1m26s
CI / Sync project version with tag (pull_request) Has been skipped
CI / Lint (ruff check) (pull_request) Successful in 4m0s
CI / Format (ruff format) (pull_request) Successful in 3m59s
CI / Tests (push) Successful in 5m41s
CI / Type check (ty) (pull_request) Successful in 4m1s
CI / Bump version, tag, and update changelog on merge to master (push) Has been skipped
CI / Tests (pull_request) Successful in 4m15s
CI / Bump version, tag, and update changelog on merge to master (pull_request) Has been skipped
Taking argmax over the n_sec classifier logits collapses secondary
multiplicity onto its conditional mode at fixed pre-step conditioning,
under-dispersing n_sec in rollouts and biasing low wherever the true
conditional count distribution is right-skewed (typical for
multiplicity).

Generalizes stage2_model.n_sec.stop_sampling (previously stop_token-only)
into stage2_model.n_sec.sampling, covering both "head" (greedy: argmax;
sample: categorical draw via torch.multinomial) and "stop_token" (unchanged:
greedy threshold / Bernoulli draw) modes. stop_sampling is kept as a
deprecated alias in NSecConfig.from_dict and migrate_config, since it
appears in existing checkpoints' model_config. Default stays "greedy" so
existing runs/checkpoints are unaffected.
2026-08-28 11:04:45 +02:00
17 changed files with 562 additions and 46 deletions
+1 -1
View File
@@ -1,5 +1,5 @@
[tool.bumpversion] [tool.bumpversion]
current_version = "0.3.11" current_version = "0.3.13"
parse = "(?P<major>\\d+)\\.(?P<minor>\\d+)\\.(?P<patch>\\d+)" parse = "(?P<major>\\d+)\\.(?P<minor>\\d+)\\.(?P<patch>\\d+)"
serialize = ["{major}.{minor}.{patch}"] serialize = ["{major}.{minor}.{patch}"]
search = "{current_version}" search = "{current_version}"
+12
View File
@@ -1,5 +1,17 @@
# Changelog # Changelog
## [0.3.13] - 2026-08-28
### Added
- Add inference-time model_config overrides with a sampling-key allowlist [gitea #87](https://git.larsbogner.de/lars/giant/issues/87)
## [0.3.12] - 2026-08-28
### Added
- Add sampled n_sec under n_sec.mode = 'head' [gitea #86](https://git.larsbogner.de/lars/giant/issues/86)
## [0.3.11] - 2026-08-26 ## [0.3.11] - 2026-08-26
### Changed ### Changed
+52 -3
View File
@@ -14,11 +14,18 @@ directly and imported from non-CLI code (`giant.analysis.router_gating`,
lazily — see that module's docstring for why). Failures raise lazily — see that module's docstring for why). Failures raise
`CheckpointCompatibilityError` with the same wording the CLI has always `CheckpointCompatibilityError` with the same wording the CLI has always
shown; the CLI layer catches it and does the `typer.echo`/`Exit(1)`. shown; the CLI layer catches it and does the `typer.echo`/`Exit(1)`.
`load_for_inference`'s `config_overrides` (gitea #87) lets a caller change a
checkpoint's `model_config` at load time, restricted to
`giant.config.INFERENCE_OVERRIDES` — the allowlist of keys that only affect
sampling, never module construction/shapes or the preprocessing normalizers/
vocab maps were fit under.
""" """
from __future__ import annotations from __future__ import annotations
from dataclasses import dataclass import copy
from dataclasses import dataclass, field
from pathlib import Path from pathlib import Path
import torch import torch
@@ -29,13 +36,45 @@ from giant.constants import K_MAX
from giant.data.loader import TopNMap from giant.data.loader import TopNMap
from giant.data.setup_cache import topnmap_from_json from giant.data.setup_cache import topnmap_from_json
from giant.data.transforms import Normalizer from giant.data.transforms import Normalizer
from giant.model.network import build_models from giant.model.network import _migrate_legacy_model_config, build_models
class CheckpointCompatibilityError(Exception): class CheckpointCompatibilityError(Exception):
"""Checkpoint is missing something `load_for_inference` needs.""" """Checkpoint is missing something `load_for_inference` needs."""
def apply_config_overrides(model_cfg: dict, overrides: dict[str, object] | None) -> dict:
"""Deep-merge dotted-path *overrides* into a checkpoint's `model_config`,
validated against `giant.config.INFERENCE_OVERRIDES` — the allowlist of
keys that only affect sampling, not module construction/shapes or the
preprocessing normalizers/vocab maps were fit under (gitea #87).
Migrates a v0.2 flat `model_config` to the nested v0.3 shape first: a
dotted path like "stage1_model.ddpm.n_steps" would otherwise silently
write into a dict that `build_models` still reads as flat (it decides
v0.2-vs-v0.3 by `"stage1_model" in model_config`), suppressing migration.
Raises `CheckpointCompatibilityError` — never a bare `ValueError` or a
downstream `load_state_dict` size mismatch — for an unknown/disallowed
path or a value that fails its allowlisted check.
"""
if not overrides:
return model_cfg
cfg = model_cfg if "stage1_model" in model_cfg else _migrate_legacy_model_config(model_cfg)
cfg = copy.deepcopy(cfg)
for path, value in overrides.items():
spec = gconfig.INFERENCE_OVERRIDES.get(path)
if spec is None:
allowed = ", ".join(sorted(gconfig.INFERENCE_OVERRIDES))
raise CheckpointCompatibilityError(f"{path!r} is not an inference-safe override — allowed paths: {allowed}")
try:
spec.check(path, value)
except ValueError as exc:
raise CheckpointCompatibilityError(str(exc)) from exc
gconfig._set_path(cfg, path, value)
return cfg
def conditioning_axes(model_cfg: dict, default: str = "embedding") -> tuple[str, str]: def conditioning_axes(model_cfg: dict, default: str = "embedding") -> tuple[str, str]:
"""(particle_conditioning, material_conditioning) for """(particle_conditioning, material_conditioning) for
`giant.data.transforms.build_cond_features`/`build_features` — from `giant.data.transforms.build_cond_features`/`build_features` — from
@@ -128,6 +167,7 @@ class InferenceContext:
model_config: dict model_config: dict
epoch: int | None epoch: int | None
best_val_loss: float | None best_val_loss: float | None
config_overrides: dict[str, object] = field(default_factory=dict)
def load_for_inference( def load_for_inference(
@@ -136,6 +176,7 @@ def load_for_inference(
command_name: str, command_name: str,
weights: str = "raw", weights: str = "raw",
require_stage2: bool = True, require_stage2: bool = True,
config_overrides: dict[str, object] | None = None,
) -> InferenceContext: ) -> InferenceContext:
"""Load *checkpoint* and reconstruct everything `predict`/`rollout` need """Load *checkpoint* and reconstruct everything `predict`/`rollout` need
to run it forward, on *device*, in `eval()` mode. to run it forward, on *device*, in `eval()` mode.
@@ -148,6 +189,13 @@ def load_for_inference(
both stages) or an acceptable `stage2 = None` result — kept as a real both stages) or an acceptable `stage2 = None` result — kept as a real
parameter since `stage{1,2}_model.active` is a real, if currently parameter since `stage{1,2}_model.active` is a real, if currently
stage1+stage2-only-in-practice, config option. stage1+stage2-only-in-practice, config option.
*config_overrides* deep-merges dotted `model_config` paths (e.g.
`{"stage2_model.n_sec.sampling": "sample"}`) before anything is
derived from `model_config` or built — see `apply_config_overrides` for
the allowlist and validation. Every derived `InferenceContext` field
(`other_policy`, `stage{1,2}_ddpm_steps`, the built modules, ...)
reflects the overridden config.
""" """
ckpt = torch.load(checkpoint, map_location="cpu", weights_only=False) ckpt = torch.load(checkpoint, map_location="cpu", weights_only=False)
for key in ("model_config", "sec_decoder"): for key in ("model_config", "sec_decoder"):
@@ -159,7 +207,7 @@ def load_for_inference(
gconfig.warn_if_checkpoint_config_mismatch(checkpoint) gconfig.warn_if_checkpoint_config_mismatch(checkpoint)
model_cfg = ckpt["model_config"] model_cfg = apply_config_overrides(ckpt["model_config"], config_overrides)
particle_conditioning, material_conditioning = conditioning_axes(model_cfg) particle_conditioning, material_conditioning = conditioning_axes(model_cfg)
pdg_topn_map = load_pdg_topn_map(ckpt) pdg_topn_map = load_pdg_topn_map(ckpt)
mat_topn_map = load_mat_topn_map(ckpt) mat_topn_map = load_mat_topn_map(ckpt)
@@ -231,4 +279,5 @@ def load_for_inference(
model_config=model_cfg, model_config=model_cfg,
epoch=ckpt.get("epoch"), epoch=ckpt.get("epoch"),
best_val_loss=ckpt.get("best_val_loss"), best_val_loss=ckpt.get("best_val_loss"),
config_overrides=dict(config_overrides) if config_overrides else {},
) )
+44 -2
View File
@@ -141,6 +141,23 @@ def _parse_router_axis_flags(specs: list[str]) -> dict[str, object]:
return out return out
def _parse_set_flags(specs: Optional[list[str]]) -> dict[str, object]:
"""Parse repeated `--set dotted.path=value` flags into a dict, typing
each value with `_coerce_scalar` the same way a TOML file's native types
would arrive. Validation against the inference-safe allowlist happens
downstream in `giant.checkpoint_io.apply_config_overrides` — this only
parses syntax.
"""
out: dict[str, object] = {}
for spec in specs or []:
path, sep, val = spec.partition("=")
if not sep:
typer.echo(f"error: --set {spec!r} must be 'dotted.path=value'", err=True)
raise typer.Exit(1)
out[path] = _coerce_scalar(val)
return out
def _router_cli_overrides( def _router_cli_overrides(
router: bool | None, router: bool | None,
router_type: str | None, router_type: str | None,
@@ -1020,6 +1037,15 @@ def predict(
help="Free-text note recorded in the prediction's YAML sidecar", help="Free-text note recorded in the prediction's YAML sidecar",
), ),
] = None, ] = None,
set_: Annotated[
Optional[list[str]],
typer.Option(
"--set",
help="Override a sampling-only model_config key on this checkpoint, "
"'dotted.path=value' (repeatable) — see giant.config.INFERENCE_OVERRIDES "
"for the allowlist, e.g. --set stage2_model.n_sec.sampling=sample",
),
] = None,
) -> None: ) -> None:
"""Run trained model on a parquet file and save predictions.""" """Run trained model on a parquet file and save predictions."""
batch_size_auto = False batch_size_auto = False
@@ -1040,8 +1066,11 @@ def predict(
typer.echo(f"device: {_device}") typer.echo(f"device: {_device}")
# --- Load checkpoint --- # --- Load checkpoint ---
config_overrides = _parse_set_flags(set_)
try: try:
ctx = load_for_inference(checkpoint, _device, "predict", weights=weights.value) ctx = load_for_inference(
checkpoint, _device, "predict", weights=weights.value, config_overrides=config_overrides
)
except CheckpointCompatibilityError as exc: except CheckpointCompatibilityError as exc:
typer.echo(f"error: {exc}", err=True) typer.echo(f"error: {exc}", err=True)
raise typer.Exit(1) raise typer.Exit(1)
@@ -1407,6 +1436,15 @@ def rollout(
Optional[int], Optional[int],
typer.Option("--seed", help="Torch/numpy seed for reproducibility"), typer.Option("--seed", help="Torch/numpy seed for reproducibility"),
] = None, ] = None,
set_: Annotated[
Optional[list[str]],
typer.Option(
"--set",
help="Override a sampling-only model_config key on this checkpoint, "
"'dotted.path=value' (repeatable) — see giant.config.INFERENCE_OVERRIDES "
"for the allowlist, e.g. --set stage2_model.n_sec.sampling=sample",
),
] = None,
) -> None: ) -> None:
"""Roll the surrogate forward into full showers (autoregressive).""" """Roll the surrogate forward into full showers (autoregressive)."""
if seed is not None: if seed is not None:
@@ -1416,8 +1454,11 @@ def rollout(
_device = torch.device(device) if device else gconfig.auto_device() _device = torch.device(device) if device else gconfig.auto_device()
typer.echo(f"device: {_device}") typer.echo(f"device: {_device}")
config_overrides = _parse_set_flags(set_)
try: try:
ctx = load_for_inference(checkpoint, _device, "rollout", weights=weights.value) ctx = load_for_inference(
checkpoint, _device, "rollout", weights=weights.value, config_overrides=config_overrides
)
except CheckpointCompatibilityError as exc: except CheckpointCompatibilityError as exc:
typer.echo(f"error: {exc}", err=True) typer.echo(f"error: {exc}", err=True)
raise typer.Exit(1) raise typer.Exit(1)
@@ -1532,6 +1573,7 @@ def rollout(
# model knob (router type/n_experts, noise_dim, vocab sizes, ...) # model knob (router type/n_experts, noise_dim, vocab sizes, ...)
# is available downstream without touching this command again. # is available downstream without touching this command again.
"model_config": dict(model_cfg), "model_config": dict(model_cfg),
"config_overrides": dict(ctx.config_overrides),
"training_epoch": ctx.epoch, "training_epoch": ctx.epoch,
"best_val_loss": ctx.best_val_loss, "best_val_loss": ctx.best_val_loss,
# [train]/[meta] from the sibling config.toml (giant.config.save_config) # [train]/[meta] from the sibling config.toml (giant.config.save_config)
+124 -13
View File
@@ -404,6 +404,16 @@ class Stage2RouterConfig(RouterConfig):
return {"tie_to_stage1": self.tie_to_stage1, **super().to_dict()} return {"tie_to_stage1": self.tie_to_stage1, **super().to_dict()}
# stage2_model.n_sec.sampling choices — single source of truth for both
# validate_config's train-time check and INFERENCE_OVERRIDES below.
STOP_SAMPLING_CHOICES = ("greedy", "sample")
# stage2_model.particle_type.other_policy choices — see ParticleTypeConfig's
# docstring for what each means; only documented there until now, since
# nothing validated it at train time.
OTHER_POLICY_CHOICES = ("sample", "modal", "drop")
@dataclass(frozen=True) @dataclass(frozen=True)
class NSecConfig: class NSecConfig:
# "head": a classifier over {0..k_max} on the condition encoding alone # "head": a classifier over {0..k_max} on the condition encoding alone
@@ -424,11 +434,15 @@ class NSecConfig:
# n_sec head was trained against Stage 1's own ConditionEncoder output and so has # n_sec head was trained against Stage 1's own ConditionEncoder output and so has
# to stay attached there, not just be labeled as such). # to stay attached there, not just be labeled as such).
owner: str = "stage2" owner: str = "stage2"
# mode="stop_token" only: how sample_secondaries_ar turns a slot's stop logit into a # How resolve_n_sec/sample_secondaries_ar turn a count-bearing head's output into an
# stop/continue decision. "greedy": sigmoid(logit) >= 0.5 (deterministic). "sample": # actual n_sec decision. mode="head": "greedy" is argmax over the classifier logits
# a Bernoulli draw at sigmoid(logit) (a real sample from the learned length # (deterministic — the conditional mode, not a sample); "sample" is a categorical draw
# distribution, at the cost of an extra RNG draw per slot). # from softmax(logits) (a real sample from the learned count distribution). mode=
stop_sampling: str = "greedy" # "stop_token": "greedy" is sigmoid(stop_logit) >= 0.5 per slot (deterministic);
# "sample" is a Bernoulli draw at sigmoid(stop_logit) per slot. Renamed from
# "stop_sampling" (gitea #86), which is still accepted as a deprecated alias since it
# appears in existing checkpoints' model_config.
sampling: str = "greedy"
@classmethod @classmethod
def from_dict(cls, d: dict | None) -> "NSecConfig": def from_dict(cls, d: dict | None) -> "NSecConfig":
@@ -437,7 +451,7 @@ class NSecConfig:
mode=d.get("mode", "head"), mode=d.get("mode", "head"),
lambda_weight=d.get("lambda", 0.1), lambda_weight=d.get("lambda", 0.1),
owner=d.get("owner", "stage2"), owner=d.get("owner", "stage2"),
stop_sampling=d.get("stop_sampling", "greedy"), sampling=d.get("sampling", d.get("stop_sampling", "greedy")),
) )
def to_dict(self) -> dict: def to_dict(self) -> dict:
@@ -445,7 +459,7 @@ class NSecConfig:
"mode": self.mode, "mode": self.mode,
"lambda": self.lambda_weight, "lambda": self.lambda_weight,
"owner": self.owner, "owner": self.owner,
"stop_sampling": self.stop_sampling, "sampling": self.sampling,
} }
@@ -1075,6 +1089,27 @@ def _set_path(d: dict, dotted: str, value) -> None:
cur[parts[-1]] = value cur[parts[-1]] = value
def _pop_path(d: dict, dotted: str) -> None:
"""Remove a dotted path from a nested dict, if present. No-op if any
component along the path is missing."""
parts = dotted.split(".")
cur = d
for part in parts[:-1]:
if not isinstance(cur, dict) or part not in cur:
return
cur = cur[part]
if isinstance(cur, dict):
cur.pop(parts[-1], None)
# Config keys renamed within v0.3 itself (not part of the v0.2->v0.3 migration
# above) — normalized by migrate_config so a config.toml still using an older
# v0.3 key name keeps passing validate_config_keys.
_RENAMED_KEYS = {
"stage2_model.n_sec.stop_sampling": "stage2_model.n_sec.sampling", # gitea #86
}
def _deep_merge(base: dict, override: dict) -> dict: def _deep_merge(base: dict, override: dict) -> dict:
"""Recursively merge `override` onto a copy of `base`. """Recursively merge `override` onto a copy of `base`.
@@ -1093,6 +1128,72 @@ def _deep_merge(base: dict, override: dict) -> dict:
return result return result
@dataclass(frozen=True)
class InferenceOverride:
"""One dotted `model_config` path that `giant.checkpoint_io.load_for_inference`
is allowed to change on an already-trained checkpoint, without retraining.
A path only belongs here if it affects neither module construction/tensor
shapes nor the data preprocessing the normalizers/vocab maps were fit
under see the module docstring on `giant.model.summary` for the class
of key this targets (`_fingerprint`'s "plain scalar attribute" leaves),
and `giant.checkpoint_io.apply_config_overrides` for where this is used.
"""
why: str
choices: tuple[str, ...] | None = None
minimum: float | None = None
numeric: bool = False # int/float leaf (vs. str, the default)
def check(self, path: str, value: object) -> None:
if self.choices is not None:
if value not in self.choices:
raise ValueError(f"{path} = {value!r} — must be one of {self.choices}")
return
if self.numeric:
if isinstance(value, bool) or not isinstance(value, (int, float)):
raise ValueError(f"{path} = {value!r} — must be a number")
if self.minimum is not None and value < self.minimum:
raise ValueError(f"{path} = {value!r} — must be >= {self.minimum}")
# Inference-safe dotted `model_config` paths — the allowlist gitea #87 asked
# for, so a typo or a shape-bearing key (e.g. "stage1_model.hidden_dim")
# raises a clear CheckpointCompatibilityError instead of surfacing as an
# opaque load_state_dict size mismatch later. Extend this table, not a
# per-call bypass, when a new inference-only key needs the capability.
INFERENCE_OVERRIDES: dict[str, InferenceOverride] = {
"stage2_model.n_sec.sampling": InferenceOverride(
why="giant.sample's n_sec head/stop-token sampling reads this at sample time only (gitea #86)",
choices=STOP_SAMPLING_CHOICES,
),
"stage1_model.ddpm.n_steps": InferenceOverride(
why="giant.model.schedule.CosineSchedule's step count, resolved at sample time",
numeric=True,
minimum=1,
),
"stage2_model.ddpm.n_steps": InferenceOverride(
why="giant.model.schedule.CosineSchedule's step count, resolved at sample time",
numeric=True,
minimum=1,
),
"stage2_model.particle_type.other_policy": InferenceOverride(
why="giant.rollout resolves an 'other'-bucket secondary's PDG code with this at rollout time",
choices=OTHER_POLICY_CHOICES,
),
"stage1_model.router.temperature": InferenceOverride(
why="giant.model.routers.EnergyRouter.temperature, a plain constructor attribute",
numeric=True,
minimum=1e-6,
),
"stage2_model.router.temperature": InferenceOverride(
why="giant.model.routers.EnergyRouter.temperature, a plain constructor attribute",
numeric=True,
minimum=1e-6,
),
}
@dataclass(frozen=True) @dataclass(frozen=True)
class FlagSpec: class FlagSpec:
"""One CLI flag's mapping into the config-overrides tree. """One CLI flag's mapping into the config-overrides tree.
@@ -1271,11 +1372,21 @@ def migrate_config(cfg: dict) -> dict:
(which additionally carries n_sec_head ownership and needs (which additionally carries n_sec_head ownership and needs
`network.build_models`'s cooperation) is a separate migration surface, `network.build_models`'s cooperation) is a separate migration surface,
deferred to the network.py refactor. deferred to the network.py refactor.
"""
if _get_path(cfg, "meta.config_version") == CONFIG_VERSION:
return copy.deepcopy(cfg)
Independently of the v0.2/v0.3 branch below, `_RENAMED_KEYS` normalizes
keys renamed within v0.3 itself (e.g. `stop_sampling` -> `sampling`,
gitea #86) so a config.toml written against an older v0.3 key name still
passes `validate_config_keys`.
"""
cfg = copy.deepcopy(cfg) cfg = copy.deepcopy(cfg)
for old_path, new_path in _RENAMED_KEYS.items():
if _get_path(cfg, old_path) is not None and _get_path(cfg, new_path) is None:
_set_path(cfg, new_path, _get_path(cfg, old_path))
_pop_path(cfg, old_path)
if _get_path(cfg, "meta.config_version") == CONFIG_VERSION:
return cfg
old_train = cfg.pop("train", {}) old_train = cfg.pop("train", {})
old_model = cfg.pop("model", {}) old_model = cfg.pop("model", {})
old_router = dict(old_model.pop("router", {})) old_router = dict(old_model.pop("router", {}))
@@ -1512,9 +1623,9 @@ def validate_config(cfg: dict, *, resume: bool = False) -> None:
"conditioning to hang an EOS decision off" "conditioning to hang an EOS decision off"
) )
stop_sampling = _get_path(cfg, "stage2_model.n_sec.stop_sampling") n_sec_sampling = _get_path(cfg, "stage2_model.n_sec.sampling")
if stop_sampling not in ("greedy", "sample"): if n_sec_sampling not in STOP_SAMPLING_CHOICES:
raise ValueError(f"stage2_model.n_sec.stop_sampling = {stop_sampling!r} — must be 'greedy' or 'sample'") raise ValueError(f"stage2_model.n_sec.sampling = {n_sec_sampling!r} — must be 'greedy' or 'sample'")
precision = _get_path(cfg, "train.precision") precision = _get_path(cfg, "train.precision")
if precision not in ("fp32", "bf16"): if precision not in ("fp32", "bf16"):
+2 -1
View File
@@ -138,7 +138,7 @@ def build_models(model_config: dict) -> dict[str, nn.Module | None]:
n_sec_head_cfg=s2_spec.heads.n_sec.to_dict(), n_sec_head_cfg=s2_spec.heads.n_sec.to_dict(),
type_head_cfg=s2_spec.heads.type.to_dict(), type_head_cfg=s2_spec.heads.type.to_dict(),
build_stop_head=stop_token, build_stop_head=stop_token,
stop_sampling=s2_spec.n_sec.stop_sampling, n_sec_sampling=s2_spec.n_sec.sampling,
stop_head_cfg=s2_spec.heads.n_sec.to_dict(), stop_head_cfg=s2_spec.heads.n_sec.to_dict(),
) )
else: else:
@@ -168,6 +168,7 @@ def build_models(model_config: dict) -> dict[str, nn.Module | None]:
cond_enc=shared_cond_enc, cond_enc=shared_cond_enc,
n_sec_head_cfg=s2_spec.heads.n_sec.to_dict(), n_sec_head_cfg=s2_spec.heads.n_sec.to_dict(),
type_head_cfg=s2_spec.heads.type.to_dict(), type_head_cfg=s2_spec.heads.type.to_dict(),
n_sec_sampling=s2_spec.n_sec.sampling,
) )
return result return result
+4 -2
View File
@@ -370,6 +370,7 @@ class Stage2OneShot(StageModel):
cond_enc: ConditionEncoder | None = None, cond_enc: ConditionEncoder | None = None,
n_sec_head_cfg: dict | None = None, n_sec_head_cfg: dict | None = None,
type_head_cfg: dict | None = None, type_head_cfg: dict | None = None,
n_sec_sampling: str = "greedy",
) -> None: ) -> None:
super().__init__( super().__init__(
pdg_vocab, pdg_vocab,
@@ -383,6 +384,7 @@ class Stage2OneShot(StageModel):
particle_type_cfg=particle_type_cfg, particle_type_cfg=particle_type_cfg,
cond_enc=cond_enc, cond_enc=cond_enc,
) )
self.n_sec_sampling = n_sec_sampling
self._build_context_fusion(x_dim, context_dim, cond_out_dim) self._build_context_fusion(x_dim, context_dim, cond_out_dim)
target = self.particle_type_cfg.target target = self.particle_type_cfg.target
type_head_out_dim = None if target == "physical" else k_max * self.type_dim type_head_out_dim = None if target == "physical" else k_max * self.type_dim
@@ -498,7 +500,7 @@ class Stage2Autoregressive(StageModel):
n_sec_head_cfg: dict | None = None, n_sec_head_cfg: dict | None = None,
type_head_cfg: dict | None = None, type_head_cfg: dict | None = None,
build_stop_head: bool = False, build_stop_head: bool = False,
stop_sampling: str = "greedy", n_sec_sampling: str = "greedy",
stop_head_cfg: dict | None = None, stop_head_cfg: dict | None = None,
) -> None: ) -> None:
super().__init__( super().__init__(
@@ -514,7 +516,7 @@ class Stage2Autoregressive(StageModel):
cond_enc=cond_enc, cond_enc=cond_enc,
) )
self.history_kind = history self.history_kind = history
self.stop_sampling = stop_sampling self.n_sec_sampling = n_sec_sampling
self.context_adapter = ContextAdapter(x_dim, context_dim) self.context_adapter = ContextAdapter(x_dim, context_dim)
self.base_fuse = nn.Sequential( self.base_fuse = nn.Sequential(
nn.Linear(cond_out_dim + context_dim, cond_out_dim), nn.Linear(cond_out_dim + context_dim, cond_out_dim),
+12 -2
View File
@@ -37,7 +37,7 @@ from dataclasses import dataclass, field
import torch.nn as nn import torch.nn as nn
from giant.config import _get_path, _set_path, leaf_paths from giant.config import INFERENCE_OVERRIDES, _get_path, _set_path, leaf_paths
from giant.model.builders import build_critics, build_models from giant.model.builders import build_critics, build_models
from giant.model.trunks import RoutedTrunk from giant.model.trunks import RoutedTrunk
@@ -111,6 +111,7 @@ class ModelSummary:
pdg_vocab: int pdg_vocab: int
mat_vocab: int mat_vocab: int
vocab_caveats: list[str] = field(default_factory=list) vocab_caveats: list[str] = field(default_factory=list)
overridable: list[str] = field(default_factory=list)
def _build_model_config(cfg: dict, pdg_vocab: int, mat_vocab: int) -> dict: def _build_model_config(cfg: dict, pdg_vocab: int, mat_vocab: int) -> dict:
@@ -140,7 +141,7 @@ def _fingerprint(modules: dict[str, nn.Module]) -> list:
exist, every parameter's/buffer's shape+dtype (never values those are exist, every parameter's/buffer's shape+dtype (never values those are
randomly initialized and irrelevant to *structure*), and every plain randomly initialized and irrelevant to *structure*), and every plain
scalar attribute any module stores on itself (e.g. `Stage2Autoregressive scalar attribute any module stores on itself (e.g. `Stage2Autoregressive
.stop_sampling`, `EnergyRouter.temperature`) this is what makes a .n_sec_sampling`, `EnergyRouter.temperature`) this is what makes a
non-parametric key's effect on construction observable.""" non-parametric key's effect on construction observable."""
sig = [] sig = []
for stage_name, module in modules.items(): for stage_name, module in modules.items():
@@ -234,6 +235,8 @@ def summarize_model(cfg: dict, pdg_vocab: int, mat_vocab: int) -> ModelSummary:
else: else:
inert.append(path) inert.append(path)
overridable = sorted(p for p in in_scope if p in INFERENCE_OVERRIDES)
return ModelSummary( return ModelSummary(
modules=modules, modules=modules,
consumed=sorted(consumed), consumed=sorted(consumed),
@@ -242,6 +245,7 @@ def summarize_model(cfg: dict, pdg_vocab: int, mat_vocab: int) -> ModelSummary:
pdg_vocab=pdg_vocab, pdg_vocab=pdg_vocab,
mat_vocab=mat_vocab, mat_vocab=mat_vocab,
vocab_caveats=_vocab_caveats(cfg), vocab_caveats=_vocab_caveats(cfg),
overridable=overridable,
) )
@@ -309,6 +313,12 @@ def render_summary(summary: ModelSummary) -> str:
else: else:
lines.append(" (none)") lines.append(" (none)")
if summary.overridable:
lines.append("")
lines.append("inference-overridable without retraining (giant predict/rollout --set):")
for path in summary.overridable:
lines.append(f" {path} ({INFERENCE_OVERRIDES[path].why})")
if summary.vocab_caveats: if summary.vocab_caveats:
lines.append("") lines.append("")
lines.append("vocab placeholder caveats:") lines.append("vocab placeholder caveats:")
+10 -3
View File
@@ -261,7 +261,7 @@ def sample_secondaries_ar(
slot's own stop logit (`predict_stop`, evaluated on the same prefix slot's own stop logit (`predict_stop`, evaluated on the same prefix
conditioning as the token itself see `predict_type`'s docstring for conditioning as the token itself see `predict_type`'s docstring for
why this needs no extra state) decides whether generation should have why this needs no extra state) decides whether generation should have
already stopped, per `sec_decoder.stop_sampling` ("greedy": threshold at already stopped, per `sec_decoder.n_sec_sampling` ("greedy": threshold at
0; "sample": a Bernoulli draw at `sigmoid(logit)`). A row's own 0; "sample": a Bernoulli draw at `sigmoid(logit)`). A row's own
`n_sec_pred` is the first slot index where this fires; once every row in `n_sec_pred` is the first slot index where this fires; once every row in
the batch has fired, the loop breaks before spending a model call on the the batch has fired, the loop breaks before spending a model call on the
@@ -345,7 +345,7 @@ def sample_secondaries_ar(
slot_idx, slot_idx,
hist=hist, hist=hist,
).squeeze(1) ).squeeze(1)
if sec_decoder.stop_sampling == "sample": if sec_decoder.n_sec_sampling == "sample":
stop_now = torch.rand(B, device=device) < torch.sigmoid(stop_logit) stop_now = torch.rand(B, device=device) < torch.sigmoid(stop_logit)
else: else:
stop_now = stop_logit >= 0.0 stop_now = stop_logit >= 0.0
@@ -496,7 +496,12 @@ def resolve_n_sec(
Raises if neither stage owns any n_sec mechanism at all the only way Raises if neither stage owns any n_sec mechanism at all the only way
that happens is `stage2_model.n_sec.mode = "truth"`, which is not a valid that happens is `stage2_model.n_sec.mode = "truth"`, which is not a valid
rollout-/predict-capable checkpoint.""" rollout-/predict-capable checkpoint.
`n_sec.mode = "head"` resolves the classifier logits per
`sec_decoder.n_sec_sampling`: "greedy" (default) takes the conditional
mode via argmax; "sample" draws a real sample from the learned count
distribution via `torch.multinomial` on the softmax see gitea #86."""
if n_sec_pred is not None: if n_sec_pred is not None:
return n_sec_pred return n_sec_pred
if getattr(sec_decoder, "stop_head", None) is not None: if getattr(sec_decoder, "stop_head", None) is not None:
@@ -508,4 +513,6 @@ def resolve_n_sec(
"'truth' is standalone-evaluation-only" "'truth' is standalone-evaluation-only"
) )
logits = sec_decoder.predict_n_sec(cond_cont, cond_cat, stage1_out) logits = sec_decoder.predict_n_sec(cond_cont, cond_cat, stage1_out)
if sec_decoder.n_sec_sampling == "sample":
return torch.multinomial(logits.softmax(dim=-1), 1).squeeze(-1)
return logits.argmax(dim=-1) return logits.argmax(dim=-1)
+1 -1
View File
@@ -1,6 +1,6 @@
[project] [project]
name = "giant" name = "giant"
version = "0.3.11" version = "0.3.13"
description = "Geant4 step-function surrogate via conditional flow matching" description = "Geant4 step-function surrogate via conditional flow matching"
readme = "README.md" readme = "README.md"
requires-python = ">=3.12" requires-python = ">=3.12"
+120
View File
@@ -14,6 +14,7 @@ from giant import config as gconfig
from giant.checkpoint_io import ( from giant.checkpoint_io import (
CheckpointCompatibilityError, CheckpointCompatibilityError,
InferenceContext, InferenceContext,
apply_config_overrides,
conditioning_axes, conditioning_axes,
load_for_inference, load_for_inference,
stage_cfg, stage_cfg,
@@ -278,3 +279,122 @@ def test_stage_cfg_new_shape_returns_subdict():
def test_stage_cfg_v02_flat_shape_returns_empty_dict(): def test_stage_cfg_v02_flat_shape_returns_empty_dict():
model_cfg = {"hidden_dim": 32, "n_blocks": 4} model_cfg = {"hidden_dim": 32, "n_blocks": 4}
assert stage_cfg(model_cfg, "stage2") == {} assert stage_cfg(model_cfg, "stage2") == {}
# ---------------------------------------------------------------------------
# config_overrides (gitea #87)
# ---------------------------------------------------------------------------
def _router_model_cfg() -> dict:
cfg = _model_cfg()
cfg["stage1_model"]["router"] = {"enabled": True, "type": "energy", "n_experts": 2}
return cfg
def test_config_override_n_sec_sampling_changes_stage2_attribute(tmp_path):
checkpoint = _write_checkpoint(tmp_path)
ctx = load_for_inference(
checkpoint,
torch.device("cpu"),
"predict",
config_overrides={"stage2_model.n_sec.sampling": "sample"},
)
assert ctx.stage2 is not None
assert ctx.stage2.n_sec_sampling == "sample"
assert ctx.config_overrides == {"stage2_model.n_sec.sampling": "sample"}
def test_config_override_ddpm_n_steps_changes_context_fields(tmp_path):
checkpoint = _write_checkpoint(tmp_path)
ctx = load_for_inference(
checkpoint,
torch.device("cpu"),
"predict",
config_overrides={"stage1_model.ddpm.n_steps": 42, "stage2_model.ddpm.n_steps": 7},
)
assert ctx.stage1_ddpm_steps == 42
assert ctx.stage2_ddpm_steps == 7
def test_config_override_other_policy_changes_context_field(tmp_path):
checkpoint = _write_checkpoint(tmp_path)
ctx = load_for_inference(
checkpoint,
torch.device("cpu"),
"predict",
config_overrides={"stage2_model.particle_type.other_policy": "modal"},
)
assert ctx.other_policy == "modal"
def test_config_override_router_temperature_changes_router_attribute(tmp_path):
checkpoint = _write_checkpoint(tmp_path, model_cfg=_router_model_cfg())
ctx = load_for_inference(
checkpoint,
torch.device("cpu"),
"predict",
config_overrides={"stage1_model.router.temperature": 1.5},
)
assert ctx.stage1 is not None
assert ctx.stage1.trunk.router.temperature == pytest.approx(1.5)
def test_config_override_no_overrides_defaults_to_empty_dict(tmp_path):
checkpoint = _write_checkpoint(tmp_path)
ctx = load_for_inference(checkpoint, torch.device("cpu"), "predict")
assert ctx.config_overrides == {}
def test_config_override_unknown_path_raises(tmp_path):
checkpoint = _write_checkpoint(tmp_path)
with pytest.raises(CheckpointCompatibilityError, match="not an inference-safe override"):
load_for_inference(
checkpoint,
torch.device("cpu"),
"predict",
config_overrides={"stage2_model.n_sec.typo": "sample"},
)
def test_config_override_shape_bearing_key_raises_up_front(tmp_path):
checkpoint = _write_checkpoint(tmp_path)
with pytest.raises(CheckpointCompatibilityError, match="not an inference-safe override"):
load_for_inference(
checkpoint,
torch.device("cpu"),
"predict",
config_overrides={"stage1_model.hidden_dim": 999},
)
def test_config_override_bad_value_raises(tmp_path):
checkpoint = _write_checkpoint(tmp_path)
with pytest.raises(CheckpointCompatibilityError, match="must be one of"):
load_for_inference(
checkpoint,
torch.device("cpu"),
"predict",
config_overrides={"stage2_model.n_sec.sampling": "maybe"},
)
def test_apply_config_overrides_no_overrides_returns_same_object():
cfg = _model_cfg()
assert apply_config_overrides(cfg, None) is cfg
assert apply_config_overrides(cfg, {}) is cfg
def test_apply_config_overrides_migrates_legacy_flat_model_config_first():
legacy_cfg = {
"pdg_vocab": len(PDG_MAP),
"mat_vocab": len(MAT_MAP),
"hidden_dim": 32,
"n_blocks": 4,
"emb_dim": 8,
"dropout": 0.1,
"k_max": 5,
}
merged = apply_config_overrides(legacy_cfg, {"stage1_model.ddpm.n_steps": 10})
assert merged["stage1_model"]["ddpm"]["n_steps"] == 10
assert merged["stage1_model"]["hidden_dim"] == 32
+40
View File
@@ -172,3 +172,43 @@ def test_predict_exits_1_on_checkpoint_missing_model_config(tmp_path):
assert result.exit_code == 1 assert result.exit_code == 1
assert "checkpoint has no model_config" in result.output assert "checkpoint has no model_config" in result.output
# ---------------------------------------------------------------------------
# --set (gitea #87)
# ---------------------------------------------------------------------------
def test_predict_set_flag_without_equals_exits_1(tmp_path):
checkpoint = tmp_path / "missing.pt"
result = runner.invoke(
app,
["predict", "dummy.parquet", "--checkpoint", str(checkpoint), "--set", "sampling"],
)
assert result.exit_code == 1
assert "must be 'dotted.path=value'" in result.output
def test_predict_set_flag_disallowed_path_surfaces_compat_error(tmp_path):
checkpoint = tmp_path / "ckpt.pt"
torch.save(
{"model_config": {"stage1_model": {}, "stage2_model": {}}, "sec_decoder": {}, "normalizer": {"sec_phys": {}}},
checkpoint,
)
result = runner.invoke(
app,
[
"predict",
"dummy.parquet",
"--checkpoint",
str(checkpoint),
"--set",
"stage1_model.hidden_dim=999",
],
)
assert result.exit_code == 1
assert "not an inference-safe override" in result.output
+25
View File
@@ -31,3 +31,28 @@ def test_rollout_exits_1_on_checkpoint_missing_model_config(tmp_path):
assert result.exit_code == 1 assert result.exit_code == 1
assert "checkpoint has no model_config" in result.output assert "checkpoint has no model_config" in result.output
def test_rollout_set_flag_disallowed_path_surfaces_compat_error(tmp_path):
checkpoint = tmp_path / "ckpt.pt"
torch.save(
{"model_config": {"stage1_model": {}, "stage2_model": {}}, "sec_decoder": {}, "normalizer": {"sec_phys": {}}},
checkpoint,
)
result = runner.invoke(
app,
[
"rollout",
"dummy.parquet",
"--checkpoint",
str(checkpoint),
"--geometry",
"dummy_geometry.pkl",
"--set",
"stage2_model.n_sec.typo=sample",
],
)
assert result.exit_code == 1
assert "not an inference-safe override" in result.output
+32 -9
View File
@@ -171,17 +171,40 @@ def test_n_sec_config_owner_defaults_to_stage2():
def test_n_sec_config_owner_round_trips(): def test_n_sec_config_owner_round_trips():
n_sec = gconfig.NSecConfig.from_dict({"mode": "head", "owner": "stage1"}) n_sec = gconfig.NSecConfig.from_dict({"mode": "head", "owner": "stage1"})
assert n_sec.owner == "stage1" assert n_sec.owner == "stage1"
assert n_sec.to_dict() == {"mode": "head", "lambda": 0.1, "owner": "stage1", "stop_sampling": "greedy"} assert n_sec.to_dict() == {"mode": "head", "lambda": 0.1, "owner": "stage1", "sampling": "greedy"}
def test_n_sec_config_stop_sampling_defaults_to_greedy(): def test_n_sec_config_sampling_defaults_to_greedy():
assert gconfig.NSecConfig().stop_sampling == "greedy" assert gconfig.NSecConfig().sampling == "greedy"
def test_n_sec_config_stop_sampling_round_trips(): def test_n_sec_config_sampling_round_trips():
n_sec = gconfig.NSecConfig.from_dict({"mode": "stop_token", "sampling": "sample"})
assert n_sec.sampling == "sample"
assert n_sec.to_dict()["sampling"] == "sample"
def test_n_sec_config_stop_sampling_alias_still_honored():
"""gitea #86: stop_sampling was renamed to sampling; old checkpoints'
model_config still carries the old key and must keep working."""
n_sec = gconfig.NSecConfig.from_dict({"mode": "stop_token", "stop_sampling": "sample"}) n_sec = gconfig.NSecConfig.from_dict({"mode": "stop_token", "stop_sampling": "sample"})
assert n_sec.stop_sampling == "sample" assert n_sec.sampling == "sample"
assert n_sec.to_dict()["stop_sampling"] == "sample" assert "stop_sampling" not in n_sec.to_dict()
def test_n_sec_config_sampling_key_wins_over_stop_sampling_alias():
n_sec = gconfig.NSecConfig.from_dict({"sampling": "sample", "stop_sampling": "greedy"})
assert n_sec.sampling == "sample"
def test_migrate_config_renames_stop_sampling_key():
cfg = {
"meta": {"config_version": gconfig.CONFIG_VERSION},
"stage2_model": {"n_sec": {"stop_sampling": "sample"}},
}
migrated = gconfig.migrate_config(cfg)
assert gconfig._get_path(migrated, "stage2_model.n_sec.sampling") == "sample"
assert gconfig._get_path(migrated, "stage2_model.n_sec.stop_sampling") is None
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
@@ -851,13 +874,13 @@ def test_validate_config_stop_token_rejected_for_stage1_owner():
assert "stop_token" in str(e) and "owner" in str(e) assert "stop_token" in str(e) and "owner" in str(e)
def test_validate_config_bad_stop_sampling_rejected(): def test_validate_config_bad_n_sec_sampling_rejected():
cfg = _cfg_with(**{"stage2_model.n_sec.stop_sampling": "bogus"}) cfg = _cfg_with(**{"stage2_model.n_sec.sampling": "bogus"})
try: try:
gconfig.validate_config(cfg) gconfig.validate_config(cfg)
assert False, "expected ValueError" assert False, "expected ValueError"
except ValueError as e: except ValueError as e:
assert "stop_sampling" in str(e) assert "sampling" in str(e)
def test_validate_config_default_precision_is_fp32(): def test_validate_config_default_precision_is_fp32():
+14
View File
@@ -12,6 +12,8 @@ from giant.cli import app
from giant.materials import MATERIAL_PROPERTIES from giant.materials import MATERIAL_PROPERTIES
from giant.model.summary import _NOT_BUILD_TIME, _built_modules, _vocab_caveats, summarize_model from giant.model.summary import _NOT_BUILD_TIME, _built_modules, _vocab_caveats, summarize_model
INFERENCE_OVERRIDES = gconfig.INFERENCE_OVERRIDES
runner = CliRunner() runner = CliRunner()
_PDG_VOCAB = 300 _PDG_VOCAB = 300
@@ -53,6 +55,16 @@ def test_not_build_time_allow_list_has_no_stale_entries():
assert not stale, f"_NOT_BUILD_TIME entries no longer in DEFAULT_CONFIG: {sorted(stale)}" assert not stale, f"_NOT_BUILD_TIME entries no longer in DEFAULT_CONFIG: {sorted(stale)}"
def test_inference_overrides_allow_list_has_no_stale_entries():
in_scope = set(gconfig.leaf_paths(gconfig.DEFAULT_CONFIG))
stale = set(INFERENCE_OVERRIDES) - in_scope
assert not stale, f"INFERENCE_OVERRIDES entries no longer in DEFAULT_CONFIG: {sorted(stale)}"
def test_default_config_overridable_lists_every_allowlisted_path(default_summary):
assert set(default_summary.overridable) == set(INFERENCE_OVERRIDES)
def test_router_disabled_by_default_so_its_fields_are_inert(default_summary): 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.n_experts" in default_summary.inert
assert "stage1_model.router.temperature" in default_summary.inert assert "stage1_model.router.temperature" in default_summary.inert
@@ -135,3 +147,5 @@ def test_cli_default_smoke():
assert "parameters" in result.output assert "parameters" in result.output
assert "trunk" in result.output assert "trunk" in result.output
assert "inert under this config" in result.output assert "inert under this config" in result.output
assert "inference-overridable without retraining" in result.output
assert "stage2_model.n_sec.sampling" in result.output
+68 -8
View File
@@ -14,6 +14,7 @@ from giant.model.network import (
stage2_trunk_sec_dim, stage2_trunk_sec_dim,
) )
from giant.sample import ( from giant.sample import (
resolve_n_sec,
sample_flow, sample_flow,
sample_secondaries, sample_secondaries,
sample_secondaries_ar, sample_secondaries_ar,
@@ -72,6 +73,7 @@ def _stage2_ar(
mat: int = 2, mat: int = 2,
k_max: int = 5, k_max: int = 5,
history: str = "markov", history: str = "markov",
n_sec_sampling: str = "greedy",
) -> Stage2Autoregressive: ) -> Stage2Autoregressive:
particle_cfg, material_cfg = _particle_material_cfg(_conditioning_for(target), emb_dim) particle_cfg, material_cfg = _particle_material_cfg(_conditioning_for(target), emb_dim)
return Stage2Autoregressive( return Stage2Autoregressive(
@@ -89,6 +91,7 @@ def _stage2_ar(
history=history, history=history,
attn_n_heads=2, attn_n_heads=2,
attn_n_layers=1, attn_n_layers=1,
n_sec_sampling=n_sec_sampling,
).eval() ).eval()
@@ -99,7 +102,7 @@ def _expected_type_dim(target: str, emb_dim: int) -> int:
def _stage2_ar_stop_token( def _stage2_ar_stop_token(
target: str, target: str,
generator: str, generator: str,
stop_sampling: str = "greedy", n_sec_sampling: str = "greedy",
emb_dim: int = 6, emb_dim: int = 6,
pdg: int = 3, pdg: int = 3,
mat: int = 2, mat: int = 2,
@@ -120,7 +123,7 @@ def _stage2_ar_stop_token(
particle_type_cfg=ParticleTypeConfig(target=target), particle_type_cfg=ParticleTypeConfig(target=target),
build_n_sec_head=False, build_n_sec_head=False,
build_stop_head=True, build_stop_head=True,
stop_sampling=stop_sampling, n_sec_sampling=n_sec_sampling,
).eval() ).eval()
@@ -267,14 +270,14 @@ def test_sample_secondaries_ar_first_slot_has_no_history():
# ── Stage2Autoregressive: n_sec.mode = "stop_token" ───────────────────────── # ── Stage2Autoregressive: n_sec.mode = "stop_token" ─────────────────────────
@pytest.mark.parametrize("stop_sampling", ["greedy", "sample"]) @pytest.mark.parametrize("n_sec_sampling", ["greedy", "sample"])
def test_sample_secondaries_ar_stop_token_forced_stop_gives_zero_secondaries(stop_sampling): def test_sample_secondaries_ar_stop_token_forced_stop_gives_zero_secondaries(n_sec_sampling):
"""A stop_head pinned to a large positive logit fires at slot 0 for """A stop_head pinned to a large positive logit fires at slot 0 for
every row under both policies (greedy: sigmoid(logit) >= 0.5; sample: every row under both policies (greedy: sigmoid(logit) >= 0.5; sample:
a Bernoulli draw at sigmoid(logit) ~= 1) the loop should break before a Bernoulli draw at sigmoid(logit) ~= 1) the loop should break before
generating any token.""" generating any token."""
B, k_max = 4, 5 B, k_max = 4, 5
decoder = _stage2_ar_stop_token("physical", "flow", stop_sampling=stop_sampling, k_max=k_max) decoder = _stage2_ar_stop_token("physical", "flow", n_sec_sampling=n_sec_sampling, k_max=k_max)
_force_stop_head_logit(decoder, 50.0) _force_stop_head_logit(decoder, 50.0)
cond_cont, cond_cat = _cond(B) cond_cont, cond_cat = _cond(B)
stage1_out = torch.randn(B, X_DIM) stage1_out = torch.randn(B, X_DIM)
@@ -283,13 +286,13 @@ def test_sample_secondaries_ar_stop_token_forced_stop_gives_zero_secondaries(sto
assert not sec_valid.any() assert not sec_valid.any()
@pytest.mark.parametrize("stop_sampling", ["greedy", "sample"]) @pytest.mark.parametrize("n_sec_sampling", ["greedy", "sample"])
def test_sample_secondaries_ar_stop_token_forced_never_stop_runs_to_k_max(stop_sampling): def test_sample_secondaries_ar_stop_token_forced_never_stop_runs_to_k_max(n_sec_sampling):
"""A stop_head pinned to a large negative logit never fires under either """A stop_head pinned to a large negative logit never fires under either
policy, so every row is capped at k_max (the safety cap, not a modeling policy, so every row is capped at k_max (the safety cap, not a modeling
ceiling).""" ceiling)."""
B, k_max = 4, 5 B, k_max = 4, 5
decoder = _stage2_ar_stop_token("physical", "flow", stop_sampling=stop_sampling, k_max=k_max) decoder = _stage2_ar_stop_token("physical", "flow", n_sec_sampling=n_sec_sampling, k_max=k_max)
_force_stop_head_logit(decoder, -50.0) _force_stop_head_logit(decoder, -50.0)
cond_cont, cond_cat = _cond(B) cond_cont, cond_cat = _cond(B)
stage1_out = torch.randn(B, X_DIM) stage1_out = torch.randn(B, X_DIM)
@@ -336,3 +339,60 @@ def test_sample_secondaries_ar_none_n_sec_pred_without_stop_head_raises():
stage1_out = torch.randn(3, X_DIM) stage1_out = torch.randn(3, X_DIM)
with pytest.raises(AssertionError): with pytest.raises(AssertionError):
sample_secondaries_ar(decoder, cond_cont, cond_cat, stage1_out, None, steps=2) sample_secondaries_ar(decoder, cond_cont, cond_cat, stage1_out, None, steps=2)
# ── resolve_n_sec: n_sec.mode = "head" sampling policy (gitea #86) ──────────
def _force_n_sec_head_bias(decoder: Stage2Autoregressive, bias: torch.Tensor) -> None:
"""Zeroes n_sec_head's weights and pins its bias, so predict_n_sec
returns `bias` (broadcast over the batch) as logits regardless of
conditioning mirrors `_force_stop_head_logit`."""
assert decoder.n_sec_head is not None
last_linear = decoder.n_sec_head[-1]
with torch.no_grad():
last_linear.weight.zero_()
last_linear.bias.copy_(bias)
@pytest.mark.parametrize("n_sec_sampling", ["greedy", "sample"])
def test_resolve_n_sec_head_mode_sharply_peaked_logits_pick_dominant_class(n_sec_sampling):
"""A logit vector overwhelmingly favoring one class gives the same
answer under both policies greedy because it's the argmax, sample
because softmax puts ~all mass on it."""
B, k_max = 8, 5
decoder = _stage2_ar("physical", "flow", k_max=k_max, n_sec_sampling=n_sec_sampling)
bias = torch.full((k_max + 1,), -50.0)
bias[2] = 50.0
_force_n_sec_head_bias(decoder, bias)
cond_cont, cond_cat = _cond(B)
stage1_out = torch.randn(B, X_DIM)
n_sec = resolve_n_sec(decoder, decoder, cond_cont, cond_cat, stage1_out, None)
assert n_sec is not None
assert torch.equal(n_sec, torch.full((B,), 2, dtype=torch.long))
def test_resolve_n_sec_head_mode_greedy_is_deterministic_under_flat_logits():
B, k_max = 32, 5
decoder = _stage2_ar("physical", "flow", k_max=k_max, n_sec_sampling="greedy")
_force_n_sec_head_bias(decoder, torch.zeros(k_max + 1))
cond_cont, cond_cat = _cond(B)
stage1_out = torch.randn(B, X_DIM)
n_sec = resolve_n_sec(decoder, decoder, cond_cont, cond_cat, stage1_out, None)
assert n_sec is not None
assert n_sec.unique().numel() == 1
def test_resolve_n_sec_head_mode_sample_varies_under_flat_logits():
"""Under a flat logit vector, a categorical draw across a large batch
should hit more than one class the whole point of gitea #86: greedy
always collapses to one, sample should not."""
torch.manual_seed(0)
B, k_max = 256, 5
decoder = _stage2_ar("physical", "flow", k_max=k_max, n_sec_sampling="sample")
_force_n_sec_head_bias(decoder, torch.zeros(k_max + 1))
cond_cont, cond_cat = _cond(B)
stage1_out = torch.randn(B, X_DIM)
n_sec = resolve_n_sec(decoder, decoder, cond_cont, cond_cat, stage1_out, None)
assert n_sec is not None
assert n_sec.unique().numel() > 1
Generated
+1 -1
View File
@@ -675,7 +675,7 @@ wheels = [
[[package]] [[package]]
name = "giant" name = "giant"
version = "0.3.11" version = "0.3.13"
source = { editable = "." } source = { editable = "." }
dependencies = [ dependencies = [
{ name = "numpy" }, { name = "numpy" },