-
Notifications
You must be signed in to change notification settings - Fork 788
[Common] Experimental CuTeDSL MXFP8 backends in C++ via TVM-FFI #3137
New issue
Have a question about this project? Sign up for a free GitHub account to open an issue and contact its maintainers and the community.
By clicking “Sign up for GitHub”, you agree to our terms of service and privacy statement. We’ll occasionally send you account related emails.
Already on GitHub? Sign in to your account
Open
kainzhong
wants to merge
55
commits into
NVIDIA:main
Choose a base branch
from
kainzhong:cutedsl_mxfp8_common
base: main
Could not load branches
Branch not found: {{ refName }}
Loading
Could not load tags
Nothing to show
Loading
Are you sure you want to change the base?
Some commits from the old base branch may be removed from the timeline,
and old review comments may become outdated.
Open
Changes from 54 commits
Commits
Show all changes
55 commits
Select commit
Hold shift + click to select a range
6c43f15
start draft
kainzhong a4f9bab
remove benchmark scripts
kainzhong 96136c2
[pre-commit.ci] auto fixes from pre-commit.com hooks
pre-commit-ci[bot] a9e8df4
add license
kainzhong e036abf
fix
kainzhong 15de391
make cutlass dsl required
kainzhong 0d4a2ba
fix linting errors
kainzhong 01b28b3
[pre-commit.ci] auto fixes from pre-commit.com hooks
pre-commit-ci[bot] 94501c6
warn failed compilation
kainzhong a95ccf7
[pre-commit.ci] auto fixes from pre-commit.com hooks
pre-commit-ci[bot] 5cd2d47
make tvm-ffi common dependency
kainzhong bf0cd8d
use a higher version for cutlass-dsl
kainzhong 1650c05
less comment
kainzhong b332fb7
skip zeroing buffer if noop flag set
kainzhong 631bf05
refactor activations
kainzhong 81c52ef
fix
kainzhong 3fc36ff
lint
kainzhong f41e448
fix
kainzhong 2af3d4e
[pre-commit.ci] auto fixes from pre-commit.com hooks
pre-commit-ci[bot] 43c7b5f
nit
kainzhong 84cfcc3
maybe it's better to prepare tvm-ffi & cutedsl kernels in common/__in…
kainzhong 815ec0f
[pre-commit.ci] auto fixes from pre-commit.com hooks
pre-commit-ci[bot] c9c1c4f
nit
kainzhong 2334ca7
[pre-commit.ci] auto fixes from pre-commit.com hooks
pre-commit-ci[bot] c4291ad
fix
kainzhong 26a2847
[pre-commit.ci] auto fixes from pre-commit.com hooks
pre-commit-ci[bot] 00784d5
fix
kainzhong 752e962
fix
kainzhong aaf6518
fix
kainzhong 3d300ce
also flush colwise scale in SMEM
kainzhong b3b0a7f
use explicit fp8 types
kainzhong 5e16681
[pre-commit.ci] auto fixes from pre-commit.com hooks
pre-commit-ci[bot] 855ba5b
nit
kainzhong a6010ca
add checks
kainzhong a2dcda7
refactor
kainzhong 13bc7ff
nit
kainzhong 812b81e
nit
kainzhong 428bc5a
nit
kainzhong 5044d27
fi
kainzhong bc804b6
[pre-commit.ci] auto fixes from pre-commit.com hooks
pre-commit-ci[bot] 5172938
nit
kainzhong 19b61ff
nit
kainzhong 5f5099b
skip comparision test if cutedsl disabled
kainzhong 1f0a686
fix
kainzhong cb2cdcb
Merge pull request #3 from janekb04/cutedsl_mxfp8_common
janekb04 2ea3eb5
add test to the qa script
kainzhong 4815598
warn only once if dlopen fails
kainzhong 3c93710
fix
kainzhong 7a10bb9
let tests cover swizzling
kainzhong a26638a
warn not chosen cutedsl kernel during testing
kainzhong e054993
improve lazyload so now the cache is uint32
kainzhong 207b043
[pre-commit.ci] auto fixes from pre-commit.com hooks
pre-commit-ci[bot] eebf0e7
let tests be verbose about warnings so we know if cutedsl kernel is n…
kainzhong 09a70a2
support swizzled scale for specialized kernel but disable it
kainzhong 1c18fd7
fix
kainzhong File filter
Filter by extension
Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
There are no files selected for viewing
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
| Original file line number | Diff line number | Diff line change |
|---|---|---|
| @@ -0,0 +1,265 @@ | ||
| # Copyright (c) 2022-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. | ||
| # | ||
| # See LICENSE for license information. | ||
|
|
||
| """Cross-backend bit-exactness tests for the CuTeDSL MXFP8 quantize kernels.""" | ||
|
|
||
| import ctypes | ||
| import os | ||
| from typing import Callable, NamedTuple, Optional | ||
|
|
||
| import pytest | ||
| import torch | ||
|
|
||
| import transformer_engine.pytorch as te | ||
| import transformer_engine_torch as tex | ||
| import tvm_ffi | ||
|
|
||
| from transformer_engine.common import _get_shared_object_file | ||
| from transformer_engine.pytorch import MXFP8Quantizer | ||
|
|
||
| recipe_available, reason_for_no_recipe = te.is_mxfp8_available(return_reason=True) | ||
|
|
||
| # The already-loaded core lib (dlopen refcounts: this returns the same handle, | ||
| # so the call mutates the same dispatcher singleton the quantize ops read). | ||
| CORE_LIB = ctypes.CDLL(str(_get_shared_object_file("core"))) | ||
| if not hasattr(CORE_LIB, "nvte_set_cutedsl_quant_backend"): | ||
| raise RuntimeError( | ||
| "libtransformer_engine.so lacks nvte_set_cutedsl_quant_backend -- rebuild the " | ||
| "Transformer Engine core library." | ||
| ) | ||
|
|
||
| # The CuTeDSL entrypoint is registered only when NVTE_ENABLE_CUTEDSL_QUANT_BACKEND | ||
| # is set (see common/__init__.py); without it there is nothing to compare against | ||
| # the CUDA path, so skip these runs. | ||
| cutedsl_enabled = os.environ.get("NVTE_ENABLE_CUTEDSL_QUANT_BACKEND", "0") != "0" | ||
| pytestmark = pytest.mark.skipif( | ||
| not (recipe_available and cutedsl_enabled), | ||
| reason=reason_for_no_recipe or "NVTE_ENABLE_CUTEDSL_QUANT_BACKEND is not set", | ||
| ) | ||
|
|
||
|
|
||
| class Fusion(NamedTuple): | ||
| """An ActivationType from the C++ test: its tex ops per ProcessingMethod and | ||
| the activation desc used in the CuTeDSL config key.""" | ||
|
|
||
| name: str | ||
| act: Optional[Callable] # CAST_ACT: act(x, quantizer) | ||
| dact: Optional[Callable] # CAST_DACT: dact(grad, act_input, quantizer) | ||
| dbias_dact: Optional[Callable] # CAST_DBIAS_DACT: dbias_dact(grad, act_input, quantizer) | ||
| desc: str | ||
|
|
||
|
|
||
| # --- Case matrix, mirroring test_cast_mxfp8.cu --- | ||
| # The C++ multi-dim sizes flattened to the 2D (rows, cols) the kernels see: | ||
| # {8,32,1024} -> (256, 1024), {16,8,4,512} -> (512, 512). The C++ list also has | ||
| # non-32-divisible shapes ({1,16}, {16,48}, {993,512}, {1024}) that exercise the | ||
| # CUDA kernels' partial-block edges; the CuTeDSL backend's contract is | ||
| # 32-divisible flat dims (the dispatcher falls back to CUDA otherwise), so those | ||
| # cases are omitted here rather than vacuously comparing CUDA against CUDA. | ||
| MATRIX_SIZES = [ | ||
| (128, 128), | ||
| (256, 1024), | ||
| (512, 512), | ||
| (8192, 7168), | ||
| ] | ||
| # (block_rows, block_cols): (1,32)=rowwise, (32,1)=colwise, (32,32)=both. | ||
| BLOCK_SIZES = [(1, 32), (32, 1), (32, 32)] | ||
| # Only GeLU activation tests are used (SiLU/ReLU/QGeLU/SReLU commented out | ||
| # in the C++ test as well). | ||
| IDENTITY = Fusion("Identity", None, None, None, "none") | ||
| GELU = Fusion("GeLU", tex.gelu, tex.dgelu, tex.dbias_dgelu, "gelu") | ||
| # SILU = Fusion("SiLU", tex.silu, tex.dsilu, tex.dbias_dsilu, "silu") | ||
| # RELU = Fusion("ReLU", tex.relu, tex.drelu, tex.dbias_drelu, "relu") | ||
| # QGELU = Fusion("QGeLU", tex.qgelu, tex.dqgelu, tex.dbias_dqgelu, "qgelu") | ||
| # SRELU = Fusion("SReLU", tex.srelu, tex.dsrelu, tex.dbias_dsrelu, "srelu") | ||
|
|
||
| # Valid (ProcessingMethod, ActivationType) pairs. The C++ test crosses the two | ||
| # axes and GTEST_SKIPs the mismatched half; only the meaningful pairs are | ||
| # generated here. A newly enabled activation adds its ACT/DACT/DBIAS_DACT pairs. | ||
| METHOD_FUSION_CASES = [ | ||
| ("CAST_ONLY", IDENTITY), | ||
| ("CAST_DBIAS", IDENTITY), | ||
| ("CAST_ACT", GELU), | ||
| ("CAST_DACT", GELU), | ||
| ("CAST_DBIAS_DACT", GELU), | ||
| ] | ||
| METHOD_FUSION_IDS = [f"{m}X{f.name}" for m, f in METHOD_FUSION_CASES] | ||
| IN_DTYPES = [torch.float32, torch.bfloat16, torch.float16] | ||
| FP8_DTYPES = [tex.DType.kFloat8E4M3, tex.DType.kFloat8E5M2] | ||
|
|
||
| # Description strings for pytest case ids and the CuTeDSL config key. | ||
| DTYPE_TO_STR = {torch.float32: "fp32", torch.bfloat16: "bf16", torch.float16: "fp16"} | ||
| FP8_TO_STR = {tex.DType.kFloat8E4M3: "e4m3", tex.DType.kFloat8E5M2: "e5m2"} | ||
|
|
||
| # Scale layout: linear, or the GEMM-swizzled layout cuBLAS consumes | ||
| # (quantizer.optimize_for_gemm -> MXFP8QuantConfig::swizzled). | ||
| SWIZZLE_MODES = [False, True] | ||
|
|
||
| get_shape_id = lambda s: f"{s[0]}x{s[1]}" | ||
| get_block_id = lambda b: f"{b[0]}x{b[1]}" | ||
| get_dtype_id = DTYPE_TO_STR.get | ||
| get_fp8_id = FP8_TO_STR.get | ||
| get_swizzle_id = lambda s: "swizzled" if s else "linear" | ||
|
|
||
|
|
||
| def set_cutedsl_backend(enabled): | ||
| CORE_LIB.nvte_set_cutedsl_quant_backend(1 if enabled else 0) | ||
|
|
||
|
|
||
| @pytest.fixture(scope="module", autouse=True) | ||
| def _restore_backend_choice_from_env(): | ||
| """Restore the flag that decides the CuTeDSL / CUDA backend choice when this pytest module is done.""" | ||
| yield | ||
| flag = os.getenv("NVTE_ENABLE_CUTEDSL_QUANT_BACKEND") | ||
| set_cutedsl_backend(flag is not None and not flag.startswith("0")) | ||
|
|
||
|
|
||
| def generate_inputs(M, N, in_dtype, seed=0): | ||
| g = torch.Generator(device="cuda").manual_seed(seed) | ||
| x = torch.empty(M, N, dtype=in_dtype, device="cuda").uniform_(-2.0, 1.0, generator=g) | ||
| ain = torch.empty(M, N, dtype=in_dtype, device="cuda").uniform_(-2.0, 1.0, generator=g) | ||
| return x, ain | ||
|
|
||
|
|
||
| def run_quantize(method, act, x, ain, rowwise, columnwise, fp8_dtype, swizzled): | ||
| """Quantize via the public dispatch; returns (mxfp8_tensor, dbias_or_None).""" | ||
| q = MXFP8Quantizer(fp8_dtype=fp8_dtype, rowwise=rowwise, columnwise=columnwise) | ||
| # Emit scales in the GEMM-swizzled layout (MXFP8QuantConfig::swizzled). | ||
| q.optimize_for_gemm = swizzled | ||
| if method == "CAST_ONLY": | ||
| return q(x), None | ||
| if method == "CAST_DBIAS": | ||
| db, out = tex.bgrad_quantize(x, q) | ||
| return out, db | ||
| if method == "CAST_ACT": | ||
| return act.act(x, q), None | ||
| if method == "CAST_DACT": | ||
| return act.dact(x, ain, q), None | ||
| if method == "CAST_DBIAS_DACT": | ||
| db, out = act.dbias_dact(x, ain, q) | ||
| return out, db | ||
| raise ValueError(f"unknown method {method!r}") | ||
|
|
||
|
|
||
| def get_cfg_key(method, act, in_dtype, fp8_dtype, rowwise, colwise, swizzled): | ||
| """Mirror of MXFP8QuantConfig::to_key (quantize_mxfp8_cutedsl.cuh): the name the CuTeDSL backend registers its compiled kernel under for this config. | ||
| Used to check if the CuTeDSL implmentation is registered | ||
| """ | ||
| with_dbias = method in ("CAST_DBIAS", "CAST_DBIAS_DACT") | ||
| with_dact = method in ("CAST_DACT", "CAST_DBIAS_DACT") | ||
| with_act = method == "CAST_ACT" | ||
| desc = "none" | ||
| if with_act: | ||
| desc = act.desc | ||
| elif with_dact: | ||
| desc = f"d{act.desc}" | ||
| # with_amax is hardcoded to False for now because there is no way to obtain this value and validate in python | ||
| flags = (rowwise, colwise, swizzled, False, with_dbias, with_dact, with_act) | ||
| return ( | ||
| "cutedsl_mxfp8_" | ||
| + DTYPE_TO_STR[in_dtype] | ||
| + "_" | ||
| + FP8_TO_STR[fp8_dtype] | ||
| + "_" | ||
| + "_".join("1" if f else "0" for f in flags) | ||
| + "_" | ||
| + desc | ||
| ) | ||
|
|
||
|
|
||
| def extract_quantized_output(out, rowwise, columnwise, swizzled): | ||
| """Extract the bytes to compare between backends. | ||
|
|
||
| Linear layout: the scale padding is uninitialized, so only the meaningful region is | ||
| compared. Swizzled layout: the meaningful scales are scattered by the swizzle, so the | ||
| top-left slice is meaningless; both backends zero the padding (see zero_scales_kernel), | ||
| so compare the whole buffer instead -- which also covers the padding-zeroing itself. | ||
| """ | ||
| parts = {} | ||
| if rowwise: | ||
| d = out._rowwise_data.view(torch.uint8) | ||
| M, N = d.shape | ||
| parts["rowwise data"] = d.clone() | ||
| s = out._rowwise_scale_inv | ||
| parts["rowwise scales"] = (s if swizzled else s[:M, : (N + 31) // 32]).clone() | ||
| if columnwise: | ||
| d = out._columnwise_data.view(torch.uint8) | ||
| M, N = d.shape | ||
| parts["colwise data"] = d.clone() | ||
| s = out._columnwise_scale_inv | ||
| parts["colwise scales"] = (s if swizzled else s[: (M + 31) // 32, :N]).clone() | ||
| return parts | ||
|
|
||
|
|
||
| def run_test_case(method, act, shape, block_size, in_dtype, fp8_dtype, swizzled=False): | ||
| """Assert the CuTeDSL and CUDA backends produce bit-identical outputs for the | ||
| same input and config. | ||
| """ | ||
| M, N = shape | ||
| rowwise = block_size[1] != 1 | ||
| columnwise = block_size[0] != 1 | ||
| x, act_input = generate_inputs(M, N, in_dtype) | ||
|
|
||
| set_cutedsl_backend(False) | ||
| out_cuda, dbias_cuda = run_quantize( | ||
| method, act, x, act_input, rowwise, columnwise, fp8_dtype, swizzled | ||
| ) | ||
| cuda_output = extract_quantized_output(out_cuda, rowwise, columnwise, swizzled) | ||
|
|
||
| set_cutedsl_backend(True) | ||
| try: | ||
| out_cutedsl, dbias_cutedsl = run_quantize( | ||
| method, act, x, act_input, rowwise, columnwise, fp8_dtype, swizzled | ||
| ) | ||
| cutedsl_output = extract_quantized_output(out_cutedsl, rowwise, columnwise, swizzled) | ||
| finally: | ||
| set_cutedsl_backend(False) | ||
|
|
||
| # Guard against a silent CUDA fallback: every config in the matrix is one the | ||
| # CuTeDSL backend supports, so its kernel must have been registered under the | ||
| # config key. If not, the backend rejected or missed the config and the | ||
| # comparison above was CUDA vs CUDA. | ||
| key = get_cfg_key(method, act, in_dtype, fp8_dtype, rowwise, columnwise, swizzled) | ||
| assert tvm_ffi.get_global_func(key, allow_missing=True) is not None, ( | ||
| f"CuTeDSL kernel not registered for {key}; the CuTeDSL backend fell back " | ||
| "to CUDA and this case compared CUDA against itself" | ||
| ) | ||
|
|
||
| layout = "swizzled" if swizzled else "linear" | ||
| tag = f"{method}/{act.name}/{M}x{N}/{DTYPE_TO_STR[in_dtype]}/{FP8_TO_STR[fp8_dtype]}/{layout}" | ||
| for name, cuda_bytes in cuda_output.items(): | ||
| assert torch.equal( | ||
| cutedsl_output[name], cuda_bytes | ||
| ), f"{tag}: {name} differ between backends" | ||
| if dbias_cuda is not None: | ||
| torch.testing.assert_close(dbias_cutedsl, dbias_cuda) | ||
|
|
||
|
|
||
| # Test cases with only cast kernels (mirrors C++ test's OperatorTest_FusedCastMXFP8_CastOnly). | ||
| @pytest.mark.parametrize("shape", MATRIX_SIZES, ids=get_shape_id) | ||
| @pytest.mark.parametrize("block_size", BLOCK_SIZES, ids=get_block_id) | ||
| @pytest.mark.parametrize("in_dtype", IN_DTYPES, ids=get_dtype_id) | ||
| @pytest.mark.parametrize("fp8_dtype", FP8_DTYPES, ids=get_fp8_id) | ||
| @pytest.mark.parametrize("swizzled", SWIZZLE_MODES, ids=get_swizzle_id) | ||
| def test_cast_only(swizzled, fp8_dtype, in_dtype, block_size, shape): | ||
| run_test_case("CAST_ONLY", IDENTITY, shape, block_size, in_dtype, fp8_dtype, swizzled) | ||
|
|
||
|
|
||
| # Test cases with varying matrix shapes and block shapes | ||
| # (OperatorTest_FusedCastMXFP8_Sizes). | ||
| @pytest.mark.parametrize("shape", MATRIX_SIZES, ids=get_shape_id) | ||
| @pytest.mark.parametrize("block_size", BLOCK_SIZES, ids=get_block_id) | ||
| @pytest.mark.parametrize("method,act", METHOD_FUSION_CASES, ids=METHOD_FUSION_IDS) | ||
| @pytest.mark.parametrize("swizzled", SWIZZLE_MODES, ids=get_swizzle_id) | ||
| def test_sizes(swizzled, method, act, block_size, shape): | ||
| run_test_case(method, act, shape, block_size, torch.bfloat16, tex.DType.kFloat8E4M3, swizzled) | ||
|
|
||
|
|
||
| # Test cases with varying dtypes (OperatorTest_FusedCastMXFP8_Dtypes). | ||
| @pytest.mark.parametrize("in_dtype", IN_DTYPES, ids=get_dtype_id) | ||
| @pytest.mark.parametrize("fp8_dtype", FP8_DTYPES, ids=get_fp8_id) | ||
| @pytest.mark.parametrize("method,act", METHOD_FUSION_CASES, ids=METHOD_FUSION_IDS) | ||
| @pytest.mark.parametrize("swizzled", SWIZZLE_MODES, ids=get_swizzle_id) | ||
| def test_dtypes(swizzled, method, act, fp8_dtype, in_dtype): | ||
| run_test_case(method, act, (256, 384), (32, 32), in_dtype, fp8_dtype, swizzled) | ||
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
| Original file line number | Diff line number | Diff line change |
|---|---|---|
| @@ -0,0 +1,24 @@ | ||
| # Copyright (c) 2022-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. | ||
| # | ||
| # See LICENSE for license information. | ||
|
|
||
| """CuTeDSL kernels for Transformer Engine. | ||
| To expose CuTeDSL kernels to C++, they should be registered via `register_cutedsl_backends()`. | ||
| They should provide a string function name which can be used to retrieve the corresponding CuTeDSL kernel function via TVM-FFI. | ||
| """ | ||
|
|
||
| import tvm_ffi | ||
|
|
||
| from transformer_engine.common.CuTeDSL.cast.mxfp8.quantize_mxfp8 import ( | ||
| get_mxfp8_quantization_function, | ||
| ) | ||
|
|
||
|
|
||
| def register_cutedsl_backends(): | ||
| """Register all available CuTeDSL backends for on-demand compilation via TVM-FFI. | ||
| The C++ dispatcher retrieves them by the names defined here. | ||
| """ | ||
| tvm_ffi.register_global_func( | ||
| "get_mxfp8_quantization_function", get_mxfp8_quantization_function, override=True | ||
| ) |
Oops, something went wrong.
Oops, something went wrong.
Add this suggestion to a batch that can be applied as a single commit.
This suggestion is invalid because no changes were made to the code.
Suggestions cannot be applied while the pull request is closed.
Suggestions cannot be applied while viewing a subset of changes.
Only one suggestion per line can be applied in a batch.
Add this suggestion to a batch that can be applied as a single commit.
Applying suggestions on deleted lines is not supported.
You must change the existing code in this line in order to create a valid suggestion.
Outdated suggestions cannot be applied.
This suggestion has been applied or marked resolved.
Suggestions cannot be applied from pending reviews.
Suggestions cannot be applied on multi-line comments.
Suggestions cannot be applied while the pull request is queued to merge.
Suggestion cannot be applied right now. Please check back later.
There was a problem hiding this comment.
Choose a reason for hiding this comment
The reason will be displayed to describe this comment to others. Learn more.
import tvm_ffiraisesImportErrorfor users without the packageThe test file imports
tvm_ffiat module level before thepytestmarkskip guard is evaluated. Pytest collects all test modules regardless of environment; users withoutapache-tvm-ffiinstalled will see a collection error instead of a clean skip. The import should be moved inside the test body or placed under atry/except ImportErrorguard that setstvm_ffi_available = False, similar to how the test already conditionally setscutedsl_enabled.