Skip to content

NNX MultiHeadAttention: is_causal=True during decode attends only to key position 0 #5591

Description

@Anu27n

is_causal=True during NNX decode attends only to key position 0

nnx.MultiHeadAttention documents is_causal as the way to overlay a causal mask, and decode already installs a cache-occupancy mask (arange(max_length) <= cur_index) so a single query can attend to every written KV slot.

Those two flags currently compose incorrectly. Decode still forwards is_causal into attention_fn. With query length T=1 and KV length S=max_length, jnp.tril(ones((1, S))) keeps only column 0. Combined with the occupancy mask, generation after the first token never attends to later cached keys.

This is easy to hit: train/prefill with is_causal=True (as in the MiniGPT example), then keep the same flag while stepping decode=True. Maintainer note on #4505 is that decode is already causal via the cache mask — so is_causal should be a no-op there, not a second tril over the padded cache.

Reproduction (current main)

import jax
import jax.numpy as jnp
from flax import nnx
import numpy as np

def make():
    return nnx.MultiHeadAttention(
        num_heads=2,
        in_features=4,
        qkv_features=4,
        decode=True,
        rngs=nnx.Rngs(0),
        kernel_init=nnx.initializers.normal(),
        bias_init=nnx.initializers.zeros_init(),
    )

seq = jax.random.normal(jax.random.key(1), (1, 4, 4))
a, b = make(), make()
a.init_cache(seq.shape)
b.init_cache(seq.shape)

out_ok, out_broken = [], []
for t in range(4):
    token = seq[:, t : t + 1, :]
    out_ok.append(a(token, decode=True, is_causal=False))
    out_broken.append(b(token, decode=True, is_causal=True))

print(jnp.max(jnp.abs(
    jnp.concatenate(out_ok, axis=1) - jnp.concatenate(out_broken, axis=1)
)))
# large; token 0 matches, later tokens diverge

Expected

Decode with is_causal=True should match decode with is_causal=False (occupancy mask only). Prefill / decode=False with is_causal=True is unchanged.

Proposed fix

Pass is_causal=is_causal and not decode into attention_fn. Happy to send a PR with a regression test.

Activity

  1. Anu27n commented on Sep 20, 2026

    @Anu27n
    Author

    I opened a fix in #5592 decode now skips the extra is_causal tril so it keeps the cache occupancy mask only.

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Metadata

Metadata

Assignees

No one assigned

    Labels

    No labels
    No labels

    Type

    No type

    Projects

    No projects

      Milestone

      No milestone

      Relationships

      None yet

      Development

      No branches or pull requests

      Issue actions