[FEATURE] Add block-causal attention for dLLM example - #2499
Conversation
|
👋 Hi! Thank you for contributing to the TileLang project. Please remember to run We appreciate you taking this step! Our team will review your contribution, and we look forward to your awesome work! 🚀 |
|
Note Reviews pausedIt 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 Use the following commands to manage reviews:
Use the checkboxes below for quick actions:
📝 WalkthroughWalkthroughChangesThe PR replaces the fixed-length implementation with a block-size-64 TileLang kernel path parameterized by Block-causal attention
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
🚥 Pre-merge checks | ✅ 4 | ❌ 1❌ Failed checks (1 warning)
✅ Passed checks (4 passed)
✨ Finishing Touches🧪 Generate unit tests (beta)
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. Comment |
There was a problem hiding this comment.
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
📒 Files selected for processing (3)
examples/block_causal_attention/block_causal_attention.pyexamples/block_causal_attention/regression_block_causal_attention.pyexamples/block_causal_attention/test_block_causal_attention.py
| 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) |
There was a problem hiding this comment.
📐 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.
| 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
|
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 |
There was a problem hiding this comment.
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.
There was a problem hiding this comment.
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!
There was a problem hiding this comment.
@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.
|
I will be fixing the remaining concerns soon :) |
There was a problem hiding this comment.
Actionable comments posted: 2
🧹 Nitpick comments (1)
examples/block_causal_attention/block_causal_attention_varlen.py (1)
317-324: 📐 Maintainability & Code Quality | 🔵 Trivial | 💤 Low valueDuplicate 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
📒 Files selected for processing (3)
examples/block_causal_attention/block_causal_attention.pyexamples/block_causal_attention/block_causal_attention_varlen.pyexamples/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
Block-causal (dLLM) attention exampleMotivationDiffusion / 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 doesTwo self-contained example modules:
Highlights:
Files
BenchmarksSetupTested on 1xH100 for bf16 dtype against FlexAttention (max-autotune Fixed-length kernel
Varlen kernelTotal of 16k tokens,
Summary
|
|
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. |
Summary
dllm_block_sizevalues.cu_seqlens, including forward/backward kernels, PyTorch wrappers, reference implementations, and validation tests.