gemma-4-e2b-it-hybrid / configuration_gemma_4_e2b_it_hybrid.py
jakmro's picture
Add files using upload-large-folder tool
94981e6 verified
Raw History Blame
1.89 kB
"""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"]