Skip to content

[FEATURE] Add block-causal attention for dLLM example - #2499

Merged
LeiWang1999 merged 7 commits into
tile-ai:mainfrom
perkyfever:block-causal-dllm-attention-example
Jul 23, 2026
Merged

LeiWang1999 merged 7 commits into
tile-ai:mainfrom
perkyfever:block-causal-dllm-attention-example

Conversation

@perkyfever

@perkyfever perkyfever commented Jun 30, 2026 •

Copy link
Copy Markdown
Contributor

Summary

  • Replaced the fixed block-causal attention implementation with TileLang forward and backward kernels parameterized by supported dllm_block_size values.
  • Added variable-length packed-sequence support using cu_seqlens, including forward/backward kernels, PyTorch wrappers, reference implementations, and validation tests.
  • Updated the test module to cover fixed-shape and variable-length attention across all supported block sizes.
  • Removed kernel caching, benchmarking, regression, and related helper routines from the fixed-shape example.

@github-actions

Copy link
Copy Markdown

👋 Hi! Thank you for contributing to the TileLang project.

Please remember to run pre-commit run --all-files in the root directory of the project to ensure your changes are properly linted and formatted. This will help ensure your contribution passes the format check.

We appreciate you taking this step! Our team will review your contribution, and we look forward to your awesome work! 🚀

@coderabbitai

coderabbitai Bot commented Jun 30, 2026 •

Copy link
Copy Markdown
Contributor

Review Change Stack

Note

Reviews paused

It looks like this branch is under active development. To avoid overwhelming you with review comments due to an influx of new commits, CodeRabbit has automatically paused this review. You can configure this behavior by changing the reviews.auto_review.auto_pause_after_reviewed_commits setting.

Use the following commands to manage reviews:

  • @coderabbitai resume to resume automatic reviews.
  • @coderabbitai review to trigger a single review.

Use the checkboxes below for quick actions:

  • ▶️ Resume reviews
  • 🔍 Trigger review
📝 Walkthrough

Walkthrough

Changes

The PR replaces the fixed-length implementation with a block-size-64 TileLang kernel path parameterized by dllm_block_size, adds packed varlen attention with forward/backward kernels, and introduces shared reference and CUDA validation entry points.

Block-causal attention

Layer / File(s) Summary
Fixed-length forward kernel
examples/block_causal_attention/block_causal_attention.py
Adds supported block-size validation, TileLang tile masking, tiled softmax computation, output generation, and LSE writes.
Fixed-length backward kernels
examples/block_causal_attention/block_causal_attention.py
Adds Delta preprocessing plus dQ, dK, and dV kernels that reuse the forward masking rules.
Fixed-length API and reference validation
examples/block_causal_attention/block_causal_attention.py
Wires the autograd wrapper and public API, adds explicit reference masking, and tests outputs and gradients across supported block sizes.
Packed varlen attention path
examples/block_causal_attention/block_causal_attention_varlen.py
Adds cu_seqlens-based forward and backward kernels, autograd dispatch, validation, and a packed reference implementation.
CUDA test entry points
examples/block_causal_attention/test_block_causal_attention.py, examples/block_causal_attention/block_causal_attention_varlen.py
Adds fixed and varlen CUDA-gated tests, block-size iteration, and module entry points.

Estimated code review effort: 4 (Complex) | ~60 minutes

Sequence Diagram(s)

sequenceDiagram
  participant Caller
  participant FixedAPI
  participant FixedAutograd
  participant FixedForward
  participant FixedBackward
  Caller->>FixedAPI: provide q, k, v, and dllm_block_size
  FixedAPI->>FixedAutograd: apply attention operation
  FixedAutograd->>FixedForward: compute masked output and LSE
  FixedForward-->>FixedAutograd: return output and LSE
  Caller->>FixedAutograd: provide dO
  FixedAutograd->>FixedBackward: compute dQ, dK, and dV
  FixedBackward-->>Caller: return gradients
Loading
🚥 Pre-merge checks | ✅ 4 | ❌ 1

