Skip to content
Merged
Show file tree
Hide file tree
Changes from 1 commit
Commits
Show all changes
45 commits
Select commit Hold shift + click to select a range
4de007e
Add SM120 NVFP4 blockscaled GEMM support
Jul 3, 2026
be805b1
Address SM120 NVFP4 review cleanup
Jul 8, 2026
0ce22e8
Remove SM120 private C-fragment store helpers
Jul 8, 2026
2bdaffb
Unify SM120 blockscaled MMA TIR helper
Jul 8, 2026
e84c318
Clean SM120 NVFP4 blockscaled fast path
Jul 8, 2026
c450de2
Prune SM120 NVFP4 debug lowering paths
Jul 8, 2026
850ebe2
Simplify SM120 NVFP4 blockscaled example
Jul 8, 2026
d4c646e
Remove SM120 fulltile debug macros and compile flags
Jul 9, 2026
1a1062d
Fix TensorCoreIntrinEmitter.mma base signature regression
Jul 9, 2026
66c8a30
Pin scale layout byte-compat with CuTeDSL blocked SF layout
Jul 9, 2026
3e4984e
Remove unreachable SM120 blockscale exploration code
Jul 10, 2026
be668c6
Scope shared-memory bit-exact sizing to packed scalar NVFP4
Jul 10, 2026
9c673b5
Add T.copy_ue4m3_scale_tile scale staging helper
Jul 10, 2026
790a3f0
Support M or N tail tiles in the SM120 NVFP4 example
Jul 10, 2026
f681a83
Slim tilelang/quantize/nvfp4.py
Jul 10, 2026
831652e
Converge quantizer kernels and drop example debug input modes
Jul 10, 2026
87d99c2
Keep tilelang.quantize to device helpers and format contract
Jul 10, 2026
cb81c7e
Import layout oracle from the nvfp4 submodule in the language test
Jul 10, 2026
c48e373
Group SM120 NVFP4 maint files and use native target API
Jul 21, 2026
c1ae801
Keep NVFP4 scale staging out of the tilelang.language surface
Jul 21, 2026
d3de762
Merge origin/main into nvf4-block-scale-sm120
Jul 21, 2026
1784d83
Drop stale scale-load metadata from the SM120 WS benchmark
Jul 21, 2026
317455c
Strip trailing blank lines left in gemm_op.py
Jul 21, 2026
b6f70a4
Trim the SF staging comment to the load-bearing caveat
Jul 21, 2026
c34e750
Consolidate NVFP4 tests into a single file
Jul 24, 2026
fceddda
Merge origin/main (language dialect split) into nvf4-block-scale-sm120
Jul 24, 2026
4ceee6c
Evaluate emitter dtype defaults lazily under the dialect facade
Jul 24, 2026
1fce4e9
Restore the scale-layout contract tests; keep the example CLI file fo…
Jul 24, 2026
97abc1f
Merge remote-tracking branch 'origin/main' into pr-2364-main-merge
LeiWang1999 Jul 28, 2026
70dec36
refactpr example
LeiWang1999 Jul 28, 2026
679f5f2
[SM120] Simplify NVFP4 benchmark
LeiWang1999 Jul 28, 2026
3c6cecc
[SM120] Simplify NVFP4 correctness comparison
LeiWang1999 Jul 28, 2026
5933fc9
[SM120] Inline NVFP4 quantizer pass configs
LeiWang1999 Jul 28, 2026
3e774df
[SM120] Move block-scaled MMA helper to instruction headers
LeiWang1999 Jul 28, 2026
53e3f20
[CUDA] Fix packed FP4 address codegen
LeiWang1999 Jul 28, 2026
2d69257
[SM120] Simplify NVFP4 lowering internals
LeiWang1999 Jul 28, 2026
d6a230d
[SM120] Generalize NVFP4 block-scaled MMA lowering
Rachmanino Jul 29, 2026
f29a144
[SM120] Fix package macro source test
LeiWang1999 Jul 30, 2026
bef4aff
Merge origin/main into nvf4-block-scale-sm120-pr-clean
LeiWang1999 Jul 30, 2026
9a8d1fc
[SM120] Remove redundant local FP4 access pointer test
LeiWang1999 Jul 30, 2026
a1b21a7
[SM120] Split block-scaled MMA emitter
LeiWang1999 Jul 30, 2026
ddd1098
Merge branch 'main' of https://github.com/tile-ai/tilelang into nvf4-…
LeiWang1999 Jul 30, 2026
11dae61
[SM120] Split block-scaled GEMM lowering
LeiWang1999 Jul 30, 2026
a0996a8
[SM120] Isolate block-scaled GEMM plumbing
LeiWang1999 Jul 30, 2026
acdd66b
Merge branch 'main' of https://github.com/tile-ai/tilelang into nvf4-…
LeiWang1999 Jul 30, 2026
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Prev Previous commit
Next Next commit
Merge remote-tracking branch 'origin/main' into pr-2364-main-merge
# Conflicts:
#	examples/dequantize_gemm/quantize/nvfp4.py
  • Loading branch information
