Skip to content

[Feature] Introduce tile scheduler - #2441

Merged
LeiWang1999 merged 2 commits into
tile-ai:mainfrom
Rachmanino:tile-schedule
Jun 24, 2026
Merged

LeiWang1999 merged 2 commits into
tile-ai:mainfrom
Rachmanino:tile-schedule

Conversation

@Rachmanino

@Rachmanino Rachmanino commented Jun 23, 2026 •

Copy link
Copy Markdown
Collaborator

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 precomputed waves/manual tile-walk logic.

Key Changes

New Modules

  • tilelang/language/meta.py (+193 lines): Implements meta-programming utilities:

    • @inline descriptor that dispatches to eager macro generation when an eager builder is active, otherwise lowers to TVMScript inline.
    • @meta_class class decorator for JIT-time stateful helper classes:
      • captures an optional prefix for buffer naming
      • best-effort renames assigned Buffer attributes as {prefix}_{attr}
      • auto-wraps selected store-emitting methods with @inline via 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-exports inline, meta_class, BaseTileScheduler, and PersistentTileScheduler.

Refactored Examples to Use T.PersistentTileScheduler

  • examples/gemm_sm100/gemm_tcgen5mma_ws_persistent.py: Refactors gemm_persistent and gemm_persistent_2cta to replace waves/manual tile-walk counters with scheduler-driven while sched.valid() loops. Loader/MMA/epilogue roles each own a scheduler instance and advance via sched.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 use cluster_size = 2.

  • examples/deepseek_v4/fp8_fp4_gemm_1d1d_sm100.py: Converts persistent scheduling logic to PersistentTileScheduler-driven loops for load, MMA, scale-factor transpose, and epilogue phases, removing explicit waves-based tile indexing and associated bounds guards.

Design Notes

  • Schedulers are implemented as Python classes decorated with @meta_class, allowing stateful tile traversal logic to be lowered into TIR while keeping coordinate decoding reusable (coord(tile_id) vs. stateful update_current_idx).
  • Persistent kernels now use scheduler state (current_iter, m_idx, n_idx, etc.) to drive role-specific pipeline progress and tile coordinate selection.

Testing & Review Notes

  • High review effort is expected due to new scheduler/meta-programming infrastructure and multiple kernel refactors.
  • The example changes demonstrate the intended usage pattern: multiple warp roles each instantiate and iterate their own persistent scheduler via while sched.valid().

C++ style / lint notes

  • This PR does not touch C++ sources or CI logic related to C++ linting; it only changes Python modules and Python example kernels.
  • The PR does not modify docs/developer_guide/cpp_style.md, so C++ style rule guidance is unaffected.
  • CI includes “C++ API Style Audit (warning only)”; since this PR does not change C++/API surface code, it should not introduce new warning-only TLCPP003/TLCPP004 findings from this change.

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.
@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 23, 2026 •

Copy link
Copy Markdown
Contributor

Review Change Stack

📝 Walkthrough

Walkthrough

Adds tilelang/language/meta.py with inline and meta_class meta-programming utilities, and tilelang/language/tile_schedule.py with BaseTileScheduler and PersistentTileScheduler built on meta_class. Re-exports these from the language module. Updates three SM100 GEMM examples to use PersistentTileScheduler per warp role instead of manually unrolled tile waves.

Changes

Persistent Tile Scheduler Infrastructure and Example Migrations

Layer / File(s) Summary
meta_class and inline TIR meta-programming primitives
tilelang/language/meta.py
Adds _InlineMethod descriptor dispatching to eager macro or TVMScript inline engine, _emits_store AST inspector detecting subscript-target assignments, prefix-based buffer naming utilities, and meta_class decorator that auto-wraps buffer-emitting methods with inline.
BaseTileScheduler and PersistentTileScheduler
tilelang/language/tile_schedule.py
Adds BaseTileScheduler with stateful T.alloc_var buffers and init/next_tile/valid control flow, and PersistentTileScheduler supporting cluster_size, swizzle_size, column_major, grid-strided init/next_tile, and stateless coord(tile_id) decode.
Language module re-exports
tilelang/language/__init__.py
Extends public API to re-export inline, meta_class, BaseTileScheduler, and PersistentTileScheduler.
SM100 GEMM examples migrated to PersistentTileScheduler
examples/gemm_sm100/gemm_tcgen5mma_ws_persistent.py, examples/blockscaled_gemm_sm100/gemm_mxfp8_blockscaled_1d1d.py, examples/deepseek_v4/fp8_fp4_gemm_1d1d_sm100.py
Replaces manual waves unrolling and swizzle math in three GEMM variants (warp-specialized, block-scaled, and DeepSeek) with per-warp-role PersistentTileScheduler instances in while sched.valid() loops; removes explicit tile-bounds checks and precomputed cluster tiling variables; adds cluster_size=2 for 2-CTA variants.