❌ Failed checks (1 warning)

Check name Status Explanation Resolution
Docstring Coverage ⚠️ Warning Docstring coverage is 0.00% which is insufficient. The required threshold is 80.00%. Write docstrings for the functions missing them to satisfy the coverage threshold.
✅ Passed checks (4 passed)
Check name Status Explanation
Description Check ✅ Passed Check skipped - CodeRabbit’s high-level summary is enabled.
Title check ✅ Passed The title clearly matches the main change: adding block-causal attention support for the dLLM example.
Linked Issues check ✅ Passed Check skipped because no linked issues were found for this pull request.
Out of Scope Changes check ✅ Passed Check skipped because no linked issues were found for this pull request.
✨ Finishing Touches
🧪 Generate unit tests (beta)
  • Create PR with unit tests

Thanks for using CodeRabbit! It's free for OSS, and your support helps us grow. If you like it, consider giving us a shout-out.

❤️ Share

Comment @coderabbitai help to get the list of available commands.

@coderabbitai coderabbitai Bot left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

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

Actionable comments posted: 4

🤖 Prompt for all review comments with AI agents
Verify each finding against current code. Fix only still-valid issues, skip the
rest with a brief reason, keep changes minimal, and validate.

Inline comments:
In `@examples/block_causal_attention/block_causal_attention.py`:
- Around line 67-70: The public `mask_block_size` path in the block causal
attention flow currently accepts unsupported values, which can later produce
incorrect masking or delayed JIT failures. Add an early validation guard in the
wrapper and the forward schedule entry points (including the `forward`/mask
setup near the diagonal tile logic) to reject any `mask_block_size` that is not
supported by the 64-token tile layout and the backward kernels’ fixed
`mask_block == 32` specialization. Keep the check close to the existing
shape/tile assertions so invalid inputs fail fast before any kernel launch.
- Around line 179-190: The kernel in block_causal_attention.py uses the
ambiguous tensor name O, which Ruff E741 flags and can break linting. Rename O
to a clearer identifier throughout the affected function and update every
corresponding use in the same kernel body, including the matching argument and
any related references near dO and the T.copy call, so the naming stays
consistent and unambiguous.
- Around line 175-191: The Delta preprocessing in `prep` still assumes every `k`
block is a full 64-wide tile, so the `T.copy` slices can run past the last valid
columns when `dim` is not a multiple of `block`. Update the `prep` loop to guard
the tail case in `T.ceildiv(dim, block)` by restricting the copied range to the
remaining valid width, or fail fast up front if non-64-wide tail dimensions are
not supported. Keep the fix localized around `prep`, `Delta`, and the two
`T.copy` calls so the block dimension handling stays consistent.
- Around line 488-505: The forward launch in
block_causal_attention/_BlockCausalAttentionTL is still passing through
non-contiguous query, key, and value tensors, which TileLang’s tensor proxy may
reject at the host-kernel boundary. Normalize query/key/value in
block_causal_attention before calling _BlockCausalAttentionTL.apply by checking
is_contiguous() and materializing contiguous tensors when needed, and keep the
backward helper’s contig logic aligned with the same rule instead of relying
only on stride(-1) == 1.
🪄 Autofix (Beta)

Fix all unresolved CodeRabbit comments on this PR:

  • Push a commit to this branch (recommended)
  • Create a new PR with the fixes

ℹ️ Review info
⚙️ Run configuration

Configuration used: Path: .coderabbit.yaml

Review profile: CHILL

Plan: Pro

Run ID: 725afb7a-3d0b-4843-a6af-dbef47cd83ac

📥 Commits

Reviewing files that changed from the base of the PR and between d861ff4 and 0e10494.

📒 Files selected for processing (3)
  • examples/block_causal_attention/block_causal_attention.py
  • examples/block_causal_attention/regression_block_causal_attention.py
  • examples/block_causal_attention/test_block_causal_attention.py

