[Feature] Introduce tile scheduler - #2441
Conversation
This commit introduces the PersistentTileScheduler, which enhances the warp-specialized GEMM implementation by managing tile scheduling through a single iteration clock. The scheduler allows for efficient tile traversal and supports both TMA and tcgen5 operations within a persistent loop. Additionally, the meta class and related functions are added to facilitate the integration of stateful scheduling in the TileLang framework. The changes also include updates to the gemm_tcgen5mma_ws_persistent.py example to utilize the new scheduler effectively.
|
👋 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! 🚀 |
📝 WalkthroughWalkthroughAdds ChangesPersistent Tile Scheduler Infrastructure and Example Migrations
Sequence DiagramsequenceDiagram
participant Warp0 as Warp-0 (TMA Loader)
participant Warp1 as Warp-1 (tcgen5 MMA)
participant EpiW as Epilogue Warp
participant Sched as PersistentTileScheduler
participant SMEM as Shared Memory
participant TMEM as TMEM Double Buffer
participant C as Output Tensor C
rect rgba(70, 130, 180, 0.5)
note over Warp0,Sched: TMA Load phase
Warp0->>Sched: sched_tma.init(block_id)
loop while sched_tma.valid()
Warp0->>Sched: read m_idx, n_idx, current_iter
Warp0->>SMEM: T.tma_copy A/B tiles (double-buffered by w&1)
Warp0->>Sched: sched_tma.next_tile()
end
end
rect rgba(60, 179, 113, 0.5)
note over Warp1,TMEM: MMA phase
Warp1->>Sched: sched_mma.init(block_id)
loop while sched_mma.valid()
Warp1->>SMEM: mbarrier_wait for A/B tiles
Warp1->>TMEM: tcgen05_gemm into TMEM[w&1]
Warp1->>Sched: sched_mma.next_tile()
end
end
rect rgba(210, 105, 30, 0.5)
note over EpiW,C: Epilogue phase
EpiW->>Sched: sched_epi.init(block_id)
loop while sched_epi.valid()
EpiW->>TMEM: read TMEM double buffer
EpiW->>C: store via TMA or direct cast
EpiW->>Sched: sched_epi.next_tile()
end
end
Estimated code review effort🎯 4 (Complex) | ⏱️ ~60 minutes Possibly related issues
Suggested reviewers
🚥 Pre-merge checks | ✅ 4 | ❌ 1❌ Failed checks (1 warning)
✅ Passed checks (4 passed)
✏️ Tip: You can configure your own custom pre-merge checks in the settings. ✨ 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: 2
🤖 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/gemm_sm100/gemm_tcgen5mma_ws_persistent.py`:
- Around line 89-107: The MMA consumer loop is using incorrect slot indices. The
TMA producer writes to shared memory slots using `phase % num_stages` (where
`phase = w * k_blocks + k`), but the MMA loop is reading from `A_shared[k %
num_stages]` and `B_shared[k % num_stages]` while waiting on `loaded[phase %
num_stages]`. Replace all `k % num_stages` indices with `phase % num_stages` in
both the tcgen05_gemm calls within the if and else blocks, so that the MMA
consumer reads from the same slots that the producer writes to and signals via
the consumed mbarrier. This aligns the 1-CTA variant with the 2-CTA variant
pattern.
In `@tilelang/language/meta.py`:
- Around line 115-193: Add a regression test that validates the buffer naming
mechanism in the meta_class decorator works correctly for both BaseTileScheduler
and PersistentTileScheduler. The test should instantiate both classes with a
specific prefix value and verify that any buffers allocated in their __init__
methods are automatically named with the pattern {prefix}_{attribute_name} (for
example, if prefix is "sched", buffers should be named "sched_m_idx",
"sched_n_idx", etc). This test ensures that the _install_prefix_naming wrapper
correctly captures the prefix argument and that __setattr__ properly names
buffers in both the base class and subclass __init__ methods, preventing
regressions if future signature changes break the prefix naming mechanism.
🪄 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: ac30f9b9-f524-4618-931f-294536ee6bb0
📒 Files selected for processing (4)
examples/gemm_sm100/gemm_tcgen5mma_ws_persistent.pytilelang/language/__init__.pytilelang/language/meta.pytilelang/language/tile_schedule.py
| def _install_prefix_naming(cls: type) -> None: | ||
| """Make ``self.<attr> = <buffer>`` auto-name the buffer using ``prefix``. | ||
|
|
||
| ``prefix`` is read from the constructor argument of the same name (if any), | ||
| so existing ``__init__(self, prefix, ...)`` signatures just work. | ||
| """ | ||
| if cls.__dict__.get("_tl_prefix_naming_installed", False): | ||
| return | ||
|
|
||
| orig_init = cls.__init__ | ||
| init_sig = inspect.signature(orig_init) | ||
|
|
||
| def __init__(self, *args, **kwargs): | ||
| prefix = None | ||
| try: | ||
| bound = init_sig.bind(self, *args, **kwargs) | ||
| bound.apply_defaults() | ||
| prefix = bound.arguments.get("prefix") | ||
| except TypeError: | ||
| prefix = None | ||
| object.__setattr__(self, "_meta_prefix", prefix) | ||
| orig_init(self, *args, **kwargs) | ||
|
|
||
| def __setattr__(self, name, value): | ||
| object.__setattr__(self, name, value) | ||
| prefix = getattr(self, "_meta_prefix", None) | ||
| if prefix and not name.startswith("_"): | ||
| _name_buffer(prefix, name, value) | ||
|
|
||
| cls.__init__ = __init__ | ||
| cls.__setattr__ = __setattr__ | ||
| cls._tl_prefix_naming_installed = True | ||
|
|
||
|
|
||
| def meta_class(cls: _C) -> _C: | ||
| """Class decorator for JIT-time stateful helpers (e.g. tile schedulers). | ||
|
|
||
| Instances exist only during JIT tracing / parsing and hold ``T.alloc_var`` | ||
| buffers as state (``T.alloc_var`` emits into the active IR frame in both | ||
| modes). The decorator has three responsibilities: | ||
|
|
||
| 1. Mark the class with ``_is_meta_class``. The lazy parser needs this to | ||
| bind ``sched = Sched(...)`` as a (non-TIR) instance in its scope instead | ||
| of trying to turn it into a constant. | ||
| 2. Auto-``inline`` every TIR-emitting method, so methods need no per-method | ||
| ``@inline``. A method is considered TIR-emitting iff it contains a | ||
| buffer store, i.e. a subscript-target assignment ``self.x[...] = ...`` | ||
| (see ``_emits_store``). Left as plain Python: | ||
|
|
||
| - dunders (``__init__`` allocates state as plain Python), | ||
| - ``staticmethod`` / ``classmethod`` / ``property``, | ||
| - methods already decorated with ``@inline``, | ||
| - methods with no buffer store -- pure compile-time helpers that only | ||
| build/return ``PrimExpr`` expressions (e.g. ``valid`` returning a | ||
| loop condition, or a ``coord(tile_id)`` decode returning ``(m, n)``). | ||
| These stay plain so they can return values and be reused statelessly. | ||
| A method that emits TIR only through control flow / nested calls (no | ||
| direct store) should be marked explicitly with ``@inline``. | ||
| 3. Auto-name state buffers ``{prefix}_{attr}`` in the generated IR, where | ||
| ``prefix`` is the constructor argument named ``prefix`` (so | ||
| ``Sched("sched", ...)`` yields ``sched_m_idx`` etc.). This wraps | ||
| ``__init__`` to capture ``prefix`` and installs a ``__setattr__`` that | ||
| names any ``Buffer`` assigned to a non-underscore attribute; non-buffers | ||
| (compile-time ints like ``num_n_tiles``) and underscore attributes are | ||
| skipped, and naming is best-effort (never fatal). It runs at | ||
| construction while an IR builder is active, so it works in both modes. | ||
| """ | ||
| cls._is_meta_class = True | ||
| _install_prefix_naming(cls) | ||
| for name, attr in list(cls.__dict__.items()): | ||
| if name.startswith("__"): | ||
| continue | ||
| if isinstance(attr, _InlineMethod): | ||
| continue | ||
| if isinstance(attr, (staticmethod, classmethod, property)): | ||
| continue | ||
| if inspect.isfunction(attr) and _emits_store(attr): | ||
| setattr(cls, name, inline(attr)) | ||
| return cls |
There was a problem hiding this comment.
📐 Maintainability & Code Quality | 🟡 Minor
🧩 Analysis chain
🏁 Script executed:
#!/bin/bash
rg -nP '`@meta_class`|def __init__\s*\(' tilelang/language/tile_schedule.py
ast-grep run --pattern 'def __init__($$$)' --lang python tilelang/language/tile_schedule.pyRepository: tile-ai/tilelang
Length of output: 3310
🏁 Script executed:
# Check _name_buffer implementation and any related tests
grep -n "_name_buffer" tilelang/language/meta.py
find . -name "*test*.py" -o -name "test_*.py" | head -20Repository: tile-ai/tilelang
Length of output: 1316
🏁 Script executed:
# Look for tests that specifically test the meta_class decorator and prefix naming
rg -l "meta_class|prefix.*naming" --type py | head -20Repository: tile-ai/tilelang
Length of output: 154
🏁 Script executed:
# Get the full context of _name_buffer and _emits_store functions in meta.py
sed -n '1,114p' tilelang/language/meta.py | head -120Repository: tile-ai/tilelang
Length of output: 4242
🏁 Script executed:
# Search for tests related to tile_schedule or buffer naming
rg -n "BaseTileScheduler|PersistentTileScheduler" --type py | head -30Repository: tile-ai/tilelang
Length of output: 154
🏁 Script executed:
# Look for test files that import or use these classes
rg -l "from.*tile_schedule import|import.*tile_schedule" --type py | head -10Repository: tile-ai/tilelang
Length of output: 154
🏁 Script executed:
# Check if there are any tests that validate prefix naming or buffer names
rg -n "prefix.*name|_meta_prefix|_name_buffer" --type py | head -20Repository: tile-ai/tilelang
Length of output: 154
🏁 Script executed:
# Get the class definitions in tile_schedule.py
ast-grep outline tilelang/language/tile_schedule.py --view expandedRepository: tile-ai/tilelang
Length of output: 645
🏁 Script executed:
# Search for any imports or usage of tile_schedule module across the codebase
rg "tile_schedule|tile scheduler" --type py -i | head -30Repository: tile-ai/tilelang
Length of output: 154
🏁 Script executed:
# Look at complete tile_schedule.py to understand the structure
wc -l tilelang/language/tile_schedule.py
cat -n tilelang/language/tile_schedule.py | head -100Repository: tile-ai/tilelang
Length of output: 5788
🏁 Script executed:
# Search for usage of BaseTileScheduler or PersistentTileScheduler in examples
rg "Scheduler|from tilelang.language.tile_schedule" examples/ --type py | head -20Repository: tile-ai/tilelang
Length of output: 1044
🏁 Script executed:
# Check the complete __init__ signatures for both classes
sed -n '40,47p' tilelang/language/tile_schedule.py
sed -n '193,204p' tilelang/language/tile_schedule.pyRepository: tile-ai/tilelang
Length of output: 803
🏁 Script executed:
# Look for any direct instantiation or usage patterns of these classes
rg -A 5 "Scheduler(" --type py | head -40Repository: tile-ai/tilelang
Length of output: 225
🏁 Script executed:
# Search for actual instantiation of PersistentTileScheduler
rg "PersistentTileScheduler\(" --type py -B 2 -A 2 | head -50Repository: tile-ai/tilelang
Length of output: 154
🏁 Script executed:
# Check if there's a test file for tile_schedule
find . -name "*test*tile*schedule*" -o -name "*tile*schedule*test*"Repository: tile-ai/tilelang
Length of output: 154
🏁 Script executed:
# Look at one example that uses PersistentTileScheduler to understand the usage
sed -n '1,150p' examples/gemm_sm100/gemm_tcgen5mma_ws_persistent.py | grep -A 10 "PersistentTileScheduler"Repository: tile-ai/tilelang
Length of output: 2426
🏁 Script executed:
# Get the full _install_prefix_naming function
sed -n '115,150p' tilelang/language/meta.pyRepository: tile-ai/tilelang
Length of output: 1410
🏁 Script executed:
# Check if there's a way to inspect generated IR to verify buffer names
rg "IRBuilder.name|alloc_var" tilelang/language --type py | head -20Repository: tile-ai/tilelang
Length of output: 1936
🏁 Script executed:
# Look at the complete meta_class decorator implementation
sed -n '145,193p' tilelang/language/meta.pyRepository: tile-ai/tilelang
Length of output: 2724
Add a test validating buffer naming for both BaseTileScheduler and PersistentTileScheduler.
The wrapped __init__/__setattr__ mechanism works correctly here: _install_prefix_naming uses per-class guard checks, so both decorators wrap their respective __init__ methods. When PersistentTileScheduler.__init__ calls super().__init__(prefix, stateful=stateful), it re-enters the wrapped base __init__, which successfully binds its signature because both prefix (positional) and stateful (keyword) match the base signature. Buffers allocated in both __init__ methods get named via __setattr__. However, no existing test validates that buffers are actually named {prefix}_* (e.g., sched_m_idx, sched_n_idx). A regression test would prevent future signature changes in subclasses from silently breaking prefix naming.
🤖 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 `@tilelang/language/meta.py` around lines 115 - 193, Add a regression test that
validates the buffer naming mechanism in the meta_class decorator works
correctly for both BaseTileScheduler and PersistentTileScheduler. The test
should instantiate both classes with a specific prefix value and verify that any
buffers allocated in their __init__ methods are automatically named with the
pattern {prefix}_{attribute_name} (for example, if prefix is "sched", buffers
should be named "sched_m_idx", "sched_n_idx", etc). This test ensures that the
_install_prefix_naming wrapper correctly captures the prefix argument and that
__setattr__ properly names buffers in both the base class and subclass __init__
methods, preventing regressions if future signature changes break the prefix
naming mechanism.
Source: Linters/SAST tools
This commit refines the implementation of the PersistentTileScheduler in the GEMM examples, specifically in the `gemm_mxfp8_blockscaled_1d1d.py` and `fp8_fp4_gemm_1d1d_sm100.py` files. The changes include the introduction of a single iteration clock for managing tile scheduling, enhancing the clarity and efficiency of the scheduling logic. Additionally, the comments have been updated to reflect the new structure and usage of the scheduler, ensuring better understanding and maintainability of the code.
There was a problem hiding this comment.
Actionable comments posted: 2
🤖 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/blockscaled_gemm_sm100/gemm_mxfp8_blockscaled_1d1d.py`:
- Around line 392-396: The calculation of bx using sched.m_idx[0] * cluster_size
+ cta_id can exceed m_blocks when M is not divisible by cluster_size * block_M,
causing out-of-bounds accesses during TMA loads and C epilogue stores. Fix this
by adding an assertion at the point where cluster_size and block_M are defined
to ensure M % (cluster_size * block_M) == 0, or alternatively add a bounds check
guard condition (bx < m_blocks) around the tensor access operations that use bx,
such as the TMA load operation and the C epilogue store operations, to prevent
accessing invalid indices.
In `@examples/deepseek_v4/fp8_fp4_gemm_1d1d_sm100.py`:
- Around line 84-88: The calculation of bx using sched.m_idx[0] * cluster_size +
cta_id can exceed the m_blocks boundary when m_blocks is not a multiple of
cluster_size, causing out-of-bounds reads from matrices A and SFA and writes to
matrix D. Add a bounds check to ensure bx is less than m_blocks before using it
to access these arrays; this prevents the tail cluster from attempting to access
memory past the valid range. Since M is a concrete int value, you can add an
explicit assertion or guard condition immediately after the bx and by
assignments in the scheduler loop.
🪄 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: c309d14b-fd09-463a-ada4-9215b1304aff
📒 Files selected for processing (5)
examples/blockscaled_gemm_sm100/gemm_mxfp8_blockscaled_1d1d.pyexamples/deepseek_v4/fp8_fp4_gemm_1d1d_sm100.pyexamples/gemm_sm100/gemm_tcgen5mma_ws_persistent.pytilelang/language/meta.pytilelang/language/tile_schedule.py
🚧 Files skipped from review as they are similar to previous changes (3)
- tilelang/language/meta.py
- tilelang/language/tile_schedule.py
- examples/gemm_sm100/gemm_tcgen5mma_ws_persistent.py
| sched = T.PersistentTileScheduler("sched_tma", m_blocks, n_blocks, swizzle_size=group_size, cluster_size=cluster_size) | ||
| sched.init(block_id // cluster_size) | ||
| while sched.valid(): | ||
| bx = sched.m_idx[0] * cluster_size + cta_id | ||
| by = sched.n_idx[0] |
There was a problem hiding this comment.
🩺 Stability & Availability | 🟠 Major
🧩 Analysis chain
🏁 Script executed:
#!/bin/bash
rg -nP -C3 'mxfp8_blockscaled_gemm_2cta_persistent|block_M\s*=' examples/blockscaled_gemm_sm100/Repository: tile-ai/tilelang
Length of output: 7196
🏁 Script executed:
# Get the outline of the mxfp8_blockscaled_gemm_2cta_persistent function
ast-grep outline examples/blockscaled_gemm_sm100/gemm_mxfp8_blockscaled_1d1d.py --view expanded | grep -A 50 "mxfp8_blockscaled_gemm_2cta_persistent"Repository: tile-ai/tilelang
Length of output: 528
🏁 Script executed:
# Read the function definition and the scheduler/bx calculation area
sed -n '324,410p' examples/blockscaled_gemm_sm100/gemm_mxfp8_blockscaled_1d1d.pyRepository: tile-ai/tilelang
Length of output: 3433
🏁 Script executed:
# Check the cluster_size and m_blocks calculation
sed -n '358,370p' examples/blockscaled_gemm_sm100/gemm_mxfp8_blockscaled_1d1d.pyRepository: tile-ai/tilelang
Length of output: 682
🏁 Script executed:
# Look for any M divisibility checks in the entry point function or callers
rg -n 'assert.*M.*%|M.*%.*block_M|assert.*m_blocks' examples/blockscaled_gemm_sm100/gemm_mxfp8_blockscaled_1d1d.pyRepository: tile-ai/tilelang
Length of output: 154
🏁 Script executed:
# Check line 501 (C store) and surrounding context for guards
sed -n '495,510p' examples/blockscaled_gemm_sm100/gemm_mxfp8_blockscaled_1d1d.pyRepository: tile-ai/tilelang
Length of output: 826
🏁 Script executed:
# Find the wrapper/entry point that calls mxfp8_blockscaled_gemm_2cta_persistent
sed -n '630,680p' examples/blockscaled_gemm_sm100/gemm_mxfp8_blockscaled_1d1d.pyRepository: tile-ai/tilelang
Length of output: 2559
🏁 Script executed:
# Check for M divisibility assertions in the main entry point
rg -B5 -A5 'def blockscaled_gemm\(' examples/blockscaled_gemm_sm100/gemm_mxfp8_blockscaled_1d1d.py | head -40Repository: tile-ai/tilelang
Length of output: 154
🏁 Script executed:
# Look at the full blockscaled_gemm function to see if it enforces M divisibility
sed -n '595,675p' examples/blockscaled_gemm_sm100/gemm_mxfp8_blockscaled_1d1d.pyRepository: tile-ai/tilelang
Length of output: 3529
🏁 Script executed:
# Check grouped_gemm caller as well
sed -n '540,580p' examples/blockscaled_gemm_sm100/grouped_gemm_mxfp8_blockscaled_1d1d.pyRepository: tile-ai/tilelang
Length of output: 1104
🏁 Script executed:
# Search for any tests or examples that show how these functions are called
rg -n 'blockscaled_gemm\(' examples/blockscaled_gemm_sm100/ --type py -A3 | head -60Repository: tile-ai/tilelang
Length of output: 1650
🏁 Script executed:
# Check if there are any guards around bx accesses (lines 400-415 for TMA load and 495-510 for store)
sed -n '398,415p' examples/blockscaled_gemm_sm100/gemm_mxfp8_blockscaled_1d1d.pyRepository: tile-ai/tilelang
Length of output: 1083
🏁 Script executed:
# Look for any validation in the kernel entry or wrapper functions
sed -n '185,200p' examples/blockscaled_gemm_sm100/gemm_mxfp8_blockscaled_1d1d.pyRepository: tile-ai/tilelang
Length of output: 565
🏁 Script executed:
# Check if there are any wrapper or high-level APIs that validate M
rg -n 'def.*blockscaled_gemm' examples/blockscaled_gemm_sm100/gemm_mxfp8_blockscaled_1d1d.py | head -10Repository: tile-ai/tilelang
Length of output: 371
🏁 Script executed:
# Search for any comments or documentation about M constraints
rg -B2 'M.*divisible|M.*multiple|cluster_size.*block_M' examples/blockscaled_gemm_sm100/gemm_mxfp8_blockscaled_1d1d.pyRepository: tile-ai/tilelang
Length of output: 154
🏁 Script executed:
# Check if there are any conditional guards around the bx TMA access or store
sed -n '390,425p' examples/blockscaled_gemm_sm100/gemm_mxfp8_blockscaled_1d1d.pyRepository: tile-ai/tilelang
Length of output: 1866
🏁 Script executed:
# Look at the C store operation more carefully for any guards
sed -n '490,510p' examples/blockscaled_gemm_sm100/gemm_mxfp8_blockscaled_1d1d.pyRepository: tile-ai/tilelang
Length of output: 1033
🏁 Script executed:
# Check git history to see if bx < m_blocks was recently removed
git log -p --all -S "bx < m_blocks" -- examples/blockscaled_gemm_sm100/gemm_mxfp8_blockscaled_1d1d.py 2>/dev/null | head -80Repository: tile-ai/tilelang
Length of output: 154
🏁 Script executed:
# Check all usages of bx in the file to see if there are any guards
rg -n 'bx' examples/blockscaled_gemm_sm100/gemm_mxfp8_blockscaled_1d1d.py | grep -v 'by\|swap'Repository: tile-ai/tilelang
Length of output: 978
🏁 Script executed:
# Let's check what the clustered scheduler actually does - look for its implementation
rg -n 'class PersistentTileScheduler|def PersistentTileScheduler' --type py 2>/dev/null | head -5Repository: tile-ai/tilelang
Length of output: 154
🏁 Script executed:
# Check if M divisibility is validated at the wrapper/API level
sed -n '18,75p' examples/blockscaled_gemm_sm100/gemm_mxfp8_blockscaled_1d1d.pyRepository: tile-ai/tilelang
Length of output: 2225
🏁 Script executed:
# Search for any M divisibility checks anywhere in the repo
rg -n 'M.*%.*block_M|assert.*M' examples/blockscaled_gemm_sm100/ --type pyRepository: tile-ai/tilelang
Length of output: 576
Out-of-bounds bx index when M is not divisible by cluster_size * block_M.
The calculation bx = sched.m_idx[0] * cluster_size + cta_id can exceed m_blocks when M is not a multiple of cluster_size * block_M (i.e., 256). For example, if M = 257, the scheduler creates m_blocks = 3, but the tail cluster allows bx = 3 when cta_id = 1, causing out-of-bounds reads in the TMA load (line 404) and stores in the C epilogue (line 501).
The test suite only exercises M = 8192, which is perfectly divisible by 256, masking this issue. Add an assertion that M % (cluster_size * block_M) == 0, or restore a bx < m_blocks guard around the tensor accesses.
🤖 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/blockscaled_gemm_sm100/gemm_mxfp8_blockscaled_1d1d.py` around lines
392 - 396, The calculation of bx using sched.m_idx[0] * cluster_size + cta_id
can exceed m_blocks when M is not divisible by cluster_size * block_M, causing
out-of-bounds accesses during TMA loads and C epilogue stores. Fix this by
adding an assertion at the point where cluster_size and block_M are defined to
ensure M % (cluster_size * block_M) == 0, or alternatively add a bounds check
guard condition (bx < m_blocks) around the tensor access operations that use bx,
such as the TMA load operation and the C epilogue store operations, to prevent
accessing invalid indices.
| sched = T.PersistentTileScheduler("sched_tma", m_blocks, n_blocks, swizzle_size=group_size, cluster_size=cluster_size) | ||
| sched.init(block_id // cluster_size) | ||
| while sched.valid(): | ||
| bx = sched.m_idx[0] * cluster_size + cta_id | ||
| by = sched.n_idx[0] |
There was a problem hiding this comment.
🩺 Stability & Availability | 🟠 Major | ⚡ Quick win
Removed bx bounds check can read/write out of bounds when m_blocks is odd.
Same issue as the block-scaled example: bx = sched.m_idx[0] * cluster_size + cta_id reaches m_blocks for cta_id == 1 in the tail cluster when m_blocks isn't a multiple of cluster_size, indexing A/SFA and the D store (Line 182) past the end. Here M is a concrete int, so a guard is cheap — add an explicit assertion (or keep the old bx < m_blocks check).
🛡️ Proposed assertion
block_M = 128
block_N = 256
block_K = 128
+ assert M % (2 * block_M) == 0 # cluster_size == 2 splits M across CTAs🤖 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/deepseek_v4/fp8_fp4_gemm_1d1d_sm100.py` around lines 84 - 88, The
calculation of bx using sched.m_idx[0] * cluster_size + cta_id can exceed the
m_blocks boundary when m_blocks is not a multiple of cluster_size, causing
out-of-bounds reads from matrices A and SFA and writes to matrix D. Add a bounds
check to ensure bx is less than m_blocks before using it to access these arrays;
this prevents the tail cluster from attempting to access memory past the valid
range. Since M is a concrete int value, you can add an explicit assertion or
guard condition immediately after the bx and by assignments in the scheduler
loop.
Overview
This PR introduces tile scheduler infrastructure to TileLang, enabling persistent kernel patterns with flexible tile iteration strategies. It adds meta-programming utilities (
@inline,@meta_class) for lowering stateful Python helpers into TIR, and new scheduler classes (BaseTileScheduler,PersistentTileScheduler) that manage tile traversal with clustering and swizzling options. It also refactors multiple SM100 persistent GEMM example kernels to use the new persistent schedulers instead of precomputedwaves/manual tile-walk logic.Key Changes
New Modules
tilelang/language/meta.py(+193 lines): Implements meta-programming utilities:@inlinedescriptor that dispatches to eagermacrogeneration when an eager builder is active, otherwise lowers to TVMScriptinline.@meta_classclass decorator for JIT-time stateful helper classes:prefixfor buffer namingBufferattributes as{prefix}_{attr}@inlinevia AST inspection.tilelang/language/tile_schedule.py(+262 lines): Adds scheduler implementations:BaseTileScheduler: core iteration logic with optional state buffers (m_idx,n_idx,linear_idx,current_iter) and control-flow methods (init,valid,next_tile, plus abstract coordinate decoding).PersistentTileScheduler: persistent/grid-strided traversal with configurable:cluster_size(clustered M tiling)swizzle_size(L2/panel swizzling behavior)column_major(axis fast/slow selection)coord(tile_id)stateless coordinate decoding and stateful worker-based iteration (init(worker_id),next_tile()).Updated Modules
tilelang/language/__init__.py(+10 lines): Re-exportsinline,meta_class,BaseTileScheduler, andPersistentTileScheduler.Refactored Examples to Use
T.PersistentTileSchedulerexamples/gemm_sm100/gemm_tcgen5mma_ws_persistent.py: Refactorsgemm_persistentandgemm_persistent_2ctato replacewaves/manual tile-walk counters with scheduler-drivenwhile sched.valid()loops. Loader/MMA/epilogue roles each own a scheduler instance and advance viasched.next_tile(), deriving tile coordinates from scheduler state.examples/blockscaled_gemm_sm100/gemm_mxfp8_blockscaled_1d1d.py: Updates the persistent 2-CTA scheduling across TMA load / MMA compute / scale-factor transpose / epilogue to scheduler-driven iteration, and adjusts persistent cluster configuration to usecluster_size = 2.examples/deepseek_v4/fp8_fp4_gemm_1d1d_sm100.py: Converts persistent scheduling logic toPersistentTileScheduler-driven loops for load, MMA, scale-factor transpose, and epilogue phases, removing explicitwaves-based tile indexing and associated bounds guards.Design Notes
@meta_class, allowing stateful tile traversal logic to be lowered into TIR while keeping coordinate decoding reusable (coord(tile_id)vs. statefulupdate_current_idx).current_iter,m_idx,n_idx, etc.) to drive role-specific pipeline progress and tile coordinate selection.Testing & Review Notes
while sched.valid().C++ style / lint notes
docs/developer_guide/cpp_style.md, so C++ style rule guidance is unaffected.