Skip to content
Open
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
17 changes: 16 additions & 1 deletion gemma/gm/nn/gemma4/_transformer.py
Original file line number Diff line number Diff line change
Expand Up @@ -463,12 +463,27 @@ def _encode_and_get_inputs(
sliding_attention_mask = None
if self.config.use_bidirectional_attention == 'vision':
bidirectional_mask = tokens == _token_utils.SOFT_TOKEN_PLACEHOLDER
sliding_attention_mask = (
sliding_bidir_mask = (
_attention_mask.make_causal_bidirectional_attention_mask(
inputs_mask,
bidirectional_mask=bidirectional_mask,
)
)
# For multi-turn with cache: expand the sliding mask to cover
# cached history tokens. History portion uses the same causal
# mask as attention_mask (those tokens are already in the KV
# cache and _create_sliding_mask will apply the window).
if (
attention_mask is not None
and attention_mask.shape[-1] > sliding_bidir_mask.shape[-1]
):
hist_width = attention_mask.shape[-1] - sliding_bidir_mask.shape[-1]
hist_mask = attention_mask[:, :, :hist_width]
sliding_attention_mask = jnp.concatenate(
[hist_mask, sliding_bidir_mask], axis=-1
)
else:
sliding_attention_mask = sliding_bidir_mask

if self.config.per_layer_input_dim:
per_layer_inputs = self.embedder.encode_per_layer_input(
Expand Down
Loading