Comment thread examples/block_causal_attention/block_causal_attention.py Outdated
Comment thread examples/block_causal_attention/block_causal_attention.py Outdated
Comment on lines +179 to +190
O: T.Tensor(shape, dtype),
dO: T.Tensor(shape, dtype),
Delta: T.Tensor([batch, heads, seq_len], accum_dtype),
):
with T.Kernel(heads, T.ceildiv(seq_len, block), batch) as (bx, by, bz):
o = T.alloc_fragment([block, block], dtype)
do = T.alloc_fragment([block, block], dtype)
acc = T.alloc_fragment([block, block], accum_dtype)
delta = T.alloc_fragment([block], accum_dtype)
T.clear(acc)
for k in range(T.ceildiv(dim, block)):
T.copy(O[bz, by * block : (by + 1) * block, bx, k * block : (k + 1) * block], o)

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

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

📐 Maintainability & Code Quality | 🟡 Minor | ⚡ Quick win

Rename O to satisfy Ruff E741.

Ruff flags O as an ambiguous variable name; this can fail lint checks.

Proposed rename
-        O: T.Tensor(shape, dtype),
+        Output: T.Tensor(shape, dtype),
         dO: T.Tensor(shape, dtype),
         Delta: T.Tensor([batch, heads, seq_len], accum_dtype),
@@
-                T.copy(O[bz, by * block : (by + 1) * block, bx, k * block : (k + 1) * block], o)
+                T.copy(Output[bz, by * block : (by + 1) * block, bx, k * block : (k + 1) * block], o)
📝 Committable suggestion

‼️ IMPORTANT
Carefully review the code before committing. Ensure that it accurately replaces the highlighted code, contains no missing lines, and has no issues with indentation. Thoroughly test & benchmark the code to ensure it meets the requirements.

Suggested change
O: T.Tensor(shape, dtype),
dO: T.Tensor(shape, dtype),
Delta: T.Tensor([batch, heads, seq_len], accum_dtype),
):
with T.Kernel(heads, T.ceildiv(seq_len, block), batch) as (bx, by, bz):
o = T.alloc_fragment([block, block], dtype)
do = T.alloc_fragment([block, block], dtype)
acc = T.alloc_fragment([block, block], accum_dtype)
delta = T.alloc_fragment([block], accum_dtype)
T.clear(acc)
for k in range(T.ceildiv(dim, block)):
T.copy(O[bz, by * block : (by + 1) * block, bx, k * block : (k + 1) * block], o)
Output: T.Tensor(shape, dtype),
dO: T.Tensor(shape, dtype),
Delta: T.Tensor([batch, heads, seq_len], accum_dtype),
):
with T.Kernel(heads, T.ceildiv(seq_len, block), batch) as (bx, by, bz):
o = T.alloc_fragment([block, block], dtype)
do = T.alloc_fragment([block, block], dtype)
acc = T.alloc_fragment([block, block], accum_dtype)
delta = T.alloc_fragment([block], accum_dtype)
T.clear(acc)
for k in range(T.ceildiv(dim, block)):
T.copy(Output[bz, by * block : (by + 1) * block, bx, k * block : (k + 1) * block], o)
🧰 Tools
🪛 Ruff (0.15.20)

[error] 179-179: Ambiguous variable name: O

(E741)

🤖 Prompt for AI Agents
Verify each finding against current code. Fix only still-valid issues, skip the
rest with a brief reason, keep changes minimal, and validate.

In `@examples/block_causal_attention/block_causal_attention.py` around lines 179 -
190, The kernel in block_causal_attention.py uses the ambiguous tensor name O,
which Ruff E741 flags and can break linting. Rename O to a clearer identifier
throughout the affected function and update every corresponding use in the same
kernel body, including the matching argument and any related references near dO
and the T.copy call, so the naming stays consistent and unambiguous.

Source: Linters/SAST tools

Comment thread examples/block_causal_attention/block_causal_attention.py Outdated
@LeiWang1999

Copy link
Copy Markdown
Member

please fix the lint :)

k_block = k_local // mask_block

