|
73 | 73 | from inference_engine.v04.restored_attention import prepare_restored_attention_kv |
74 | 74 |
|
75 | 75 |
|
| 76 | +def resolve_text_config(config: Any) -> Any: |
| 77 | + """Return the text-decoder sub-config for multimodal HF configs. |
| 78 | +
|
| 79 | + Gemma 4 (``Gemma4Config``) is a multimodal composite whose decoder |
| 80 | + dimensions live under ``config.text_config`` (with sibling |
| 81 | + ``vision_config`` / ``audio_config``). Flat text-only configs |
| 82 | + (e.g. Gemma 3) expose the attributes directly. This helper returns |
| 83 | + ``config.text_config`` when present, else ``config`` itself, so |
| 84 | + callers can read ``num_hidden_layers`` / ``num_key_value_heads`` / |
| 85 | + ``head_dim`` uniformly. |
| 86 | + """ |
| 87 | + return getattr(config, "text_config", None) or config |
| 88 | + |
| 89 | + |
| 90 | +def get_verifier_decoder(model: Any) -> Any: |
| 91 | + """Locate the decoder module exposing ``.layers`` / ``.embed_tokens``. |
| 92 | +
|
| 93 | + Handles two HF layouts: |
| 94 | + * flat: ``model.model`` (Gemma 3, Llama, ...) |
| 95 | + * multimodal: ``model.model.language_model`` (Gemma 4 conditional |
| 96 | + generation — text decoder nested beside vision/audio |
| 97 | + towers) |
| 98 | + """ |
| 99 | + base = getattr(model, "model", model) |
| 100 | + lm = getattr(base, "language_model", None) |
| 101 | + if lm is not None and hasattr(lm, "layers"): |
| 102 | + return lm |
| 103 | + if hasattr(base, "layers"): |
| 104 | + return base |
| 105 | + for attr in ("language_model", "text_model", "decoder"): |
| 106 | + sub = getattr(base, attr, None) |
| 107 | + if sub is not None and hasattr(sub, "layers"): |
| 108 | + return sub |
| 109 | + raise AttributeError( |
| 110 | + "could not locate decoder layers on verifier model " |
| 111 | + f"(type={type(model).__name__})" |
| 112 | + ) |
| 113 | + |
| 114 | + |
76 | 115 | @dataclasses.dataclass |
77 | 116 | class CrossModelLayerMapping: |
78 | 117 | """How drafter K/V layers project to verifier K/V layers under f_θ. |
@@ -149,8 +188,9 @@ def __init__( |
149 | 188 |
|
150 | 189 | def _validate_dimensions(self) -> None: |
151 | 190 | cfg = self.f_theta.config |
152 | | - # Verifier dimensions |
153 | | - v_cfg = self.verifier_model.config |
| 191 | + # Verifier dimensions (resolve multimodal text sub-config, e.g. |
| 192 | + # Gemma 4's config.text_config) |
| 193 | + v_cfg = resolve_text_config(self.verifier_model.config) |
154 | 194 | v_layers = getattr(v_cfg, "num_hidden_layers", None) |
155 | 195 | v_kv_heads = getattr(v_cfg, "num_key_value_heads", None) |
156 | 196 | v_head_dim = getattr(v_cfg, "head_dim", None) |
@@ -275,7 +315,7 @@ def forward( |
275 | 315 |
|
276 | 316 | # Patch verifier attention forwards to inject K/V at evicted |
277 | 317 | # positions. Restore originals after the forward. |
278 | | - layers = self.verifier_model.model.layers |
| 318 | + layers = get_verifier_decoder(self.verifier_model).layers |
279 | 319 | originals: List[Callable] = [] |
280 | 320 | try: |
281 | 321 | for layer_idx, layer in enumerate(layers): |
|
0 commit comments