"""Configuration for Cactus-Compute/gemma-4-e2b-it-hybrid. `Gemma4E2BItHybridConfig` is the stock Gemma-4 text config plus the hyperparameters of the handoff probe that scores every generation with ``confidence = 1 - p_wrong``. Everything the base model needs is inherited unchanged, so the base weights load with identical keys. """ from transformers.models.gemma4.configuration_gemma4 import Gemma4TextConfig class Gemma4E2BItHybridConfig(Gemma4TextConfig): """Gemma-4 text config extended with handoff-probe hyperparameters. Extra fields (all serialized to ``config.json``): - ``probe_layer`` (`int`, defaults to 28): zero-indexed decoder layer whose output feeds the probe. Must be ``< num_hidden_layers``. - ``probe_feature_size`` (`int`, defaults to 1536): width of the captured hidden states; must equal ``hidden_size``. - ``probe_max_tokens`` (`int`, defaults to 1024): the probe scores at most the first ``probe_max_tokens`` generated-token rows. - ``probe_proj_dim`` (`int`, defaults to 32): width of the probe's attention projection (fixed by the released checkpoint). """ model_type = "gemma-4-e2b-it-hybrid" probe_layer: int = 28 probe_feature_size: int = 1536 probe_max_tokens: int = 1024 probe_proj_dim: int = 32 def __post_init__(self, **kwargs): super().__post_init__(**kwargs) if not 0 <= self.probe_layer < self.num_hidden_layers: raise ValueError( f"probe_layer={self.probe_layer} must be in [0, num_hidden_layers=" f"{self.num_hidden_layers})" ) # Note: probe_feature_size == hidden_size is enforced by the model, not # here — `to_diff_dict()` default-constructs the config class, so this # class must stay constructible with pure defaults. __all__ = ["Gemma4E2BItHybridConfig"]