same_region = T.if_then_else(q_clean == k_clean, 1, 0)
block_diagonal = T.if_then_else(q_block == k_block, 1, 0) * same_region

Copy link
Copy Markdown

Choose a reason for hiding this comment

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

Should it be block_diagonal only for noisy regions?

In the current implementation it seems to be block_diagonal both for noisy and clean part.

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

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

This code implements the block-wise mask. Based on this picture (left part) it should be only noisy.

Right now the code also enables block-diagonal part for the clean region, but those positions are allowed by clean block-causal masking anyway, so it's effectively redundant rather than wrong.

Though, the masking logic could be clearer if gated to noisy part only. Thanks for pointing this out!

Copy link
Copy Markdown

Choose a reason for hiding this comment

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

@perkyfever
I see, I thought the original intent was to have a mask like this
изображение
But if the intended mask is shown on the left pic you attached, then it should be correct.

@perkyfever

Copy link
Copy Markdown
Contributor Author

I will be fixing the remaining concerns soon :)

Godofnothing

This comment was marked as resolved.

Comment thread examples/block_causal_attention/block_causal_attention.py Outdated
Comment thread examples/block_causal_attention/block_causal_attention.py Outdated
Comment thread examples/block_causal_attention/block_causal_attention.py Outdated

@coderabbitai coderabbitai Bot left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

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

Actionable comments posted: 2

🧹 Nitpick comments (1)
examples/block_causal_attention/block_causal_attention_varlen.py (1)

317-324: 📐 Maintainability & Code Quality | 🔵 Trivial | 💤 Low value

Duplicate comment block.

Lines 317-320 and 321-324 are near-identical copy-paste; keep one.

🧹 Remove the redundant block
-                # Attending query tiles split into two *contiguous* groups so each loop index is
-                # affine (varlen host/device split rejects the non-affine merged schedule):
-                #   noisy group -- noisy key: just the diagonal {by}; clean key: [local, region_tiles)
-                #   clean group -- clean key only: [region_tiles + local, 2 * region_tiles)
                 # Attending query tiles split into two *contiguous* groups so each loop index is
                 # affine (the varlen host/device split rejects a non-affine merged schedule):
                 #   noisy group -- noisy key: just the diagonal {by}; clean key: [local, region_tiles)
                 #   clean group -- clean key only: [region_tiles + local, 2 * region_tiles)
🤖 Prompt for AI Agents
Verify each finding against current code. Fix only still-valid issues, skip the
rest with a brief reason, keep changes minimal, and validate.

In `@examples/block_causal_attention/block_causal_attention_varlen.py` around
lines 317 - 324, Remove the duplicate explanatory comment in the block around
the varlen attention scheduling logic, retaining a single accurate description
of the two contiguous query-tile groups and their affine loop requirements.
🤖 Prompt for all review comments with AI agents
Verify each finding against current code. Fix only still-valid issues, skip the
rest with a brief reason, keep changes minimal, and validate.

Inline comments:
In `@examples/block_causal_attention/block_causal_attention_varlen.py`:
- Around line 130-160: Rename the prep kernel parameter O in
_bwd_preprocess_varlen_template to a descriptive non-ambiguous name, and update
the corresponding T.copy access in prep while leaving all other parameters and
behavior unchanged.
- Around line 430-461: Remove the public block_size parameter and its host-side
divisibility checks, or thread it consistently through
_BlockCausalAttentionVarlenTL.forward/backward and
_bwd_preprocess_varlen_template so all kernels use the requested tile size.
Ensure non-64 callers no longer receive behavior that silently uses the
templates’ fixed 64-tile path.

---

Nitpick comments:
In `@examples/block_causal_attention/block_causal_attention_varlen.py`:
- Around line 317-324: Remove the duplicate explanatory comment in the block
around the varlen attention scheduling logic, retaining a single accurate
description of the two contiguous query-tile groups and their affine loop
requirements.
🪄 Autofix (Beta)

