Skip to content

Align the JAX causal mask bottom-right when queries and keys differ - #172

Merged
jessegrabowski merged 2 commits into
pymc-devs:mainfrom
jessegrabowski:fix-jax-causal-alignment
Oct 6, 2026
Merged

jessegrabowski merged 2 commits into
pymc-devs:mainfrom
jessegrabowski:fix-jax-causal-alignment

Conversation

@jessegrabowski

@jessegrabowski jessegrabowski commented Oct 6, 2026 •

Copy link
Copy Markdown
Member

The JAX dispatch of AttentionLayer now builds the bottom-right causal triangle itself when the query is shorter than its keys. jax.nn.dot_product_attention aligns is_causal top-left, so decoding from a key-value cache attended only to the first keys.

Closes #120


📚 Documentation preview 📚: https://pytensor-ml--172.org.readthedocs.build/en/172/

@jessegrabowski
jessegrabowski merged commit 6a9c8c9 into pymc-devs:main Oct 6, 2026
14 checks passed
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

JAX aligns the causal mask top-left, so generation with a KV cache attends to the wrong prefix

1 participant