Fix conditioning="physical" so it can actually generalize past training vocab
The whole point of conditioning="physical" is generalizing to a species/material outside the training menu, but two independent code paths still hard-required training-vocab membership: - giant/data/transforms.py: build_cond_features unconditionally raised KeyError on an out-of-vocab pdg/material. _vectorized_map_lookup gains a strict=False mode (dummy index instead of raising), used only under conditioning="physical" where ConditionEncoder never reads cond_cat anyway; "embedding" mode is untouched and still raises, since cond_cat IS the conditioning signal there. - giant/rollout.py: the known_pdg termination gate still killed a track on step 1 for any pdg outside pdg_map, regardless of conditioning mode. Now skipped entirely under conditioning="physical". - giant/model/network.py: PdgRouter/ProcessRouter always build their own training-vocab nn.Embedding independent of conditioning, silently reintroducing the same limitation at the routing layer. build_models now raises loudly if conditioning="physical" is paired with either router type, rather than silently building a model that can't generalize the way it claims to. This unblocks the held-out-species/material generalization experiment against the multi-material dataset (see CLAUDE.md roadmap). Each fix has a regression test, including an end-to-end rollout test seeded with a resolvable-but-out-of-vocab PDG code. Co-Authored-By: Claude Sonnet 5 <noreply@anthropic.com>
This commit is contained in:
@@ -339,12 +339,22 @@ def sorted_membership(values: np.ndarray, sorted_arr: np.ndarray) -> np.ndarray:
|
|||||||
return sorted_arr[idx] == values
|
return sorted_arr[idx] == values
|
||||||
|
|
||||||
|
|
||||||
def _vectorized_map_lookup(values: np.ndarray, mapping: dict) -> np.ndarray:
|
def _vectorized_map_lookup(
|
||||||
|
values: np.ndarray, mapping: dict, strict: bool = True
|
||||||
|
) -> np.ndarray:
|
||||||
"""Vectorized equivalent of `np.array([mapping[v] for v in values], dtype=np.int64)`.
|
"""Vectorized equivalent of `np.array([mapping[v] for v in values], dtype=np.int64)`.
|
||||||
|
|
||||||
Replaces a per-element Python dict lookup with one `searchsorted` call.
|
Replaces a per-element Python dict lookup with one `searchsorted` call.
|
||||||
Raises `KeyError` if any value in `values` isn't a key of `mapping`,
|
Raises `KeyError` if any value in `values` isn't a key of `mapping`,
|
||||||
matching the dict-comprehension it replaces (never silently misassigns).
|
matching the dict-comprehension it replaces (never silently misassigns)
|
||||||
|
— unless `strict=False`, in which case unmapped values get a dummy index
|
||||||
|
of 0 instead. Only pass `strict=False` where the caller has independently
|
||||||
|
verified the resulting index is never actually read (e.g.
|
||||||
|
`build_cond_features` under `conditioning="physical"`, where
|
||||||
|
`ConditionEncoder` ignores `cond_cat` entirely); it exists so a rollout
|
||||||
|
can be seeded with a species/material outside the training vocab without
|
||||||
|
a spurious `KeyError`, which is the entire point of physical-property
|
||||||
|
conditioning.
|
||||||
"""
|
"""
|
||||||
keys = np.asarray(list(mapping.keys()))
|
keys = np.asarray(list(mapping.keys()))
|
||||||
vals = np.asarray(list(mapping.values()), dtype=np.int64)
|
vals = np.asarray(list(mapping.values()), dtype=np.int64)
|
||||||
@@ -355,6 +365,10 @@ def _vectorized_map_lookup(values: np.ndarray, mapping: dict) -> np.ndarray:
|
|||||||
pos = np.clip(pos, 0, len(keys_sorted) - 1)
|
pos = np.clip(pos, 0, len(keys_sorted) - 1)
|
||||||
found = keys_sorted[pos] == values
|
found = keys_sorted[pos] == values
|
||||||
if not found.all():
|
if not found.all():
|
||||||
|
if not strict:
|
||||||
|
out = np.zeros(values.shape, dtype=np.int64)
|
||||||
|
out[found] = vals_sorted[pos[found]]
|
||||||
|
return out
|
||||||
missing = np.unique(values[~found])
|
missing = np.unique(values[~found])
|
||||||
raise KeyError(f"value(s) not in mapping: {missing[:10].tolist()}")
|
raise KeyError(f"value(s) not in mapping: {missing[:10].tolist()}")
|
||||||
return vals_sorted[pos]
|
return vals_sorted[pos]
|
||||||
@@ -693,8 +707,15 @@ def build_cond_features(
|
|||||||
[cond_cont, _physical_cond_columns(data, conditioning)]
|
[cond_cont, _physical_cond_columns(data, conditioning)]
|
||||||
).astype(np.float32)
|
).astype(np.float32)
|
||||||
|
|
||||||
pdg_idx = _vectorized_map_lookup(data["pdg"], pdg_map)
|
# In "physical" mode cond_cat is only a reporting/router convenience —
|
||||||
mat_idx = _vectorized_map_lookup(data["material"], mat_map)
|
# ConditionEncoder never reads it (giant/model/network.py) — so a
|
||||||
|
# species/material outside the training vocab (the whole point of
|
||||||
|
# physical-property conditioning) gets a dummy index instead of raising.
|
||||||
|
# In "embedding" mode cond_cat IS the conditioning signal, so an unmapped
|
||||||
|
# value must still raise loudly rather than silently misassign.
|
||||||
|
strict = conditioning == "embedding"
|
||||||
|
pdg_idx = _vectorized_map_lookup(data["pdg"], pdg_map, strict=strict)
|
||||||
|
mat_idx = _vectorized_map_lookup(data["material"], mat_map, strict=strict)
|
||||||
cond_cat = np.column_stack([pdg_idx, mat_idx])
|
cond_cat = np.column_stack([pdg_idx, mat_idx])
|
||||||
|
|
||||||
if cond_normalizer is not None:
|
if cond_normalizer is not None:
|
||||||
|
|||||||
+45
-4
@@ -1228,7 +1228,40 @@ def _parse_composed_axes(router_cfg: dict) -> list[dict]:
|
|||||||
return [axes[i] for i in range(len(axes))]
|
return [axes[i] for i in range(len(axes))]
|
||||||
|
|
||||||
|
|
||||||
def _build_router_from_cfg(router_cfg: dict, pdg_vocab: int, mat_vocab: int) -> Router:
|
# Router types that read cond_cat's pdg index through their own
|
||||||
|
# nn.Embedding(pdg_vocab, ...), regardless of the trunk's `conditioning`
|
||||||
|
# mode — see _check_router_conditioning_compat.
|
||||||
|
_VOCAB_SCOPED_ROUTER_TYPES = ("pdg", "process")
|
||||||
|
|
||||||
|
|
||||||
|
def _check_router_conditioning_compat(router_types: list[str], conditioning: str) -> None:
|
||||||
|
"""Reject a router axis that reintroduces a training-vocab PDG lookup
|
||||||
|
under `conditioning="physical"`.
|
||||||
|
|
||||||
|
`PdgRouter`/`ProcessRouter` always build their own dataset-scoped
|
||||||
|
`nn.Embedding(pdg_vocab, ...)` (network.py's PdgRouter/ProcessRouter),
|
||||||
|
independent of `ConditionEncoder`'s `conditioning` mode. Pairing either
|
||||||
|
with `conditioning="physical"` would silently reintroduce a
|
||||||
|
training-menu-scoped lookup at the routing layer — defeating the entire
|
||||||
|
point of physical-property conditioning, which is to generalize to a
|
||||||
|
species/material outside that menu (see giant/rollout.py's
|
||||||
|
`build_cond_features(strict=...)` gate for the same concern on the
|
||||||
|
trunk side). Raised loudly at model-build time rather than left to
|
||||||
|
surface as a confusing rollout/generalization-benchmark result.
|
||||||
|
"""
|
||||||
|
bad = sorted(set(router_types) & set(_VOCAB_SCOPED_ROUTER_TYPES))
|
||||||
|
if bad and conditioning == "physical":
|
||||||
|
raise ValueError(
|
||||||
|
f"router type(s) {bad} always use a training-vocab PDG embedding, "
|
||||||
|
"which is incompatible with conditioning='physical' (whose whole "
|
||||||
|
"point is generalizing beyond that vocab) — pick a different "
|
||||||
|
"router type (e.g. 'energy') or use conditioning='embedding'."
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _build_router_from_cfg(
|
||||||
|
router_cfg: dict, pdg_vocab: int, mat_vocab: int, conditioning: str = "embedding"
|
||||||
|
) -> Router:
|
||||||
"""Resolve one `model.router` config into a `Router`, single-axis or composed.
|
"""Resolve one `model.router` config into a `Router`, single-axis or composed.
|
||||||
|
|
||||||
`router_cfg["type"] == "composed"` reads `axis{i}_{field}` flat keys
|
`router_cfg["type"] == "composed"` reads `axis{i}_{field}` flat keys
|
||||||
@@ -1243,9 +1276,12 @@ def _build_router_from_cfg(router_cfg: dict, pdg_vocab: int, mat_vocab: int) ->
|
|||||||
"""
|
"""
|
||||||
shared_vocab = dict(pdg_vocab=pdg_vocab, mat_vocab=mat_vocab)
|
shared_vocab = dict(pdg_vocab=pdg_vocab, mat_vocab=mat_vocab)
|
||||||
if router_cfg["type"] == "composed":
|
if router_cfg["type"] == "composed":
|
||||||
router = build_composed_router(_parse_composed_axes(router_cfg), **shared_vocab)
|
axes = _parse_composed_axes(router_cfg)
|
||||||
|
_check_router_conditioning_compat([a["type"] for a in axes], conditioning)
|
||||||
|
router = build_composed_router(axes, **shared_vocab)
|
||||||
router.gumbel = bool(router_cfg.get("gumbel", False))
|
router.gumbel = bool(router_cfg.get("gumbel", False))
|
||||||
return router
|
return router
|
||||||
|
_check_router_conditioning_compat([router_cfg["type"]], conditioning)
|
||||||
router_kwargs = {
|
router_kwargs = {
|
||||||
k: v for k, v in router_cfg.items() if k not in ("enabled", "type", "n_experts")
|
k: v for k, v in router_cfg.items() if k not in ("enabled", "type", "n_experts")
|
||||||
}
|
}
|
||||||
@@ -1303,13 +1339,18 @@ def build_models(model_config: dict) -> tuple[nn.Module, nn.Module]:
|
|||||||
dropout=model_config.get("dropout", 0.1),
|
dropout=model_config.get("dropout", 0.1),
|
||||||
conditioning=model_config.get("conditioning", "embedding"),
|
conditioning=model_config.get("conditioning", "embedding"),
|
||||||
)
|
)
|
||||||
|
conditioning = shared["conditioning"]
|
||||||
stage1 = RoutedDenoisingMLP(
|
stage1 = RoutedDenoisingMLP(
|
||||||
router=_build_router_from_cfg(router_cfg, pdg_vocab, mat_vocab),
|
router=_build_router_from_cfg(
|
||||||
|
router_cfg, pdg_vocab, mat_vocab, conditioning
|
||||||
|
),
|
||||||
k_max=model_config.get("k_max", K_MAX),
|
k_max=model_config.get("k_max", K_MAX),
|
||||||
**shared,
|
**shared,
|
||||||
)
|
)
|
||||||
sec_decoder = RoutedSecondaryDecoder(
|
sec_decoder = RoutedSecondaryDecoder(
|
||||||
router=_build_router_from_cfg(router_cfg, pdg_vocab, mat_vocab),
|
router=_build_router_from_cfg(
|
||||||
|
router_cfg, pdg_vocab, mat_vocab, conditioning
|
||||||
|
),
|
||||||
**shared,
|
**shared,
|
||||||
)
|
)
|
||||||
return stage1, sec_decoder
|
return stage1, sec_decoder
|
||||||
|
|||||||
+10
-1
@@ -407,7 +407,16 @@ def _step_chunk(
|
|||||||
tr["_material"] = material
|
tr["_material"] = material
|
||||||
tr["_layer_id"] = layer_id
|
tr["_layer_id"] = layer_id
|
||||||
|
|
||||||
known_pdg = np.array([int(p) in pdg_map for p in tr["pdg"]], dtype=bool)
|
if conditioning == "physical":
|
||||||
|
# Under physical-property conditioning, mass/charge (already resolved
|
||||||
|
# on every track — see the cond_dict comment below) drive the model,
|
||||||
|
# not a training-vocab PDG embedding — build_cond_features passes
|
||||||
|
# strict=False for exactly this mode, so an out-of-vocab species no
|
||||||
|
# longer raises. Terminating on it here would defeat the entire
|
||||||
|
# point of physical conditioning: generalizing to a held-out species.
|
||||||
|
known_pdg = np.ones(n, dtype=bool)
|
||||||
|
else:
|
||||||
|
known_pdg = np.array([int(p) in pdg_map for p in tr["pdg"]], dtype=bool)
|
||||||
|
|
||||||
# --- Pre-step termination gates (in priority order; each track picks one) ---
|
# --- Pre-step termination gates (in priority order; each track picks one) ---
|
||||||
stop = np.zeros(n, dtype=bool)
|
stop = np.zeros(n, dtype=bool)
|
||||||
|
|||||||
@@ -119,6 +119,21 @@ def test_rollout_physical_conditioning_end_to_end(fake_material_props):
|
|||||||
assert set(np.unique(rec["pdg"]).tolist()) <= set(PDG_MAP.keys())
|
assert set(np.unique(rec["pdg"]).tolist()) <= set(PDG_MAP.keys())
|
||||||
|
|
||||||
|
|
||||||
|
def test_rollout_physical_conditioning_generalizes_to_out_of_vocab_pdg(
|
||||||
|
fake_material_props,
|
||||||
|
):
|
||||||
|
"""A real, giant.particles-resolvable species outside the training PDG
|
||||||
|
vocab (muon, 13) must run through physical-property conditioning rather
|
||||||
|
than terminate via TERM_UNKNOWN_PDG — that generalization is the entire
|
||||||
|
point of "physical" mode (see build_cond_features(strict=...))."""
|
||||||
|
seeds = _seeds(6)
|
||||||
|
seeds["pdg"] = np.full(6, 13, dtype=np.int64)
|
||||||
|
assert 13 not in PDG_MAP
|
||||||
|
rec = _run(seeds=seeds, conditioning="physical")
|
||||||
|
assert len(rec["event_id"]) > 0
|
||||||
|
assert TERM_UNKNOWN_PDG not in set(rec["termination_reason"].tolist())
|
||||||
|
|
||||||
|
|
||||||
def test_seed_frontier_track_ids():
|
def test_seed_frontier_track_ids():
|
||||||
seeds = _seeds(3)
|
seeds = _seeds(3)
|
||||||
fr, counts = make_seed_frontier(**seeds)
|
fr, counts = make_seed_frontier(**seeds)
|
||||||
|
|||||||
@@ -1,5 +1,6 @@
|
|||||||
"""Tests for the mixture-of-experts routing prototype (giant/model/network.py)."""
|
"""Tests for the mixture-of-experts routing prototype (giant/model/network.py)."""
|
||||||
|
|
||||||
|
import pytest
|
||||||
import torch
|
import torch
|
||||||
|
|
||||||
from giant.constants import COND_DIM, K_MAX, SEC_DIM, X_DIM
|
from giant.constants import COND_DIM, K_MAX, SEC_DIM, X_DIM
|
||||||
@@ -484,6 +485,49 @@ def test_build_models_routed_with_pdg_router():
|
|||||||
assert stage1.router.pdg_emb.num_embeddings == 4
|
assert stage1.router.pdg_emb.num_embeddings == 4
|
||||||
|
|
||||||
|
|
||||||
|
def test_build_models_rejects_pdg_router_with_physical_conditioning():
|
||||||
|
"""conditioning="physical" is meant to generalize beyond the training PDG
|
||||||
|
vocab; PdgRouter always uses a training-vocab nn.Embedding regardless of
|
||||||
|
conditioning, so the combination must raise rather than silently building
|
||||||
|
a model that can't actually generalize the way it claims to."""
|
||||||
|
model_config = dict(
|
||||||
|
pdg_vocab=4,
|
||||||
|
mat_vocab=2,
|
||||||
|
emb_dim=16,
|
||||||
|
dropout=0.1,
|
||||||
|
k_max=K_MAX,
|
||||||
|
expert_hidden_dim=16,
|
||||||
|
expert_n_blocks=2,
|
||||||
|
conditioning="physical",
|
||||||
|
router={"enabled": True, "type": "pdg", "n_experts": 3},
|
||||||
|
)
|
||||||
|
with pytest.raises(ValueError, match="physical"):
|
||||||
|
build_models(model_config)
|
||||||
|
|
||||||
|
|
||||||
|
def test_build_models_rejects_composed_router_with_pdg_axis_and_physical_conditioning():
|
||||||
|
model_config = dict(
|
||||||
|
pdg_vocab=4,
|
||||||
|
mat_vocab=2,
|
||||||
|
emb_dim=16,
|
||||||
|
dropout=0.1,
|
||||||
|
k_max=K_MAX,
|
||||||
|
expert_hidden_dim=16,
|
||||||
|
expert_n_blocks=2,
|
||||||
|
conditioning="physical",
|
||||||
|
router={
|
||||||
|
"enabled": True,
|
||||||
|
"type": "composed",
|
||||||
|
"axis0_type": "energy",
|
||||||
|
"axis0_n_experts": 2,
|
||||||
|
"axis1_type": "pdg",
|
||||||
|
"axis1_n_experts": 3,
|
||||||
|
},
|
||||||
|
)
|
||||||
|
with pytest.raises(ValueError, match="physical"):
|
||||||
|
build_models(model_config)
|
||||||
|
|
||||||
|
|
||||||
# ── ProcessRouter ────────────────────────────────────────────────────────────
|
# ── ProcessRouter ────────────────────────────────────────────────────────────
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -574,6 +574,51 @@ def test_vectorized_map_lookup_raises_keyerror_on_missing_value():
|
|||||||
_vectorized_map_lookup(values, mapping)
|
_vectorized_map_lookup(values, mapping)
|
||||||
|
|
||||||
|
|
||||||
|
def test_vectorized_map_lookup_strict_false_dummy_indexes_unmapped_values():
|
||||||
|
"""strict=False must leave found values untouched and only dummy-index
|
||||||
|
(0) the unmapped ones — never raise, and never disturb a value that IS
|
||||||
|
in the mapping (e.g. one that happens to map to a nonzero index)."""
|
||||||
|
mapping = {1: 5, 2: 7}
|
||||||
|
values = np.array([1, 99, 2, 100])
|
||||||
|
result = _vectorized_map_lookup(values, mapping, strict=False)
|
||||||
|
np.testing.assert_array_equal(result, [5, 0, 7, 0])
|
||||||
|
|
||||||
|
|
||||||
|
def test_build_cond_features_physical_mode_tolerates_out_of_vocab_pdg_and_material():
|
||||||
|
"""conditioning="physical" must not KeyError on a pdg/material outside
|
||||||
|
the training-dataset vocab (mat_map/pdg_map) — that's the entire point
|
||||||
|
of the mode (see giant.rollout's known_pdg gate for the paired fix).
|
||||||
|
"embedding" mode must still raise, since cond_cat IS the conditioning
|
||||||
|
signal there. Note this is specifically about the dataset-scoped
|
||||||
|
vocab index, not giant.materials' physical-properties table — a
|
||||||
|
material must still be a real, known Geant4 material (e.g. "G4_Pb",
|
||||||
|
just not one *this* mat_map happened to include) for "physical" mode
|
||||||
|
to derive its Z_eff/A_eff/density/X0/λ_int; a genuinely unknown
|
||||||
|
material name correctly still raises via giant.materials, same as the
|
||||||
|
documented G4_LYSO precedent — that's a separate, intentional guard."""
|
||||||
|
pdg_map = {11: 0, 22: 1}
|
||||||
|
mat_map = {"G4_AIR": 0}
|
||||||
|
data = {
|
||||||
|
"pre_pos": np.zeros((1, 3), dtype=np.float32),
|
||||||
|
"pre_E": np.array([10.0], dtype=np.float32),
|
||||||
|
"pre_dir": np.array([[0.0, 0.0, 1.0]], dtype=np.float32),
|
||||||
|
"layer_id": np.array([0], dtype=np.int32),
|
||||||
|
"pdg": np.array([13], dtype=np.int64), # not in pdg_map
|
||||||
|
"material": np.array(["G4_Pb"], dtype=object), # not in mat_map
|
||||||
|
"mass": np.array([105.7], dtype=np.float32),
|
||||||
|
"charge": np.array([-1.0], dtype=np.float32),
|
||||||
|
}
|
||||||
|
|
||||||
|
cond_cont, cond_cat = build_cond_features(
|
||||||
|
data, pdg_map, mat_map, conditioning="physical"
|
||||||
|
)
|
||||||
|
assert cond_cont.shape[-1] == COND_DIM
|
||||||
|
np.testing.assert_array_equal(cond_cat, [[0, 0]]) # dummy indices, no raise
|
||||||
|
|
||||||
|
with pytest.raises(KeyError):
|
||||||
|
build_cond_features(data, pdg_map, mat_map, conditioning="embedding")
|
||||||
|
|
||||||
|
|
||||||
# ── _WelfordAccumulator ──────────────────────────────────────────────────────
|
# ── _WelfordAccumulator ──────────────────────────────────────────────────────
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user