Skip to content

Commit a5c9e54

Browse files
Add Gemma4 Assistant Model and Speculative Decoding (#2735)
* initial chnages * address gemini comments * fix XLA error * chnages to draft acceptance logic * revert cast value * speculative fix * match scale * fix descrepencies * nit * fix torch GPU error * address review comments * restore override layers
1 parent 99baf9a commit a5c9e54

18 files changed

Lines changed: 2334 additions & 79 deletions

keras_hub/api/models/__init__.py

Lines changed: 3 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -334,6 +334,9 @@
334334
from keras_hub.src.models.gemma3n.gemma3n_tokenizer import (
335335
Gemma3nTokenizer as Gemma3nTokenizer,
336336
)
337+
from keras_hub.src.models.gemma4.gemma4_assistant_causal_lm import (
338+
Gemma4AssistantCausalLM as Gemma4AssistantCausalLM,
339+
)
337340
from keras_hub.src.models.gemma4.gemma4_audio_encoder import (
338341
Gemma4AudioEncoder as Gemma4AudioEncoder,
339342
)

keras_hub/api/samplers/__init__.py

Lines changed: 3 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -14,5 +14,8 @@
1414
from keras_hub.src.samplers.serialization import deserialize as deserialize
1515
from keras_hub.src.samplers.serialization import get as get
1616
from keras_hub.src.samplers.serialization import serialize as serialize
17+
from keras_hub.src.samplers.speculative_sampler import (
18+
SpeculativeSampler as SpeculativeSampler,
19+
)
1720
from keras_hub.src.samplers.top_k_sampler import TopKSampler as TopKSampler
1821
from keras_hub.src.samplers.top_p_sampler import TopPSampler as TopPSampler

keras_hub/src/models/gemma4/gemma4_assistant_causal_lm.py

Lines changed: 396 additions & 0 deletions
Large diffs are not rendered by default.
Lines changed: 166 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,166 @@
1+
import os
2+
3+
import numpy as np
4+
from absl.testing import parameterized
5+
from keras import ops
6+
7+
from keras_hub.src.models.gemma4.gemma4_assistant_causal_lm import (
8+
Gemma4AssistantCausalLM,
9+
)
10+
from keras_hub.src.models.gemma4.gemma4_backbone import Gemma4Backbone
11+
from keras_hub.src.models.gemma4.gemma4_causal_lm import Gemma4CausalLM
12+
from keras_hub.src.tests.test_case import TestCase
13+
14+
15+
class Gemma4AssistantTest(TestCase, parameterized.TestCase):
16+
def setUp(self):
17+
self.backbone = Gemma4Backbone(
18+
vocabulary_size=256,
19+
num_layers=4,
20+
num_query_heads=4,
21+
num_key_value_heads=1,
22+
hidden_dim=8,
23+
intermediate_dim=16,
24+
head_dim=4,
25+
global_head_dim=8,
26+
image_size=16,
27+
layer_types=[
28+
"sliding_attention",
29+
"sliding_attention",
30+
"sliding_attention",
31+
"full_attention",
32+
],
33+
)
34+
# backbone_hidden_size=16 matches the target hidden_dim used in tests.
35+
self.model = Gemma4AssistantCausalLM(
36+
preprocessor=None,
37+
backbone=self.backbone,
38+
backbone_hidden_size=16,
39+
num_centroids=4,
40+
centroid_intermediate_top_k=2,
41+
use_ordered_embeddings=True,
42+
)
43+
44+
def test_call_with_cache(self):
45+
batch_size = 2
46+
target_num_layers = 6
47+
max_head_dim = 8 # max(head_dim=4, global_head_dim=8)
48+
target_kv_heads = 1
49+
cache_seq = 5
50+
51+
target_cache = np.zeros(
52+
(
53+
batch_size,
54+
target_num_layers,
55+
2,
56+
cache_seq,
57+
target_kv_heads,
58+
max_head_dim,
59+
),
60+
dtype="float32",
61+
)
62+
target_cache = ops.convert_to_tensor(target_cache)
63+
64+
last_token_embedding = ops.convert_to_tensor(
65+
np.random.randn(batch_size, 1, 16).astype("float32")
66+
)
67+
last_hidden_state = ops.convert_to_tensor(
68+
np.random.randn(batch_size, 1, 16).astype("float32")
69+
)
70+
71+
logits, next_hidden = self.model.call_with_cache(
72+
last_token_embedding=last_token_embedding,
73+
last_hidden_state=last_hidden_state,
74+
target_cache=target_cache,
75+
cache_update_index=cache_seq - 1,
76+
)
77+
78+
self.assertEqual(ops.shape(logits), (batch_size, 1, 256))
79+
self.assertEqual(ops.shape(next_hidden), (batch_size, 1, 16))
80+
81+
def test_speculative_generate(self):
82+
target_backbone = Gemma4Backbone(
83+
vocabulary_size=256,
84+
num_layers=6,
85+
num_query_heads=4,
86+
num_key_value_heads=1,
87+
hidden_dim=16,
88+
intermediate_dim=32,
89+
head_dim=8,
90+
image_size=16,
91+
layer_types=[
92+
"sliding_attention",
93+
"sliding_attention",
94+
"sliding_attention",
95+
"sliding_attention",
96+
"sliding_attention",
97+
"full_attention",
98+
],
99+
)
100+
target_model = Gemma4CausalLM(
101+
preprocessor=None,
102+
backbone=target_backbone,
103+
)
104+
105+
batch_size = 1
106+
max_length = 20
107+
seq_len = 5
108+
token_ids_raw = np.random.randint(0, 100, (batch_size, seq_len))
109+
token_ids = np.zeros((batch_size, max_length), dtype="int32")
110+
token_ids[:, :seq_len] = token_ids_raw
111+
padding_mask = np.zeros((batch_size, max_length), dtype="bool")
112+
padding_mask[:, :seq_len] = True
113+
token_ids = ops.convert_to_tensor(token_ids)
114+
padding_mask = ops.convert_to_tensor(padding_mask)
115+
116+
output = target_model.generate(
117+
{
118+
"token_ids": token_ids,
119+
"padding_mask": padding_mask,
120+
},
121+
assistant_model=self.model,
122+
stop_token_ids=None,
123+
)
124+
self.assertIsNotNone(output)
125+
126+
def test_model_saving(self):
127+
import keras
128+
129+
path = os.path.join(self.get_temp_dir(), "model.keras")
130+
self.model.save(path)
131+
loaded_model = keras.saving.load_model(path)
132+
133+
self.assertIsInstance(loaded_model, Gemma4AssistantCausalLM)
134+
135+
batch_size = 2
136+
target_num_layers = 6
137+
max_head_dim = 8
138+
target_kv_heads = 1
139+
cache_seq = 5
140+
target_cache = ops.zeros(
141+
(
142+
batch_size,
143+
target_num_layers,
144+
2,
145+
cache_seq,
146+
target_kv_heads,
147+
max_head_dim,
148+
)
149+
)
150+
last_token_embedding = ops.zeros((batch_size, 1, 16))
151+
last_hidden_state = ops.zeros((batch_size, 1, 16))
152+
153+
logits_orig, h_orig = self.model.call_with_cache(
154+
last_token_embedding=last_token_embedding,
155+
last_hidden_state=last_hidden_state,
156+
target_cache=target_cache,
157+
cache_update_index=cache_seq - 1,
158+
)
159+
logits_loaded, h_loaded = loaded_model.call_with_cache(
160+
last_token_embedding=last_token_embedding,
161+
last_hidden_state=last_hidden_state,
162+
target_cache=target_cache,
163+
cache_update_index=cache_seq - 1,
164+
)
165+
self.assertAllClose(logits_orig, logits_loaded)
166+
self.assertAllClose(h_orig, h_loaded)

keras_hub/src/models/gemma4/gemma4_backbone.py

Lines changed: 66 additions & 22 deletions
Original file line numberDiff line numberDiff line change
@@ -76,6 +76,11 @@ class Gemma4Backbone(Backbone):
7676
attention pattern. The last layer in each group of this many
7777
consecutive layers uses global attention; all others use local
7878
(sliding-window) attention. Defaults to `6`.
79+
layer_types: list of str or `None`. Explicit specification of the
80+
attention type for every layer sequentially
81+
(e.g. `"full_attention"`, `"sliding_attention"`). When `None`,
82+
type sequence is derived from `sliding_window_pattern`.
83+
Defaults to `None`.
7984
global_head_dim: int or `None`. Per-head dimension used specifically
8085
for global attention layers. When `None`, `head_dim` is used
8186
for all layers. Defaults to `None`.
@@ -194,6 +199,7 @@ def __init__(
194199
use_sliding_window_attention=True,
195200
sliding_window_size=512,
196201
sliding_window_pattern=6,
202+
layer_types=None,
197203
global_head_dim=None,
198204
local_rope_scaling_factor=1.0,
199205
global_rope_scaling_factor=1.0,
@@ -269,6 +275,7 @@ def __init__(
269275

270276
self.vision_encoder = vision_encoder
271277
self.audio_encoder = audio_encoder
278+
self.layer_types = layer_types
272279
text_only_model = vision_encoder is None and audio_encoder is None
273280
if vision_encoder is not None:
274281
self.interleave_embeddings = Gemma4InterleaveEmbeddings(
@@ -294,34 +301,62 @@ def __init__(
294301
# The last `num_kv_shared_layers` layers reuse K/V from the most
295302
# recent non-shared layer of the same attention type.
296303
_first_kv_shared = num_layers - num_kv_shared_layers
304+
_kv_source = {}
297305
if num_kv_shared_layers > 0:
298-
_non_shared_types = [
299-
"global"
300-
if (j % sliding_window_pattern) == (sliding_window_pattern - 1)
301-
else "local"
302-
for j in range(_first_kv_shared)
303-
]
304-
# Map each shared layer index → the absolute index of its KV source.
305-
_kv_source = {}
306-
for j in range(_first_kv_shared, num_layers):
307-
_is_g = (j % sliding_window_pattern) == (
308-
sliding_window_pattern - 1
309-
)
310-
_type = "global" if _is_g else "local"
311-
for k in range(len(_non_shared_types) - 1, -1, -1):
312-
if _non_shared_types[k] == _type:
313-
_kv_source[j] = k
314-
break
315-
else:
316-
_kv_source = {}
306+
if _first_kv_shared == 0:
307+
# Assistant mode: all layers share KV externally.
308+
# Mirror HF's hardcoded logic: assign arbitrary indices
309+
# 0 and 1 based on type.
310+
for j in range(num_layers):
311+
if layer_types is not None:
312+
_is_g = layer_types[j] == "full_attention"
313+
else:
314+
_is_g = (j % sliding_window_pattern) == (
315+
sliding_window_pattern - 1
316+
)
317+
_kv_source[j] = 1 if _is_g else 0
318+
else:
319+
if layer_types is not None:
320+
_non_shared_types = [
321+
"global" if t == "full_attention" else "local"
322+
for t in layer_types[:_first_kv_shared]
323+
]
324+
else:
325+
_non_shared_types = [
326+
"global"
327+
if (j % sliding_window_pattern)
328+
== (sliding_window_pattern - 1)
329+
else "local"
330+
for j in range(_first_kv_shared)
331+
]
332+
# Map shared layer index → absolute index of its KV source.
333+
for j in range(_first_kv_shared, num_layers):
334+
if layer_types is not None:
335+
_type = (
336+
"global"
337+
if layer_types[j] == "full_attention"
338+
else "local"
339+
)
340+
else:
341+
_is_g = (j % sliding_window_pattern) == (
342+
sliding_window_pattern - 1
343+
)
344+
_type = "global" if _is_g else "local"
345+
for k in range(len(_non_shared_types) - 1, -1, -1):
346+
if _non_shared_types[k] == _type:
347+
_kv_source[j] = k
348+
break
317349

318350
self.transformer_layers = []
319351
for i in range(num_layers):
320352
# A layer is global when it's the last in each group of
321353
# `sliding_window_pattern` consecutive layers.
322-
is_global = (i % sliding_window_pattern) == (
323-
sliding_window_pattern - 1
324-
)
354+
if layer_types is not None:
355+
is_global = layer_types[i] == "full_attention"
356+
else:
357+
is_global = (i % sliding_window_pattern) == (
358+
sliding_window_pattern - 1
359+
)
325360
sliding_window = use_sliding_window_attention and not is_global
326361
rope_wavelength = (
327362
(global_rope_wavelength or 1_000_000.0)
@@ -379,6 +414,14 @@ def __init__(
379414
)
380415
self.transformer_layers.append(layer)
381416

417+
if self.layer_types is None:
418+
self.layer_types = [
419+
"full_attention"
420+
if (i % sliding_window_pattern) == (sliding_window_pattern - 1)
421+
else "sliding_attention"
422+
for i in range(num_layers)
423+
]
424+
382425
self.layer_norm = RMSNormalization(
383426
epsilon=layer_norm_epsilon,
384427
dtype=dtype,
@@ -692,6 +735,7 @@ def get_config(self):
692735
),
693736
"sliding_window_size": self.sliding_window_size,
694737
"sliding_window_pattern": self.sliding_window_pattern,
738+
"layer_types": self.layer_types,
695739
"global_head_dim": self.global_head_dim,
696740
"local_rope_scaling_factor": self.local_rope_scaling_factor,
697741
"global_rope_scaling_factor": self.global_rope_scaling_factor,

0 commit comments

Comments
 (0)