Skip to content

Commit 634acea

Browse files
K3: support Gemma4 multimodal nested config/decoder in f_theta train + cross-model verifier
Co-authored-by: FluffyAIcode <FluffyAIcode@users.noreply.github.com>
1 parent c404aee commit 634acea

2 files changed

Lines changed: 55 additions & 9 deletions

File tree

inference_engine/v04/cross_model_dlm_verifier.py

Lines changed: 43 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -73,6 +73,45 @@
7373
from inference_engine.v04.restored_attention import prepare_restored_attention_kv
7474

7575

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+
76115
@dataclasses.dataclass
77116
class CrossModelLayerMapping:
78117
"""How drafter K/V layers project to verifier K/V layers under f_θ.
@@ -149,8 +188,9 @@ def __init__(
149188

150189
def _validate_dimensions(self) -> None:
151190
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)
154194
v_layers = getattr(v_cfg, "num_hidden_layers", None)
155195
v_kv_heads = getattr(v_cfg, "num_key_value_heads", None)
156196
v_head_dim = getattr(v_cfg, "head_dim", None)
@@ -275,7 +315,7 @@ def forward(
275315

276316
# Patch verifier attention forwards to inject K/V at evicted
277317
# positions. Restore originals after the forward.
278-
layers = self.verifier_model.model.layers
318+
layers = get_verifier_decoder(self.verifier_model).layers
279319
originals: List[Callable] = []
280320
try:
281321
for layer_idx, layer in enumerate(layers):

scripts/research/k3_f_theta_train.py

Lines changed: 12 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -82,7 +82,11 @@
8282
import torch.nn.functional as F
8383

8484
from inference_engine.v04.f_theta import FThetaConfig, FThetaProjection
85-
from inference_engine.v04.cross_model_dlm_verifier import _capture_drafter_kv
85+
from inference_engine.v04.cross_model_dlm_verifier import (
86+
_capture_drafter_kv,
87+
get_verifier_decoder,
88+
resolve_text_config,
89+
)
8690
from inference_engine.v04.dflash_drafter import DFlashDrafter
8791

8892

@@ -189,7 +193,7 @@ def _capture_verifier_kv(
189193
(verifier_k, verifier_v) of shape [num_v_layers, T, verifier_kv_dim]
190194
each, on the verifier's device.
191195
"""
192-
layers = verifier_model.model.layers
196+
layers = get_verifier_decoder(verifier_model).layers
193197
num_layers = len(layers)
194198
k_capture: List[torch.Tensor] = [None] * num_layers
195199
v_capture: List[torch.Tensor] = [None] * num_layers
@@ -386,14 +390,16 @@ def main() -> int:
386390
for p in drafter.parameters():
387391
p.requires_grad_(False)
388392

389-
# Derive f_θ config from drafter + verifier shapes
393+
# Derive f_θ config from drafter + verifier shapes. Gemma 4's config
394+
# nests decoder dims under .text_config, so resolve it first.
395+
v_cfg = resolve_text_config(verifier.config)
390396
f_cfg = FThetaConfig(
391397
drafter_num_layers=drafter.cfg.num_hidden_layers,
392398
drafter_num_kv_heads=drafter.cfg.num_key_value_heads,
393399
drafter_head_dim=drafter.cfg.head_dim,
394-
verifier_num_layers=verifier.config.num_hidden_layers,
395-
verifier_num_kv_heads=verifier.config.num_key_value_heads,
396-
verifier_head_dim=verifier.config.head_dim,
400+
verifier_num_layers=v_cfg.num_hidden_layers,
401+
verifier_num_kv_heads=v_cfg.num_key_value_heads,
402+
verifier_head_dim=v_cfg.head_dim,
397403
rank=args.rank,
398404
)
399405
print(f"[f_theta-train] f_θ config: {f_cfg}", file=sys.stderr)

0 commit comments

Comments
 (0)