Skip to content

Commit d3a64c0

Browse files
K3: Gemma4-faithful cross-model restore forward (per-layer KV, v_norm, RoPE unsqueeze_dim=2, v_proj-None, evicted slicing) + gemma4 helpers import + tests
Co-authored-by: FluffyAIcode <FluffyAIcode@users.noreply.github.com>
1 parent 4a4d96d commit d3a64c0

4 files changed

Lines changed: 157 additions & 53 deletions

File tree

inference_engine/v04/cross_model_dlm_verifier.py

Lines changed: 56 additions & 30 deletions
Original file line numberDiff line numberDiff line change
@@ -247,14 +247,15 @@ def _validate_dimensions(self) -> None:
247247
@torch.no_grad()
248248
def project_drafter_kv(
249249
self, input_ids: torch.Tensor,
250-
) -> Tuple[torch.Tensor, torch.Tensor]:
250+
) -> Tuple[List[torch.Tensor], List[torch.Tensor]]:
251251
"""Run the drafter forward over input_ids, project K/V through f_θ.
252252
253253
Returns
254254
-------
255-
(verifier_k, verifier_v) tensors of shape
256-
``[B, T, verifier_num_layers, verifier_num_kv_heads, verifier_head_dim]``
257-
on the f_θ device.
255+
(verifier_k, verifier_v): per-layer lists of length
256+
``verifier_num_layers``, element ``i`` shaped
257+
``[B, T, layer_kv_heads[i], verifier_head_dim]`` on the f_θ
258+
device. Per-layer KV-head counts can differ (Gemma 4).
258259
259260
These are the per-position-per-verifier-layer K/V that the
260261
cross-model verifier injects at evicted positions during its
@@ -309,9 +310,9 @@ def forward(
309310
if not evicted_positions:
310311
return self.verifier_model(input_ids=input_ids, use_cache=False)
311312

312-
# f_θ projection
313-
verifier_k_full, verifier_v_full = self.project_drafter_kv(input_ids)
314-
# verifier_k_full shape: [B, T, L_v, num_kv_heads_v, head_dim_v]
313+
# f_θ projection → per-layer lists (layers can have different
314+
# KV-head counts on Gemma 4).
315+
verifier_k_layers, verifier_v_layers = self.project_drafter_kv(input_ids)
315316

316317
# Patch verifier attention forwards to inject K/V at evicted
317318
# positions. Restore originals after the forward.
@@ -325,8 +326,8 @@ def forward(
325326
attn,
326327
layer_idx=layer_idx,
327328
evicted_positions=evicted_positions,
328-
verifier_k_at_layer=verifier_k_full[:, :, layer_idx],
329-
verifier_v_at_layer=verifier_v_full[:, :, layer_idx],
329+
verifier_k_at_layer=verifier_k_layers[layer_idx],
330+
verifier_v_at_layer=verifier_v_layers[layer_idx],
330331
apply_rotary_pos_emb=apply_rotary_pos_emb,
331332
eager_attention_forward=eager_attention_forward,
332333
all_attention_functions=all_attention_functions,
@@ -351,47 +352,72 @@ def _make_patched_forward(
351352
instead of using the verifier's own k_proj / v_proj at those
352353
positions.
353354
354-
The patched forward replicates the standard verifier attention
355-
layer (Q, K, V projections + RoPE + GQA + softmax) with one
356-
change: after K, V are computed at every position, K and V at
357-
evicted positions are OVERWRITTEN with the f_θ-projected values
358-
(after k_norm + RoPE applied to match the standard pipeline).
355+
Mirrors Gemma 4's ``Gemma4TextAttention.forward`` exactly (RoPE
356+
applied per-tensor on the ``[B, T, H, D]`` layout with
357+
``unsqueeze_dim=2``; ``q_norm`` / ``k_norm`` / ``v_norm``;
358+
``v_proj`` is ``None`` on full-attention layers where V = raw K)
359+
with one change: at evicted positions K/V come from the
360+
f_θ-projected values (after k_norm + RoPE for K, after v_norm
361+
for V), so the verifier attends over full context despite
362+
holding only sink+window in its local cache.
359363
"""
360364
def _patched_forward(
361365
hidden_states: torch.Tensor,
362366
position_embeddings: Tuple[torch.Tensor, torch.Tensor],
363367
attention_mask: Optional[torch.Tensor] = None,
368+
shared_kv_states: Any = None,
364369
past_key_values=None,
365370
cache_position=None,
366371
**kwargs,
367372
) -> Tuple[torch.Tensor, Optional[torch.Tensor]]:
368-
B, T, _ = hidden_states.shape
369-
370-
input_shape = (B, T)
371-
hidden_shape = (*input_shape, -1, attn_module.head_dim)
372-
373-
query_states = attn_module.q_proj(hidden_states).view(*hidden_shape).transpose(1, 2)
374-
key_states = attn_module.k_proj(hidden_states).view(*hidden_shape).transpose(1, 2)
375-
value_states = attn_module.v_proj(hidden_states).view(*hidden_shape).transpose(1, 2)
373+
input_shape = hidden_states.shape[:-1] # [B, T]
374+
head_dim = attn_module.head_dim
375+
hidden_shape = (*input_shape, -1, head_dim)
376+
cos, sin = position_embeddings
376377

