"""Single source of truth for the conditioning arrays' column layout (gitea #37). `cond_cont` and `cond_cat` are built in `giant.data.transforms` and consumed in `giant.model.encoders` / `giant.model.routers`. Their column order used to be written down independently on each side, kept in sync only by parallel comments — so getting it wrong produced silently mis-indexed columns rather than an exception, and adding a conditioning axis meant a coordinated multi-file edit. `CondLayout` owns that order. Both sides construct one from the same `conditioning.particle.type` / `conditioning.material.type` pair and read named slices off it, so the layout is stated exactly once. This module depends only on `giant.constants`, so both the data and model packages can import it. """ from dataclasses import dataclass from typing import ClassVar from giant.constants import COND_DIM, COND_DIM_BASE, MATERIAL_PHYS_DIM, PARTICLE_PHYS_DIM # The three per-axis conditioning modes. Mirrors giant.config.Conditioning, # which this module deliberately does not import (giant.config pulls in the # whole model package). AXIS_TYPES = ("physical", "embedding", "onehot") @dataclass(frozen=True) class CondLayout: """Column layout of `cond_cont`/`cond_cat` for one (particle, material) mode pair. `cond_cont` is unconditionally `COND_DIM` wide regardless of mode: the base block, then the particle physical block, then the material physical block. An axis that isn't `"physical"` gets its block zero-filled and never reads it (see `giant.data.transforms._physical_cond_columns`), so the widths are mode-independent and only the *meaning* of a block changes. `cond_cat` is 2 to 4 wide. Columns `PDG_COL`/`MAT_COL` are always the dense training-vocab index; an axis in `"onehot"` mode appends one more column holding its top-N-plus-other class index, particle before material. """ particle_type: str material_type: str # cond_cat's dense-vocab columns, present in every mode. Under # "physical"/"onehot" they are a reporting/router convenience the # ConditionEncoder never reads; under "embedding" they are the signal. PDG_COL: ClassVar[int] = 0 MAT_COL: ClassVar[int] = 1 def __post_init__(self) -> None: if self.particle_type not in AXIS_TYPES: raise ValueError(f"unknown conditioning.particle.type {self.particle_type!r}") if self.material_type not in AXIS_TYPES: raise ValueError(f"unknown conditioning.material.type {self.material_type!r}") @classmethod def from_types(cls, particle_type: str, material_type: str) -> "CondLayout": """Named constructor — the entry point both sides use.""" return cls(particle_type=particle_type, material_type=material_type) # --- cond_cont --------------------------------------------------------- @property def base(self) -> slice: """pre_pos(3), log(pre_E)(1), pre_dir(3), layer_id(1).""" return slice(0, COND_DIM_BASE) @property def particle_phys(self) -> slice: """log(mass), charge — see `giant.particles`.""" return slice(COND_DIM_BASE, COND_DIM_BASE + PARTICLE_PHYS_DIM) @property def material_phys(self) -> slice: """Z_eff, A_eff, log(density), log(X0), log(lambda_int) — see `giant.materials`.""" start = COND_DIM_BASE + PARTICLE_PHYS_DIM return slice(start, start + MATERIAL_PHYS_DIM) @property def cont_dim(self) -> int: return COND_DIM # --- cond_cat ---------------------------------------------------------- @property def particle_topn_col(self) -> int | None: """Column of the particle top-N class index, or `None` if not `"onehot"`.""" return self.MAT_COL + 1 if self.particle_type == "onehot" else None @property def material_topn_col(self) -> int | None: """Column of the material top-N class index, or `None` if not `"onehot"`. Comes after the particle top-N column when both axes are `"onehot"`. """ if self.material_type != "onehot": return None return self.MAT_COL + (2 if self.particle_type == "onehot" else 1) @property def cat_dim(self) -> int: """Total `cond_cat` width: 2, plus one column per `"onehot"` axis.""" return self.MAT_COL + 1 + (self.particle_type == "onehot") + (self.material_type == "onehot")