Skip to content

Add tensor-descriptor version of prefill_attention - #28

Open
ohadeytan wants to merge 1 commit into
mainfrom
prefill-attn-td
Open

ohadeytan wants to merge 1 commit into
mainfrom
prefill-attn-td

Conversation

@ohadeytan

Copy link
Copy Markdown
Collaborator

Summary

Adds a tensor-descriptor (TD) version of the prefill_attention flash-attention kernel, converting it from the block_ptr/pointer-arithmetic form to the tl.make_tensor_descriptor API via the td-convert skill.

Changes

  • block_ptr.py → tensor_descriptor.py: _fwd_kernel rewritten on tensor descriptors. Pointer arithmetic and masked loads/stores are replaced with descriptors whose shape/block_shape handle out-of-bounds rows and columns (zero-fill on load, clamp on store).
  • wrapper.py: updated to launch the TD kernel.
  • conversion-notes.md: records the conversion decisions.
  • tests/triton/test_prefill_attention_td.py: numerical-equivalence test against the original.
  • Removes the superseded block_ptr kernel, its kernel.ktir, and the old block_ptr-era tests.

Notes

This TD kernel is intended as a clean baseline: it will serve as the shared starting point for a follow-up spyre-aware conversion (the two operate on the same descriptor form).

🤖 Generated with Claude Code

Convert the flash-attention prefill kernel from block_ptr/pointer form to
the tl.make_tensor_descriptor API (td-convert skill):

- block_ptr.py -> tensor_descriptor.py (_fwd_kernel on descriptors; masks
  handled by descriptor block_shape / shape bounds).
- wrapper.py updated to launch the TD kernel.
- conversion-notes.md records the conversion decisions.
- tests/triton/test_prefill_attention_td.py: numerical-equivalence test.
- Remove the superseded block_ptr kernel, its kernel.ktir, and the old
  block_ptr-era tests.

Co-Authored-By: Claude Opus 4.8 (1M context) <noreply@anthropic.com>
Signed-off-by: Ohad Eytan <ohad.eytan1@ibm.com>
@ohadeytan
ohadeytan requested review from fabianlim and kiszk July 9, 2026 09:29
acc = acc / l_i[:, None]
acc = acc.to(Out.dtype.element_ty)

# The store clamps at shape=[cur_batch_seq_len, ...], so the partial tail

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

would be good to try to lower it with the triton fork to see if these index comparisons and tl.where will be lowered properly.

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.

2 participants