mirror of
https://github.com/tile-ai/tilelang.git
synced 2026-10-02 06:34:36 +08:00
* [Feature] Add Producer-Consumer Warp Specialization and T.tma_copy() API
This PR introduces Producer-Consumer Warp Specialization for sm90+ TMA
pipelines and a new T.tma_copy() API for explicit mbarrier management.
Key changes:
1. **ProducerConsumerWarpSpecialized pass** (new):
Splits pipelined TMA loops into producer (TMA loads) and consumer
(compute) warp groups with back-pressure barriers for buffer reuse.
Works with num_stages >= 1. Producer warp (128 threads) handles
arrive_expect_tx + tma_load; consumer warps handle compute + arrive
on back-pressure barriers.
2. **T.tma_copy() API** (new):
Fire-and-forget TMA copy with a required `barrier` parameter. Unlike
T.copy() which emits producer+wait pairs, T.tma_copy() emits only
arrive_expect_tx + tma_load. User manages synchronization via
T.mbarrier_wait_parity().
3. **MultiVersionBuffer barrier expansion**:
Extended to handle `shared.barrier` scope buffers, expanding them for
pipelining (1D size multiplication instead of prepending a dimension).
Auto-computes mbarrier parity as `(k // num_stages) % 2`.
4. **Pass ordering** (phase.py):
LowerSharedBarrier now runs after MultiVersionBuffer in the TMA path
so barrier buffers retain their `shared.barrier` scope during expansion.
Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
* fix for annotate reg alloc
* Enhance debugging and performance measurement in examples
- Added print statements to output kernel source for both forward and backward sparse MLA examples, aiding in debugging.
- Updated `sparse_mla_fwd_pipelined.py` to run performance regression tests and print average latency, improving performance tracking.
- Disabled auto cache in `example_mha_sink_fwd_bhsd.py` for better control during testing.
* fix
* Enhance debugging and performance tracking in examples
- Added print statements to output kernel source in various example scripts, including `example_gqa_sink_bwd_bhsd.py`, `example_blocksparse_gemm.py`, and `example_dequant_groupedgemm_bf16_mxfp4_hopper.py`, aiding in debugging.
- Updated default `window_size` in `example_gqa_sink_bwd_bhsd.py` for improved configuration.
- Disabled auto cache in `example_gemm_schedule.py` to enhance control during testing.
- Introduced a new test for preserving protected auto-injected wait groups in the async copy optimization process.
* clean
* Remove unnecessary print statements for kernel source in example scripts
- Eliminated print statements from `example_gqa_sink_bwd_bhsd.py`, `example_blocksparse_gemm.py`, and `example_dequant_groupedgemm_bf16_mxfp4_hopper.py` to clean up output and improve readability.
- These changes streamline the examples while maintaining functionality.
* clean
* lint fix
* Refactor T.tma_copy() and barrier synchronization in examples and tests
- Updated `example_warp_specialize_gemm_copy_1_gemm_0.py` to use `T.tma_copy()` for optimized memory copying with explicit barrier management.
- Enhanced test cases in `test_tilelang_issue_tma_no_ws.py` and `test_tilelang_language_tma_copy.py` to reflect changes in barrier synchronization, ensuring proper functionality with the new T.tma_copy() API.
- Removed unnecessary print statements and adjusted barrier allocations for clarity and performance.
- Cleaned up the `test_tilelang_transform_inject_tma_barrier.py` file by deleting it, as it was no longer needed.
* Add buffer data mapping for remapped buffers in multi_version_buffer_rewriter
- Introduced a check to ensure that the data variable of a buffer is discoverable, allowing the barrier_init annotation update to find the remapped buffer correctly.
- This change enhances the functionality of the buffer remapping process, ensuring proper synchronization in the transformation pipeline.
* Implement TMA im2col load synchronization with mbarrier allocation
- Added support for allocating mbarriers to synchronize TMA im2col loads, enhancing the pipeline's barrier management.
- Introduced logic to handle mbarrier creation and integration into the TMA copy process, ensuring proper synchronization and data transfer.
- Updated the handling of the mbarrier in the TMA copy statement to improve performance and maintain functionality within the thread-gated block.
* Implement TMA im2col load synchronization with mbarrier allocation
- Added support for allocating mbarriers to synchronize TMA im2col loads, enhancing the pipeline's barrier management.
- Introduced logic to handle mbarrier initialization and integration within the TMA copy process, ensuring proper synchronization during memory operations.
- Updated the handling of the mbarrier in the im2col operation to improve performance and maintain consistency in multi-threaded environments.
* disable tma lower for test_cutedsl_barrier
* Refactor code formatting in test_tilelang_jit_cutedsl.py
- Improved readability by adjusting the formatting of the `tilelang.compile` function call.
- Enabled the main testing function in the script to run without being commented out.
* Enhance sparse_mla_fwd.py and thread_storage_sync.cc with code cleanup and new functionality
- Added a call to disable cache in the main execution block of `sparse_mla_fwd.py` for improved performance.
- Commented out unused benchmarking code in `test_sparse_mla_fwd` to streamline testing.
- Introduced new utility functions in `thread_storage_sync.cc` to analyze and optimize shared memory access patterns, enhancing synchronization management.
- Improved readability and organization of the code by restructuring and commenting on key sections.
* Add thread storage synchronization utilities and optimizations
- Introduced functions to identify shared-like scopes, buffers, and pointer variables, enhancing memory management.
- Implemented a `BarrierTransparentStmtAnalyzer` to analyze statement transparency concerning shared memory.
- Developed a `ThreadStorageSyncOptimizer` to optimize sequences of statements by removing unnecessary synchronization calls, improving performance in multi-threaded environments.
- Enhanced handling of shared memory access in various call types, ensuring proper synchronization during memory operations.
* remove optimize thread_storage_sync
* Refactor pipeline planning to enhance cp.async group scheduling
- Introduced a last-use index for cp.async groups to improve scheduling accuracy.
- Removed unnecessary tracking of first consumers, simplifying the scheduling logic.
- Updated the handling of cp.async and commit statements to ensure proper group ordering based on last-use.
- Added a new test to verify that cp.async groups are ordered correctly by their last-use index.
* lint fix
* Add CUDA requirements to TMA copy tests and update main execution call
- Added decorators to require CUDA and a minimum compute version for the `test_tma_copy_pipeline_2_stages` and `test_tma_copy_pipeline_3_stages` functions.
- Updated the main execution block to call `tilelang.testing.main()` instead of directly invoking the test functions, improving test execution management.
* Refactor attention sink examples by removing pipelined implementations
- Updated imports in `benchmark_gqa_sink_fwd.py` and `benchmark_mha_sink_fwd.py` to use the non-pipelined versions of the attention sink functions.
- Deleted the pipelined implementations of `example_gqa_sink_fwd_bhsd_wgmma_pipelined.py` and `example_mha_sink_fwd_bhsd_wgmma_pipelined.py`.
- Cleaned up regression and test files to remove references to the deleted pipelined examples, ensuring consistency across the codebase.
* performance fix
* Enhance example_blocksparse_gemm.py and inject_pipeline.cc with debugging and utility improvements
- Added a print statement to output the kernel source in `example_blocksparse_gemm.py` for debugging purposes.
- Updated `inject_pipeline.cc` to include a new utility function for stripping TMA copy write buffer attributes, improving code clarity and functionality.
- Modified `producer_consumer_ws.cc` to utilize the new utility function, enhancing the handling of TMA barriers.
- Expanded the `_compile_tvm_ffi` function in the test file to accept additional keyword arguments, improving flexibility for test configurations.
- Introduced a new test to verify that pure TMA warp specialization correctly emits barriers, ensuring compliance with expected behavior.
* lint fix
* Refactor guard variable declarations in producer_consumer_ws.cc for clarity
- Changed guard variable declarations from non-const to const references to improve performance and readability.
- This minor adjustment enhances the handling of producer guards in the context of TMA barriers.
* lint fix
* fix
* lint fix
* perf fix
* Add print statement for kernel source and update main execution in example_dequant_gemm_w4a8.py
- Introduced a print statement to output the kernel source for debugging purposes.
- Updated the main execution block to call `run_regression_perf()` and print the latency, enhancing performance measurement capabilities.
* fix
* Add debugging output and update execution flow in example_tilelang_block_sparse_attn.py
- Added a print statement to display the kernel source for better debugging.
- Updated the main execution block to disable cache and call `run_regression_perf()`, printing the resulting latency for performance measurement.
- Enhanced the producer-consumer logic in producer_consumer_ws.cc with new utility functions to manage local variable access and improve code clarity.
- Introduced a new test case in test_tilelang_issue_tma_no_ws.py to ensure consumer-only local initializations do not leak into the producer, verifying correct behavior in TMA contexts.
* Revert "Add debugging output and update execution flow in example_tilelang_block_sparse_attn.py"
This reverts commit 4995504098.
* Enhance producer-consumer logic and add tests for TMA behavior
- Updated the logic in producer_consumer_ws.cc to avoid reserving unused pre-loop barriers, improving efficiency.
- Introduced a new test case in test_tilelang_transform_producer_consumer_ws.py to verify that consumer-only local initializations do not leak into the producer.
- Added assertions in test_tilelang_issue_tma_no_ws.py to ensure correct behavior of TMA contexts.
* Add local access tracking and buffer management in producer_consumer_ws.cc
- Introduced LocalAccessSummary and LocalLiveSet structures to track read/write buffers and variable definitions.
- Implemented BufferDataToBufferCollector and LocalAccessCollector classes for efficient buffer data collection and access summary generation.
- Enhanced the logic for handling TMA copy write buffer data extraction, improving the producer-consumer interaction in TMA contexts.
* fix
* fix
* fix
* Refactor producer-consumer logic and enhance kernel execution in tilelang
- Updated the SparseFlashAttn class to improve kernel invocation clarity by separating kernel definition and execution.
- Enhanced producer-consumer workflow in producer_consumer_ws.cc by introducing optional guards for wait statements, ensuring correct behavior in concurrent contexts.
- Added utility functions to manage loop body lets, improving code readability and maintainability.
- Expanded test coverage in test_tilelang_transform_producer_consumer_ws.py to verify the preservation of guarded waits and backpressure logic.
* Enhance test coverage and execution flow in attention sink examples
- Added pytest timeout decorators to tests in test_example_attention_sink.py and test_example_blocksparse_attention.py to prevent long-running tests.
- Updated the main execution block in test_example_attention_sink.py to disable cache and directly call the sliding window test function, improving test execution clarity.
* update
* lint fix
* lint fix
* Refactor and enhance barrier handling in tilelang transformations
- Removed unused VisitStmt_ method from fuse_mbarrier_arrive_expect_tx.cc to streamline code.
- Introduced ProducerSimtCopyDetector in producer_consumer_ws.cc to identify SIMT copy patterns, improving buffer access tracking.
- Enhanced test coverage in test_tilelang_transform_fuse_mbarrier_arrive_expect_tx.py and test_tilelang_transform_producer_consumer_ws.py to validate new functionality and ensure correct barrier behavior.
* lint fix
* lint fix
* Refactor example scripts and clean up utility functions
- Removed commented-out print statements in example_warp_specialize_gemm_copy_1_gemm_0.py and example_warp_specialize_gemm_softpipe_stage2.py for cleaner code.
- Simplified buffer scope checks in utils.h by removing inline functions and directly using scope comparisons in relevant functions.
- Enhanced buffer handling in lower_tile_op.cc and producer_consumer_ws.cc by streamlining condition checks for buffer types.
* Remove unnecessary cache disabling in tilelang tests and examples
- Eliminated calls to tilelang.disable_cache() in example_gqa_decode_varlen_logits.py, test_tilelang_issue_tma_no_ws.py, and test_tilelang_language_tma_copy.py for cleaner code and improved performance.
- Updated test execution flow in test_tilelang_issue_tma_no_ws.py to directly call the main testing function.
* fix
* Add LocalAccessSummary collection for loop bounds in producer_consumer_ws.cc
- Introduced CollectExpr method in LocalAccessCollector to gather access summaries for expressions.
- Updated pre-loop liveness assignment to include variables used in pipeline loop bounds, ensuring accurate classification of scalar setups.
* enhance
* fix
* enhance
* fix
* fix
---------
Co-authored-by: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
Attention Sink
We compare with an optimized version of the official Triton implementation here.
Algorithm
Forward
The only change from vanilla FlashAttention is that sinks should be taken into consideration in the softmax, which requires an extra rescaling at the epilogue stage.
Backward
Based on detailed mathematical derivation, interestingly, the backward computation process of dQ, dK, dv is almost identical to that in vanilla FlashAttention, except for that the specific meanings of lse differ. We only need to compute dsinks additionally, which is given by:
dsink_h=-\sum_{b}\sum_{q}P_{b, h, q}Delta_{b, h, q}
where P_{b, h, q} is the proportion of sink_h in the softmax in the $b$-th block, $h$-th head and $q$-th query(row).
Benchmark of forward process
Benchmark Environment
- Hardware: NVIDIA H800
- CUDA version: 12.9
- Triton Version: 3.4.0
Results
- dtype=bfloat16
- batch_size=1, heads=64, kv_heads=8 (the setting of GPT-OSS-120B)
- Full attention is adopted.
| SEQ_LEN | headdim | Triton TFLOPs | TileLang TFLOPs | Speedup |
|---|---|---|---|---|
| 2048 | 64 | 232.98 | 281.89 | 1.21x |
| 2048 | 128 | 321.55 | 417.98 | 1.30x |
| 4096 | 64 | 280.70 | 349.47 | 1.25x |
| 4096 | 128 | 369.61 | 497.13 | 1.35x |
| 8192 | 64 | 299.04 | 385.56 | 1.29x |
| 8192 | 128 | 399.39 | 507.93 | 1.27x |
| 16384 | 64 | 309.46 | 400.62 | 1.29x |
| 16384 | 128 | 418.99 | 549.11 | 1.31x |
The backward performance will be further optimized in the future.