Sequence Diagram

sequenceDiagram
    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
Loading

Estimated code review effort

🎯 4 (Complex) | ⏱️ ~60 minutes

Possibly related issues

Suggested reviewers

  • LeiWang1999

🐇 Hippity-hop, the tiles align,
No more manual waves in line!
A scheduler loops where warps once toiled,
Swizzle and cluster — nothing foiled.
while sched.valid() — the bunny's refrain,
Persistent kernels hop on the same lane! 🎉

🚥 Pre-merge checks | ✅ 4 | ❌ 1

❌ Failed checks (1 warning)

Check name Status Explanation Resolution
Docstring Coverage ⚠️ Warning Docstring coverage is 20.69% 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.
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.
Title check ✅ Passed The title accurately summarizes the main change: introducing tile scheduler functionality.

✏️ Tip: You can configure your own custom pre-merge checks in the settings.

✨ 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: 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

📥 Commits

Reviewing files that changed from the base of the PR and between 74d4f0e and 958430b.

📒 Files selected for processing (4)
  • examples/gemm_sm100/gemm_tcgen5mma_ws_persistent.py
  • tilelang/language/__init__.py
  • tilelang/language/meta.py
  • tilelang/language/tile_schedule.py

Comment thread examples/gemm_sm100/gemm_tcgen5mma_ws_persistent.py
Comment thread tilelang/language/meta.py
Comment on lines +115 to +193
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

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

🧩 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.py

Repository: 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 -20

Repository: 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 -20

Repository: 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 -120

Repository: 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 -30

Repository: 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 -10

Repository: 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 -20

Repository: 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 expanded

Repository: 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 -30

Repository: 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 -100

Repository: 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 -20

Repository: 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.py

Repository: 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 -40

Repository: tile-ai/tilelang

Length of output: 225


🏁 Script executed:

# Search for actual instantiation of PersistentTileScheduler
rg "PersistentTileScheduler\(" --type py -B 2 -A 2 | head -50

Repository: 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.py

Repository: 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 -20

Repository: tile-ai/tilelang

Length of output: 1936


🏁 Script executed:

# Look at the complete meta_class decorator implementation
sed -n '145,193p' tilelang/language/meta.py

Repository: 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.

@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

🤖 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

📥 Commits

Reviewing files that changed from the base of the PR and between 958430b and 6874615.

📒 Files selected for processing (5)
  • examples/blockscaled_gemm_sm100/gemm_mxfp8_blockscaled_1d1d.py
  • examples/deepseek_v4/fp8_fp4_gemm_1d1d_sm100.py
  • examples/gemm_sm100/gemm_tcgen5mma_ws_persistent.py
  • tilelang/language/meta.py
  • tilelang/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

Comment on lines +392 to +396
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]

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.

🩺 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.py

Repository: 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.py

Repository: 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.py

Repository: 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.py

Repository: 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.py

Repository: 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 -40

Repository: 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.py

Repository: 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.py

Repository: 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 -60

Repository: 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.py

Repository: 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.py

Repository: 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 -10

Repository: 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.py

Repository: 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.py

Repository: 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.py

Repository: 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 -80

Repository: 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 -5

Repository: 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.py

Repository: 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 py

Repository: 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.

Comment on lines +84 to +88
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]

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.

🩺 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.

@LeiWang1999 LeiWang1999 changed the title [WIP][Feature] Introduce tile scheduler [Feature] Introduce tile scheduler Jun 24, 2026
@LeiWang1999
LeiWang1999 merged commit 8a68bbc into tile-ai:main Jun 24, 2026
9 of 10 checks passed
@Rachmanino
Rachmanino deleted the tile-schedule branch August 3, 2026 06:25
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