LeiWang1999 committed Jul 28, 2026
commit 97abc1f400ecb85df7eaa67b889348066aebd3e1
2 changes: 1 addition & 1 deletion .github/workflows/ci.yml
Original file line number Diff line number Diff line change
Expand Up @@ -47,7 +47,7 @@ jobs:

- name: Setup Python 3.10
id: setup-pylowest
uses: actions/setup-python@v6
uses: actions/setup-python@v7
with:
python-version: "3.10"
update-environment: true
Expand Down
2 changes: 1 addition & 1 deletion .github/workflows/publish-docs.yml
Original file line number Diff line number Diff line change
Expand Up @@ -31,7 +31,7 @@ jobs:
submodules: recursive

- name: Setup Python
uses: actions/setup-python@v6
uses: actions/setup-python@v7
with:
python-version: "3.10"

Expand Down
2 changes: 1 addition & 1 deletion 3rdparty/tvm
Submodule tvm updated from db4770 to d04be2
7 changes: 6 additions & 1 deletion benchmark/matmul/benchmark_matmul_sp.py
Original file line number Diff line number Diff line change
@@ -1,13 +1,18 @@
import argparse
import itertools
import logging
from pathlib import Path
import sys

import torch
from triton.testing import do_bench

sys.path.insert(0, str(Path(__file__).resolve().parents[2] / "examples" / "gemm_sp"))

import tilelang.language as T
from tilelang.autotuner import autotune
from tilelang import jit
from tilelang.utils.sparse import get_e_factor
from sparse_utils import get_e_factor

logger = logging.getLogger(__name__)
logger.setLevel(logging.DEBUG)
Expand Down
7 changes: 6 additions & 1 deletion benchmark/matmul/benchmark_matmul_sp_compress.py
Original file line number Diff line number Diff line change
@@ -1,8 +1,13 @@
import argparse
from pathlib import Path
import sys

import torch
from tilelang.profiler import do_bench
from tilelang.utils.sparse import compress, randn_semi_sparse, randint_semi_sparse, torch_compress

sys.path.insert(0, str(Path(__file__).resolve().parents[2] / "examples" / "gemm_sp"))

from sparse_utils import compress, randn_semi_sparse, randint_semi_sparse, torch_compress

SUPPORTED_DTYPE_NAMES = ["float16", "bfloat16", "float32", "int8"]
SUPPORTED_META_DTYPE_NAMES = ["int8", "int16", "int32"]
Expand Down
6 changes: 3 additions & 3 deletions docs/deeplearning_operators/matmul_sparse.md
Original file line number Diff line number Diff line change
Expand Up @@ -38,10 +38,10 @@ To utilize sparse Tensor Cores, a dense tensor must first be **compressed** into