378+
# Query (Gemma 4: norm → RoPE on [B,T,Hq,D] → transpose)
379+
query_states = attn_module.q_proj(hidden_states).view(hidden_shape)
377380
query_states = attn_module.q_norm(query_states)
378-
key_states = attn_module.k_norm(key_states)
381+
query_states = apply_rotary_pos_emb(
382+
query_states, cos, sin, unsqueeze_dim=2,
383+
)
384+
query_states = query_states.transpose(1, 2) # [B, Hq, T, D]
379385

380-
cos, sin = position_embeddings
381-
query_states, key_states = apply_rotary_pos_emb(
382-
query_states, key_states, cos, sin,
386+
# Key / Value. Full-attention layers have v_proj=None ⇒ the
387+
# value is the raw k_proj output (pre k_norm), per Gemma 4.
388+
k_lin = attn_module.k_proj(hidden_states).view(hidden_shape)
389+
if getattr(attn_module, "v_proj", None) is not None:
390+
v_lin = attn_module.v_proj(hidden_states).view(hidden_shape)
391+
else:
392+
v_lin = k_lin
393+
key_states = attn_module.k_norm(k_lin)
394+
key_states = apply_rotary_pos_emb(
395+
key_states, cos, sin, unsqueeze_dim=2,
383396
)
397+
key_states = key_states.transpose(1, 2) # [B, Hkv, T, D]
398+
value_states = attn_module.v_norm(v_lin).transpose(1, 2) # [B, Hkv, T, D]
384399

385400
# Inject f_θ K/V at evicted positions.
386401
# verifier_k_at_layer shape: [B, T, num_kv_heads_v, head_dim_v]
387-
# K/V from k_proj also at all T positions; we overwrite the
388-
# evicted slice with f_θ output (after k_norm + RoPE).
402+
# (pre-norm pre-RoPE). prepare_restored_attention_kv applies
403+
# k_norm + RoPE to K; V must be v_norm'd here to match the
404+
# local branch (Gemma 4 runs V through v_norm).
389405
if evicted_positions:
406+
idx = torch.tensor(
407+
evicted_positions, device=key_states.device, dtype=torch.long,
408+
)
409+
cap_k_pre = verifier_k_at_layer.index_select(1, idx).to(
410+
device=key_states.device, dtype=key_states.dtype,
411+
)
412+
cap_v_pre = verifier_v_at_layer.index_select(1, idx).to(
413+
device=value_states.device, dtype=value_states.dtype,
414+
)
415+
cap_v_norm = attn_module.v_norm(cap_v_pre)
390416
key_states, value_states = prepare_restored_attention_kv(
391417
K_local=key_states,
392418
V_local=value_states,
393-
captured_K_pre_norm=verifier_k_at_layer,
394-
captured_V=verifier_v_at_layer,
419+
captured_K_pre_norm=cap_k_pre,
420+
captured_V=cap_v_norm,
395421
evicted_positions=evicted_positions,
396422
k_norm=attn_module.k_norm,
397423
position_embeddings=(cos, sin),

scripts/research/k3_integrated_niah_eval.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -125,7 +125,7 @@ def main() -> int:
125125

126126
# ---------- Verifier (CUDA bf16) ----------
127127
from transformers import AutoModelForCausalLM, AutoTokenizer
128-
from transformers.models.gemma3.modeling_gemma3 import ( # type: ignore
128+
from transformers.models.gemma4.modeling_gemma4 import ( # type: ignore
129129
apply_rotary_pos_emb, eager_attention_forward, ALL_ATTENTION_FUNCTIONS,
130130
)
131131

tests/inference_engine/v04/test_cross_model_dlm_verifier.py

Lines changed: 56 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -77,7 +77,9 @@ def __init__(self) -> None:
7777
self.o_proj = nn.Linear(32, 32, bias=False)
7878
self.q_norm = nn.Identity()
7979
self.k_norm = nn.Identity()
80+
self.v_norm = nn.Identity() # Gemma 4 runs V through v_norm
8081
self.head_dim = 8
82+
self.num_key_value_groups = 1
8183
self.scaling = 8 ** -0.5
8284
self.attention_dropout = 0.0
8385
self.sliding_window = None
@@ -217,11 +219,15 @@ def test_returns_correct_shape(self):
217219
B, T = 1, 6
218220
ids = torch.randint(0, 64, (B, T), dtype=torch.long)
219221
v_k, v_v = v.project_drafter_kv(ids)
220-
assert tuple(v_k.shape) == (
221-
B, T, f_cfg.verifier_num_layers,
222-
f_cfg.verifier_num_kv_heads, f_cfg.verifier_head_dim,
222+
# Per-layer list contract (layers may have heterogeneous KV heads).
223+
assert len(v_k) == f_cfg.verifier_num_layers
224+
assert len(v_v) == f_cfg.verifier_num_layers
225+
per_layer = (
226+
B, T, f_cfg.verifier_num_kv_heads, f_cfg.verifier_head_dim,
223227
)
224-
assert tuple(v_v.shape) == tuple(v_k.shape)
228+
for ko, vo in zip(v_k, v_v):
229+
assert tuple(ko.shape) == per_layer
230+
assert tuple(vo.shape) == per_layer
225231

226232

227233
class TestNoEvictPath:
@@ -268,6 +274,52 @@ def _counted(ids):
268274
assert calls["drafter"] == 0
269275

270276

277+
def _gemma4_style_rope(x, cos, sin, unsqueeze_dim=2):
278+
"""Mirror Gemma 4's apply_rotary_pos_emb(x, cos, sin, unsqueeze_dim)."""
279+
c = cos.unsqueeze(unsqueeze_dim)
280+
s = sin.unsqueeze(unsqueeze_dim)
281+
half = x.shape[-1] // 2
282+
rot = torch.cat([-x[..., half:], x[..., :half]], dim=-1)
283+
return x * c + rot * s
284+
285+
286+
def _synthetic_eager(module, q, k, v, mask, dropout=0.0, scaling=1.0,
287+
sliding_window=None, **kw):
288+
"""Minimal eager attention returning [B, T, H, D] like HF's."""
289+
attn = torch.matmul(q, k.transpose(-1, -2)) * scaling
290+
if mask is not None:
291+
attn = attn + mask[..., : k.shape[-2]]
292+
attn = attn.softmax(dim=-1)
293+
out = torch.matmul(attn, v) # [B, H, T, D]
294+
return out.transpose(1, 2).contiguous(), None # [B, T, H, D]
295+
296+
297+
class TestPatchedForwardRestore:
298+
"""Exercise the Gemma 4-style patched attention forward end-to-end on
299+
the synthetic verifier with real evictions, so signature/shape/RoPE
300+
regressions are caught without the 26B model."""
301+
302+
def test_forward_with_eviction_runs_and_injects(self):
303+
f_cfg = _tiny_f_theta_config()
304+
f_theta = FThetaProjection(f_cfg)
305+
drafter = DFlashDrafter(_tiny_drafter_config())
306+
verifier = _SyntheticVerifier()
307+
verifier.embed_tokens = torch.nn.Embedding(64, 16)
308+
verifier.get_input_embeddings = lambda: verifier.embed_tokens
309+
v = CrossModelDLMRestoredVerifier(
310+
verifier_model=verifier, drafter=drafter, f_theta=f_theta,
311+
sink_size=1, window_size=2, # sink+window = 3
312+
)
313+
B, T = 1, 6 # evicted positions [1, 2, 3]
314+
ids = torch.randint(0, 64, (B, T), dtype=torch.long)
315+
out = v.forward(
316+
ids,
317+
apply_rotary_pos_emb=_gemma4_style_rope,
318+
eager_attention_forward=_synthetic_eager,
319+
)
320+
assert out.logits.shape == (B, T, 64)
321+
322+
271323
class TestExports:
272324

273325
def test_module_exposes_classes(self):

tests/inference_engine/v04/test_f_theta.py

Lines changed: 44 additions & 18 deletions
Original file line numberDiff line numberDiff line change
@@ -77,19 +77,42 @@ def test_forward_k_shape(self):
7777
B, T = 2, 7
7878
x = torch.randn(B, T, c.encoder_in_features)
7979
y = m.forward_k(x)
80-
assert tuple(y.shape) == (
81-
B, T, c.verifier_num_layers, c.verifier_num_kv_heads, c.verifier_head_dim,
82-
)
80+
assert isinstance(y, list)
81+
assert len(y) == c.verifier_num_layers
82+
for layer_out in y:
83+
assert tuple(layer_out.shape) == (
84+
B, T, c.verifier_num_kv_heads, c.verifier_head_dim,
85+
)
8386

8487
def test_forward_v_shape(self):
8588
c = _tiny_config()
8689
m = FThetaProjection(c)
8790
B, T = 1, 3
8891
x = torch.randn(B, T, c.encoder_in_features)
8992
y = m.forward_v(x)
90-
assert tuple(y.shape) == (
91-
B, T, c.verifier_num_layers, c.verifier_num_kv_heads, c.verifier_head_dim,
93+
assert isinstance(y, list)
94+
assert len(y) == c.verifier_num_layers
95+
for layer_out in y:
96+
assert tuple(layer_out.shape) == (
97+
B, T, c.verifier_num_kv_heads, c.verifier_head_dim,
98+
)
99+
100+
def test_heterogeneous_layer_kv_heads(self):
101+
"""Per-layer KV-head counts (Gemma 4: 8 sliding, 4 full)."""
102+
c = FThetaConfig(
103+
drafter_num_layers=2, drafter_num_kv_heads=2, drafter_head_dim=4,
104+
verifier_num_layers=4, verifier_num_kv_heads=8, verifier_head_dim=8,
105+
rank=16, verifier_layer_kv_heads=(8, 4, 8, 4),
92106
)
107+
assert c.layer_kv_dims == (64, 32, 64, 32)
108+
m = FThetaProjection(c)
109+
B, T = 1, 3
110+
x = torch.randn(B, T, c.encoder_in_features)
111+
y = m.forward_k(x)
112+
assert [t.shape[2] for t in y] == [8, 4, 8, 4]
113+
# JSON round-trip preserves the per-layer field
114+
c2 = FThetaConfig.from_json_dict(c.to_json_dict())
115+
assert c2 == c
93116

94117
def test_forward_k_rejects_wrong_rank(self):
95118
c = _tiny_config()
@@ -122,9 +145,12 @@ def test_returns_paired_k_v(self):
122145
for _ in range(c.drafter_num_layers)
123146
]
124147
k_out, v_out = m.forward_kv_pack(k_per_layer, v_per_layer)
125-
expected = (B, T, c.verifier_num_layers, c.verifier_num_kv_heads, c.verifier_head_dim)
126-
assert tuple(k_out.shape) == expected
127-
assert tuple(v_out.shape) == expected
148+
assert len(k_out) == c.verifier_num_layers
149+
assert len(v_out) == c.verifier_num_layers
150+
per_layer = (B, T, c.verifier_num_kv_heads, c.verifier_head_dim)
151+
for ko, vo in zip(k_out, v_out):
152+
assert tuple(ko.shape) == per_layer
153+
assert tuple(vo.shape) == per_layer
128154

129155
def test_rejects_wrong_layer_count(self):
130156
c = _tiny_config()
@@ -162,8 +188,10 @@ def test_consistency_with_explicit_concat(self):
162188
with torch.no_grad():
163189
k_out_direct = m.forward_k(k_concat)
164190
v_out_direct = m.forward_v(v_concat)
165-
assert torch.allclose(k_out_pack, k_out_direct, atol=1e-6)
166-
assert torch.allclose(v_out_pack, v_out_direct, atol=1e-6)
191+
for kp, kd in zip(k_out_pack, k_out_direct):
192+
assert torch.allclose(kp, kd, atol=1e-6)
193+
for vp, vd in zip(v_out_pack, v_out_direct):
194+
assert torch.allclose(vp, vd, atol=1e-6)
167195

168196

169197
class TestParameterCount:
@@ -218,8 +246,10 @@ def test_save_and_load_preserves_outputs(self, tmp_path):
218246
y_k_2 = m2.forward_k(x_k)
219247
y_v_2 = m2.forward_v(x_v)
220248

221-
assert torch.allclose(y_k_1, y_k_2)
222-
assert torch.allclose(y_v_1, y_v_2)
249+
for a, b in zip(y_k_1, y_k_2):
250+
assert torch.allclose(a, b)
251+
for a, b in zip(y_v_1, y_v_2):
252+
assert torch.allclose(a, b)
223253

224254
def test_load_rejects_missing_config(self, tmp_path):
225255
# Write only weights, no config
@@ -267,12 +297,8 @@ def test_gradients_flow_for_k_path(self):
267297
m = FThetaProjection(c)
268298
B, T = 1, 3
269299
x = torch.randn(B, T, c.encoder_in_features, requires_grad=False)
270-
target = torch.randn(
271-
B, T, c.verifier_num_layers,
272-
c.verifier_num_kv_heads, c.verifier_head_dim,
273-
)
274-
out = m.forward_k(x)
275-
loss = ((out - target) ** 2).mean()
300+
out = m.forward_k(x) # list of [B, T, H, D]
301+
loss = sum(((o) ** 2).mean() for o in out)
276302
loss.backward()
277303
# encoder_k should have a grad
278304
assert m.encoder_k.weight.grad is not None

0 commit comments

Comments
 (0)