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.
is_causal=Trueduring NNX decode attends only to key position 0nnx.MultiHeadAttentiondocumentsis_causalas 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_causalintoattention_fn. With query lengthT=1and KV lengthS=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 steppingdecode=True. Maintainer note on #4505 is that decode is already causal via the cache mask — sois_causalshould be a no-op there, not a second tril over the padded cache.Reproduction (current
main)Expected
Decode with
is_causal=Trueshould match decode withis_causal=False(occupancy mask only). Prefill /decode=Falsewithis_causal=Trueis unchanged.Proposed fix
Pass
is_causal=is_causal and not decodeintoattention_fn. Happy to send a PR with a regression test.