Both `PyTorch` and `vLLM` use `CUTLASS` as their computation backend (see references [here](https://github.com/pytorch/pytorch/blob/a8d6afb511a69687bbb2b7e88a3cf67917e1697e/aten/src/ATen/native/sparse/cuda/SparseSemiStructuredOps.cu#L47) and [here](https://github.com/vllm-project/vllm/blob/a5dd03c1ebc5e4f56f3c9d3dc0436e9c582c978f/csrc/sparse/cutlass/sparse_scaled_mm_c3x.cuh#L116)), leveraging `CUTLASS`’s built-in compressor (or reimplementing it in `PyTorch`).

A compressor is provided in `tilelang.utils.sparse`. Pass in a dense 2:4-sparse tensor and optionally a metadata dtype to get back the compressed values and metadata:
A compressor is provided with the sparse GEMM example in `examples/gemm_sp/sparse_utils.py`. Pass in a dense 2:4-sparse tensor and optionally a metadata dtype to get back the compressed values and metadata:

```python
from tilelang.utils.sparse import compress
from examples.gemm_sp.sparse_utils import compress
A_sparse, E = compress(A) # default: int16 metadata for fp16/bf16
A_sparse, E = compress(A.t().contiguous()) # compress the transposed layout
```
Expand All @@ -56,7 +56,7 @@ The default metadata dtype for fp16/bf16 is `int16` with an E-factor of 16 (one

```python
import tilelang.language as T
from tilelang.utils.sparse import get_e_factor
from examples.gemm_sp.sparse_utils import get_e_factor

def matmul_sp(
M, N, K,
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -76,7 +76,7 @@ def matmul(
- num_bits (default 4) is the bit-width of the quantized elements; storage_dtype is uint8 and num_elems_per_byte = 8 // num_bits.
- QK = K // num_elems_per_byte and Block_QK = block_K // num_elems_per_byte determine B and shared-buffer shapes.
- Asserts that K % (block_K * split) == 0; K must be divisible by block_K * split for the tiling to be valid.
- When fast_dequant is True, a valid mxfp intrinsic group (C source and function name) must be available via tilelang.quantize.get_mxfp_intrin_group.
- When fast_dequant is True, a valid mxfp intrinsic group (C source and function name) must be available via quantize.get_mxfp_intrin_group.
- The kernel launches a 2D grid over ceildiv(N, block_N) and ceildiv(M, block_M) and uses `threads` threads per block with `num_stages` pipeline stages.

Parameters that alter kernel layout/behavior (brief):
Expand All @@ -101,7 +101,7 @@ def matmul(
B_dequantize_shared_shape = (block_N, block_K)
assert K % (block_K * split) == 0

from tilelang.quantize import get_mxfp_intrin_group
from quantize import get_mxfp_intrin_group

# fast_dequant_bf16_fp4_twiddling
# It requires that the 2 consecutive uint8 elements (16bits) contains 4 fp4 elements in a bit-twiddling way.
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -19,7 +19,7 @@
from tvm import tirx
import torch

from tilelang.quantize import get_mxfp_intrin_group
from quantize import get_mxfp_intrin_group
from dequantize_utils import torch_convert_bit_twiddling, torch_convert


Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -154,7 +154,7 @@ def matmul(
B_dequantize_shared_shape = (block_N, block_K)
assert K % (block_K * split) == 0

from tilelang.quantize import get_mxfp_intrin_group
from quantize import get_mxfp_intrin_group

# fast_dequant_bf16_fp4_twiddling
mxfp_intrin_info = get_mxfp_intrin_group(
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -28,7 +28,7 @@ def matmul(
threads,
num_bits=4,
):
from tilelang.quantize import _tir_packed_to_unsigned_convert
from quantize import _tir_packed_to_unsigned_convert

num_elems_per_byte = 8 // num_bits
storage_dtype = T.int8
Expand Down
8 changes: 4 additions & 4 deletions examples/dequantize_gemm/example_dequant_gemv_fp16xint4.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,7 +3,7 @@
from typing import Optional, Callable, Any
import torch
from tilelang import DataType
from tilelang.quantize import (
from quantize import (
_tir_packed_int_to_int_convert,
)

Expand Down Expand Up @@ -58,7 +58,7 @@ def dequantize_gemv(
if fast_decoding is True:
# Lazy import to decrease the startup time
# as intrin registry may take a while to load
from tilelang.quantize import get_lop3_intrin_group
from quantize import get_lop3_intrin_group

lop3_intrin_info = get_lop3_intrin_group(
out_dtype=in_dtype,
Expand Down Expand Up @@ -199,7 +199,7 @@ def main() -> None:
C = torch.zeros(M, N, dtype=getattr(torch, accum_dtype)).cuda()

if fast_decoding:
from tilelang.quantize.utils import interleave_weight
from quantize.utils import interleave_weight

qB = interleave_weight(qB, num_bits, in_dtype)
kernel(A, qB, C)
Expand Down Expand Up @@ -261,7 +261,7 @@ def run_regression_perf():
C = torch.zeros(M, N, dtype=getattr(torch, accum_dtype)).cuda()

if fast_decoding:
from tilelang.quantize.utils import interleave_weight
from quantize.utils import interleave_weight

qB = interleave_weight(qB, num_bits, in_dtype)
kernel(A, qB, C)
Expand Down
File renamed without changes.
File renamed without changes.
File renamed without changes.
File renamed without changes.
4 changes: 2 additions & 2 deletions examples/gemm_sm120/sm120_nvfp4_blockscaled_gemm.py
Comment thread
LeiWang1999 marked this conversation as resolved.
Original file line number Diff line number Diff line change
Expand Up @@ -25,7 +25,7 @@
import tilelang
import tilelang.language as T
from tilelang.profiler import do_bench
from tilelang.quantize import swizzle_blockscaled_chunk_kmajor_scale_words
from examples.dequantize_gemm.quantize import swizzle_blockscaled_chunk_kmajor_scale_words


_SM120_SCALE_LAYOUT = "blockscaled_chunk_kmajor"
Expand All @@ -50,7 +50,7 @@ def sm120_nvfp4_blockscaled_gemm(
# Tail tiles: TMA loads zero-fill out-of-bounds rows and the C store is
# predicated, so M (or N) may be arbitrary as long as the other dimension
# is a multiple of its tile. The scale source is padded to full 128-row
# tiles (the packers in tilelang.quantize do this automatically).
# tiles (the helpers in examples.dequantize_gemm.quantize do this automatically).
# N must keep 16-byte aligned bf16 rows, the same contiguous-dim rule as
# CuTeDSL's is_valid_tensor_alignment.
assert N % 8 == 0, "N must be a multiple of 8 (16-byte aligned output rows)"
Expand Down
2 changes: 1 addition & 1 deletion examples/gemm_sp/example_gemm_sp.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,7 +3,7 @@
import tilelang
import tilelang.language as T

from tilelang.utils.sparse import compress, randn_semi_sparse, get_e_factor
from sparse_utils import compress, randn_semi_sparse, get_e_factor
from tilelang.profiler import do_bench

import torch
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -5,7 +5,7 @@

# 2:4 sparsity layout metadata lives in the leaf module so the mma_sp code
# generator can import it without dragging this torch/@tilelang.jit recipe into
# the bootstrap closure. Re-exported here for backward compatibility.
# the bootstrap closure. Import it here so this example uses the same contract.
from tilelang.cuda.intrinsics.sparse_layout import ( # noqa: F401
GROUP_CONFIG,
get_e_factor,
Expand Down
Comment thread
LeiWang1999 marked this conversation as resolved.
Original file line number Diff line number Diff line change
Expand Up @@ -29,14 +29,14 @@
import tilelang.language as T
from tilelang.carver.arch import driver
from tilelang.profiler import do_bench
from tilelang.quantize import (
from examples.dequantize_gemm.quantize import (
swizzle_blockscaled_chunk_kmajor_scale_words,
unswizzle_blockscaled_chunk_kmajor_scale_words,
)

sys.path.insert(0, str(Path(__file__).resolve().parent))
from tilelang_nvfp4_quantizer import tilelang_quantize_bf16_to_nvfp4_blockscaled # noqa: E402
from tilelang.quantize.nvfp4 import (
from examples.dequantize_gemm.quantize.nvfp4 import (
blockscaled_chunk_kmajor_tile_source_coords,
blockscaled_chunk_kmajor_word_offset,
decode_ue4m3_scale_bytes,
Expand Down Expand Up @@ -159,7 +159,7 @@ def tilelang_nvfp4_blockscaled_gemm(
def copy_blockscaled_chunk_kmajor_scale_tile(SF, SF_shared, tile_row, block_rows, ko, stage, tx):
# Producer-warp-group staging: 128 producer lanes stride the tile,
# with the layout addressing shared via
# tilelang.quantize.nvfp4.blockscaled_chunk_kmajor_tile_source_coords.
# examples.dequantize_gemm.quantize.nvfp4.blockscaled_chunk_kmajor_tile_source_coords.
scale_tile_words = block_rows * sf_words_per_block_k
scale_lane = tx - 256
for scale_iter in T.serial((scale_tile_words + 127) // 128):
Expand Down
6 changes: 3 additions & 3 deletions maint/gemm/gemm_sm120/tilelang_nvfp4_quantizer.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,13 +2,13 @@

This is the performance quantizer used by the SM120 NVFP4 GEMM maint
benchmark. It builds on the device-side format helpers and layout contract in
``tilelang.quantize.nvfp4``; the library itself only ships those helpers plus
a torch reference implementation (``quantize_bf16_to_nvfp4_blockscaled``).
``examples.dequantize_gemm.quantize.nvfp4`` alongside the torch reference
implementation (``quantize_bf16_to_nvfp4_blockscaled``).
"""

import tilelang
import tilelang.language as T
from tilelang.quantize.nvfp4 import (
from examples.dequantize_gemm.quantize.nvfp4 import (
_BLOCKSCALED_CHUNK_ROWS,
_BLOCKSCALED_CHUNK_WORDS,
_FP4_E2M1_MAX,
Expand Down
2 changes: 1 addition & 1 deletion pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -58,7 +58,7 @@ nvcc = [
# On Windows, ship nvrtc through the runtime deps so `pip install tilelang`
# works without a host CUDA install (TileLang JITs CUDA kernels via nvrtc).
"nvidia-cuda-nvrtc>=13; platform_system == 'Windows'",
# cuRAND is required when kernels use tilelang.language.random
# cuRAND is required when kernels use tilelang.cuda.language.random
"nvidia-curand>=10.4; platform_system == 'Windows'",
]

Expand Down
Empty file.
37 changes: 37 additions & 0 deletions reproducer_vec_type.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,37 @@
"""
Minimal reproducer for the missing vec_type arithmetic operators in common.h.

Without the fix, compiling a CPU kernel with `c` target fails:
error: no match for 'operator-' (operand types are 'float4' and 'float4')

Requires: tilelang, torch
"""

import tilelang
import tilelang.language as T
import torch


@T.prim_func
def vec_add(
A: T.Tensor((256,), "float32"),
B: T.Tensor((256,), "float32"),
C: T.Tensor((256,), "float32"),
):
for i in T.Parallel(256):
C[i] = A[i] + B[i]


def main():
f = tilelang.compile(vec_add,>
A = torch.randn(256)
B = torch.randn(256)
C = torch.zeros(256)
f(A, B, C)
ref = A + B
assert (C - ref).abs().max().item() < 1e-5
print("OK: vec_add compiled and ran successfully")


if __name__ == "__main__":
main()
2 changes: 2 additions & 0 deletions src/backend/common/op/finalize_reducer.h
Original file line number Diff line number Diff line change
Expand Up @@ -6,6 +6,7 @@
#ifndef TVM_TL_BACKEND_COMMON_OP_FINALIZE_REDUCER_H_
#define TVM_TL_BACKEND_COMMON_OP_FINALIZE_REDUCER_H_

#include "backend/common/op/reduce.h"
#include "op/finalize_reducer.h"
#include "support/check.h"

Expand Down Expand Up @@ -52,6 +53,7 @@ template <typename Impl> struct FinalizeReducerLowerer {
auto op_str = op_names[static_cast<int>(op.op)];

int reducing_threads = extent;
reduce::CheckAllReduceWidth(reducing_threads, 1, "tl.finalize_reducer");
auto thread_offset = lower_args.thread_bounds->min;

int64_t layout_batch_size = 1;
Expand Down
23 changes: 23 additions & 0 deletions src/backend/common/op/reduce.h
Original file line number Diff line number Diff line change
Expand Up @@ -104,6 +104,25 @@ inline int GetPreferedVectorizedSize(DataType dt,
return 1;
}

inline void CheckAllReduceWidth(int reducing_threads, int scale,
const char *op_name) {
ICHECK_GT(reducing_threads, 0)
<< op_name << ": AllReduce threads must be positive, got "
<< reducing_threads;
ICHECK_GT(scale, 0) << op_name << ": AllReduce scale must be positive, got "
<< scale;
ICHECK_EQ(reducing_threads % scale, 0)
<< op_name << ": AllReduce threads (" << reducing_threads
<< ") must be divisible by scale (" << scale << ")";
int logical_width = reducing_threads / scale;
int shift = 0;
ICHECK(tirx::is_const_power_of_two_integer(Integer(logical_width), &shift))
<< op_name << ": XOR-butterfly AllReduce requires logical_width "
<< "(threads / scale) to be a positive power of two, got "
<< logical_width << " (threads=" << reducing_threads
<< ", scale=" << scale << ")";
}

inline PrimExpr MakeInitValue(const ReduceOpNode &op, int vsize = 1) {
auto dst_dtype = op.dst->dtype;
auto is_int = dst_dtype.is_int();
Expand Down Expand Up @@ -909,6 +928,8 @@ template <typename Impl> struct ReduceLowerer {

for (const auto &thread_step : reduce_plan.thread_steps) {
int reducing_threads = thread_step.ReducingThreads();
reduce::CheckAllReduceWidth(reducing_threads, thread_step.scale,
"tl.reduce");
int block_threads =
static_cast<int>(*as_const_int(lower_args.thread_bounds->extent));
auto thread_offset = lower_args.thread_bounds->min;
Expand Down Expand Up @@ -1101,6 +1122,8 @@ template <typename Impl> struct ReduceLowerer {

for (const auto &thread_step : reduce_plan.thread_steps) {
int reducing_threads = thread_step.ReducingThreads();
reduce::CheckAllReduceWidth(reducing_threads, thread_step.scale,
"tl.reduce");
auto thread_offset = lower_args.thread_bounds->min;
std::string allreduce = Impl::MakeScalarAllReduce(
reduce::MakeCodegenReducer(op).value(), reducing_threads,
Expand Down
Loading
You are viewing a condensed version of this merge commit. You can view the full changes here.