tile-ai / tile-ai/tilelang

[BUG] T.Pipelined with a buffer-loaded if condition fails in InjectSoftwarePipeline

Open
#2,876 0 comments 0 reactions 0 assignees View on GitHub
bug
Dominant language
Python
Stars
7.4k
Forks
745
Avg merge
1d 1h
Merged PRs (30d)
104

Description

### Required prerequisites

- [x] I have read the documentation .
- [x] I have searched the [Issue Tracker](https://github.com/tile-ai/tilelang/issues) that this hasn't already been reported. (comment there if it has.)

### What version of TileLang are you using?

0.1.13

### System information

```
python -c "import sys, tilelang, torch; print(sys.version, sys.platform); print(tilelang.__version__); print(torch.__version__)"
3.10.12 (main, Sep 11 2024, 15:47:36) [GCC 11.4.0] linux
0.1.13
2.11.0+cu128
```

### Problem description

A `T.Pipelined` loop containing an `if` condition derived from a buffer-loaded scalar fails during `InjectSoftwarePipeline` in TileLang 0.1.13.

The same pattern compiled successfully with TileLang 0.1.9.

The original use case iterates over all Q blocks and skips blocks that do not cover the current KV tile:

```python
for q_tile in T.Pipelined(all_q_blocks, num_stages=...):
q_start = q_block_starts[q_tile]
q_end = q_block_ends[q_tile]

if (q_start <= kv_tile) & (kv_tile < q_end):
...
```

### Reproducible example code

```python
import tilelang
import tilelang.language as T

@T.prim_func
def pipeline_if_kernel(
q: T.Tensor((32, 16), T.bfloat16),
starts: T.Tensor((2,), T.int32),
out: T.Tensor((2, 2), T.float32),
):
with T.Kernel(2, threads=128) as kv_tile:
q_shared = T.alloc_shared((16, 16), T.bfloat16)
scores = T.alloc_fragment((16, 16), T.float32)
start = T.alloc_var(T.int32)

for q_tile in T.Pipelined(2, num_stages=1):
start = starts[q_tile]

if start <= kv_tile:
T.copy(
q[q_tile * 16 : (q_tile + 1) * 16, :],
q_shared,
)
T.clear(scores)
T.gemm(
q_shared,
q_shared,
scores,
transpose_B=True,
)
out[kv_tile, q_tile] = scores[0, 0]

def run_pipeline_passes():
target = tilelang.tvm.target.Target(
{"kind": "cuda", "arch": "sm_90"}
)

mod = tilelang.tvm.IRModule.from_expr(pipeline_if_kernel)
mod = tilelang.tvm.tirx.transform.BindTarget(target)(mod)
mod = tilelang.transform.MaterializeKernelLaunch()(mod)
mod = tilelang.transform.AddWrapperForSingleBufStore()(mod)
mod = tilelang.transform.LegalizeNegativeIndex()(mod)
mod = tilelang.transform.InjectAssumes()(mod)
mod = tilelang.transform.Simplify()(mod)
mod = tilelang.transform.LayoutReducer()(mod)
mod = tilelang.transform.IfStmtBinding()(mod)
mod = tilelang.transform.PipelinePlanning()(mod)

tilelang.transform.InjectSoftwarePipeline()(mod)

if __name__ == "__main__":
print(f"TileLang version: {tilelang.__version__}")
run_pipeline_passes()
```

Run with:

```bash
python repro_tilelang_pipeline_if.py
```

### Traceback

```pytb
File "python/tvm_ffi/cython/function.pxi", line 968, in tvm_ffi.core.Function.__call__
File "", line 0, in tvm::transform::Pass::operator()(tvm::IRModule) const
File "", line 0, in tvm::transform::Pass::operator()(tvm::IRModule, tvm::transform::PassContext const&) const
File "", line 0, in tvm::tirx::transform::PrimFuncPassNode::operator()(tvm::IRModule, tvm::transform::PassContext const&) const
File "", line 0, in tvm::tl::software_pipeline::InjectPipeline(tvm::tirx::PrimFunc const&)
File "", line 0, in tvm::tl::software_pipeline::PipelineInjector::Inject(tvm::tirx::PrimFunc const&)
File "", line 0, in tvm::tirx::StmtMutator::VisitStmt(tvm::tirx::Stmt const&)
File "", line 0, in tvm::NodeFunctor*)>::operator()(tvm::ffi::ObjectRef const&, tvm::tirx::StmtFunctor*) const
File "", line 0, in tvm::tirx::StmtFunctor::InitVTable()::{lambda(tvm::ffi::ObjectRef const&, tvm::tirx::StmtFunctor*)#13}::_FUN(tvm::ffi::ObjectRef const&, tvm::tirx::StmtFunctor*)
File "", line 0, in tvm::tirx::StmtMutator::VisitStmt_(tvm::tirx::SBlockRealizeNode const*)
File "", line 0, in tvm::tirx::StmtMutator::VisitStmt(tvm::tirx::Stmt const&)
File "", line 0, in tvm::NodeFunctor*)>::operator()(tvm::ffi::ObjectRef const&, tvm::tirx::StmtFunctor*) const
File "", line 0, in tvm::tirx::StmtFunctor::InitVTable()::{lambda(tvm::ffi::ObjectRef const&, tvm::tirx::StmtFunctor*)#12}::_FUN(tvm::ffi::ObjectRef const&, tvm::tirx::StmtFunctor*)
File "", line 0, in tvm::tl::software_pipeline::PipelineInjector::VisitStmt_(tvm::tirx::SBlockNode const*)
File "", line 0, in tvm::tirx::StmtMutator::VisitStmt_(tvm::tirx::SBlockNode const*)
File "", line 0, in tvm::tirx::StmtMutator::VisitStmt(tvm::tirx::Stmt const&)
File "", line 0, in tvm::NodeFunctor*)>::operator()(tvm::ffi::ObjectRef const&, tvm::tirx::StmtFunctor*) const
File "", line 0, in tvm::tirx::StmtFunctor::InitVTable()::{lambda(tvm::ffi::ObjectRef const&, tvm::tirx::StmtFunctor*)#2}::_FUN(tvm::ffi::ObjectRef const&, tvm::tirx::StmtFunctor*)
File "", line 0, in tvm::tirx::StmtMutator::VisitStmt_(tvm::tirx::AttrStmtNode const*)
File "", line 0, in tvm::tirx::StmtMutator::VisitStmt(tvm::tirx::Stmt const&)
File "", line 0, in tvm::NodeFunctor*)>::operator()(tvm::ffi::ObjectRef const&, tvm::tirx::StmtFunctor*) const
File "", line 0, in tvm::tirx::StmtFunctor::InitVTable()::{lambda(tvm::ffi::ObjectRef const&, tvm::tirx::StmtFunctor*)#2}::_FUN(tvm::ffi::ObjectRef const&, tvm::tirx::StmtFunctor*)
File "", line 0, in tvm::tirx::StmtMutator::VisitStmt_(tvm::tirx::AttrStmtNode const*)
File "", line 0, in tvm::tirx::StmtMutator::VisitStmt(tvm::tirx::Stmt const&)
File "", line 0, in tvm::NodeFunctor*)>::operator()(tvm::ffi::ObjectRef const&, tvm::tirx::StmtFunctor*) const
File "", line 0, in tvm::tirx::StmtFunctor::InitVTable()::{lambda(tvm::ffi::ObjectRef const&, tvm::tirx::StmtFunctor*)#2}::_FUN(tvm::ffi::ObjectRef const&, tvm::tirx::StmtFunctor*)
File "", line 0, in tvm::tirx::StmtMutator::VisitStmt_(tvm::tirx::AttrStmtNode const*)
File "", line 0, in tvm::tirx::StmtMutator::VisitStmt(tvm::tirx::Stmt const&)
File "", line 0, in tvm::NodeFunctor*)>::operator()(tvm::ffi::ObjectRef const&, tvm::tirx::StmtFunctor*) const
File "", line 0, in tvm::tirx::StmtFunctor::InitVTable()::{lambda(tvm::ffi::ObjectRef const&, tvm::tirx::StmtFunctor*)#2}::_FUN(tvm::ffi::ObjectRef const&, tvm::tirx::StmtFunctor*)
File "", line 0, in tvm::tirx::StmtMutator::VisitStmt_(tvm::tirx::AttrStmtNode const*)
File "", line 0, in tvm::tirx::StmtMutator::VisitStmt(tvm::tirx::Stmt const&)
File "", line 0, in tvm::NodeFunctor*)>::operator()(tvm::ffi::ObjectRef const&, tvm::tirx::StmtFunctor*) const
File "", line 0, in tvm::tirx::StmtFunctor::InitVTable()::{lambda(tvm::ffi::ObjectRef const&, tvm::tirx::StmtFunctor*)#13}::_FUN(tvm::ffi::ObjectRef const&, tvm::tirx::StmtFunctor*)
File "", line 0, in tvm::tirx::StmtMutator::VisitStmt_(tvm::tirx::SBlockRealizeNode const*)
File "", line 0, in tvm::tirx::StmtMutator::VisitStmt(tvm::tirx::Stmt const&)
File "", line 0, in tvm::NodeFunctor*)>::operator()(tvm::ffi::ObjectRef const&, tvm::tirx::StmtFunctor*) const
File "", line 0, in tvm::tirx::StmtFunctor::InitVTable()::{lambda(tvm::ffi::ObjectRef const&, tvm::tirx::StmtFunctor*)#12}::_FUN(tvm::ffi::ObjectRef const&, tvm::tirx::StmtFunctor*)
File "", line 0, in tvm::tl::software_pipeline::PipelineInjector::VisitStmt_(tvm::tirx::SBlockNode const*)
File "", line 0, in tvm::tirx::StmtMutator::VisitStmt_(tvm::tirx::SBlockNode const*)
File "", line 0, in tvm::tirx::StmtMutator::VisitStmt(tvm::tirx::Stmt const&)
File "", line 0, in tvm::NodeFunctor*)>::operator()(tvm::ffi::ObjectRef const&, tvm::tirx::StmtFunctor*) const
File "", line 0, in tvm::tirx::StmtFunctor::InitVTable()::{lambda(tvm::ffi::ObjectRef const&, tvm::tirx::StmtFunctor*)#4}::_FUN(tvm::ffi::ObjectRef const&, tvm::tirx::StmtFunctor*)
File "", line 0, in tvm::tl::software_pipeline::PipelineInjector::VisitStmt_(tvm::tirx::ForNode const*)
File "/project/src/transform/inject_pipeline.cc", line 3130, in void tvm::tl::software_pipeline::PipelineInjector::ValidatePipelineBody(const tvm::tl::software_pipeline::PipelineInfo&, const tvm::ffi::Array&)
tvm.error.InternalError: Check failed: src_info.order < dst_info.order (4 vs. 2) : ValueError: two statements with buffer access dependency in the same stage of the software pipeline cannot be reordered
```

### Expected behavior

- TileLang 0.1.9: compiles successfully
- TileLang 0.1.13: fails in `InjectSoftwarePipeline`

The kernel should compile successfully, as it did with TileLang 0.1.9.
An `if` condition inside `T.Pipelined` should be supported when its condition is loaded from a buffer through a local scalar variable.

### Additional context

Possibly related

- Previous support/fix for `if` statements inside pipelined loops: https://github.com/tile-ai/tilelang/pull/1799
- Related TileLang change: https://github.com/tile-ai/tilelang/commit/af30ac215226d068da689cac11da40e10b2b08e6

Contributor guide

Open the contributing guide

Research direction

Run repro_tilelang_pipeline_if.py with the reported pass sequence and inspect src/transform/inject_pipeline.cc around ValidatePipelineBody, reached through InjectSoftwarePipeline. Compare the current behavior with TileLang 0.1.9 and the earlier if-statement support in PR 1799; done means the buffer-loaded condition compiles successfully in the pipelined loop without the dependency-order error.

Written by the indexing model from the issue text.

Assessment

Tech stack
cpp, python
Domain
compilers
Issue type
Bug
Difficulty
4/5
Estimated time
3-5 days
Activity status
Quiet
Clarity
Mostly clear
Newbie friendliness
50/100

Get new issues in your inbox

A short digest of beginner-friendly GitHub issues.