Fix all unresolved CodeRabbit comments on this PR:

  • Push a commit to this branch (recommended)
  • Create a new PR with the fixes

ℹ️ Review info
⚙️ Run configuration

Configuration used: Path: .coderabbit.yaml

Review profile: CHILL

Plan: Pro

Run ID: f7c07c28-123e-43af-8fb7-6e5a4cb3e4ad

📥 Commits

Reviewing files that changed from the base of the PR and between fe45685 and c83a086.

📒 Files selected for processing (3)
  • examples/block_causal_attention/block_causal_attention.py
  • examples/block_causal_attention/block_causal_attention_varlen.py
  • examples/block_causal_attention/test_block_causal_attention.py
🚧 Files skipped from review as they are similar to previous changes (1)
  • examples/block_causal_attention/test_block_causal_attention.py

Comment thread examples/block_causal_attention/block_causal_attention_varlen.py Outdated
Comment thread examples/block_causal_attention/block_causal_attention_varlen.py
@perkyfever

perkyfever commented Jul 12, 2026 •

Copy link
Copy Markdown
Contributor Author

Block-causal (dLLM) attention example

Motivation

Diffusion / block-diffusion LLMs (e.g. BD3-LM) train by attending over a sequence that is a noisy half followed by a clean half, under a specific block-causal mask. That mask is block-sparse, so a dense implementation (SDPA / a full attention mask, or an uncompiled FlexAttention mask) does a lot of wasted work.

What it does

Two self-contained example modules:

  • block_causal_attention.py -- fixed-length ([batch, seq_len, heads, dim]) forward + backward.
  • block_causal_attention_varlen.py -- the packed / variable-length (cu_seqlens) twin, sharing
    the mask helper and reference from the fixed module.

Highlights:

  • Arbitrary diffusion block size {1, 2, 4, 8, 16, 32, 64};
  • The dLLM mask (within each half, tokens grouped into blocks of dllm_block): block_diagonal, offset_causal and block_causal. See the left image from here;
  • Sparse tile schedule -- each query tile visits only the key tiles it can attend;
  • FlashAttention-style: single-pass online-softmax forward + dQ query-parallel + dK/dV key-parallel + register-accumulated. No atomics => deterministic;
  • fp16 and bf16, validated against a PyTorch reference (forward output and gradients).

Files

  • block_causal_attention.py / block_causal_attention_varlen.py -- kernels + torch reference.
  • test_block_causal_attention.py -- correctness tests (all block sizes, fixed + varlen).

Benchmarks

Setup

Tested on 1xH100 for bf16 dtype against FlexAttention (max-autotune torch.compile & Triton backend) + BlockMask (max-autotune torch.compile).

Fixed-length kernel

seq_len=8192, heads=32, head_dim=192, batch=2:

dllm_block fwd+bwd (ours) fwd+bwd (flex) speedup
1 13.42 25.47 1.90x
8 13.55 25.97 1.92x
32 13.27 26.01 1.96x
64 13.51 26.06 1.93x

Varlen kernel

Total of 16k tokens, dllm_block=32:

docs fwd+bwd (ours) fwd+bwd (flex) speedup
2048x8 6.66 10.06 1.51x
4096x4 10.92 15.43 1.41x
8192x2 19.33 26.15 1.35x

Summary

  • 1.90-1.96x faster than FlexAttention on fwd+bwd (fixed-length, uniform across block size);
  • 1.35-1.51x faster than FlexAttention on fwd+bwd for packed inputs.

@perkyfever
perkyfever requested a review from LeiWang1999 July 12, 2026 23:03
@perkyfever

perkyfever commented Jul 13, 2026 •

Copy link
Copy Markdown
Contributor Author

Hi @LeiWang1999,

I believe, all the concerns raised so far seem to have been addressed. As suggested, I explored other examples and tried to align my implementation with them. I have also added more context in the comment above.

Please let me know if there’s anything else to be updated.

@LeiWang1999
LeiWang1999 merged commit 22baf2e into tile-ai:main Jul 23, 2026
3 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.

3 participants