-
Notifications
You must be signed in to change notification settings - Fork 805
[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
96
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 all commits
Commits
Show all changes
96 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 7942da1
nit
kainzhong 0c83fb3
fix inline ptx
kainzhong b111740
more fix
kainzhong 547b84a
revert to mul
kainzhong 53a8ab5
nit
kainzhong fd9b1d8
[pre-commit.ci] auto fixes from pre-commit.com hooks
pre-commit-ci[bot] 6e02798
add jax tests
kainzhong ce9ef37
also ignore noop when dbias
kainzhong 36f8129
[pre-commit.ci] auto fixes from pre-commit.com hooks
pre-commit-ci[bot] 12b478d
Merge branch 'main' into cutedsl_mxfp8_common
kainzhong d9b6733
relax dim check for the kernels
kainzhong 2e4bf7c
reject 2D quant
kainzhong c078b31
run more tests with CuTeDSL path
kainzhong 92de056
[pre-commit.ci] auto fixes from pre-commit.com hooks
pre-commit-ci[bot] 428913e
make rowwise only handle divisible cases
kainzhong 9c328d9
fallback to CUDA if not TMA aligned
kainzhong 0ad3831
[pre-commit.ci] auto fixes from pre-commit.com hooks
pre-commit-ci[bot] c1aef2f
specialize skip masking version
kainzhong 208ddea
nit
kainzhong 775ec31
[pre-commit.ci] auto fixes from pre-commit.com hooks
pre-commit-ci[bot] 39d69c9
add type annotations
kainzhong 13c3fa1
[pre-commit.ci] auto fixes from pre-commit.com hooks
pre-commit-ci[bot] 1170a0e
optimize masking
kainzhong b291b0f
skip masking for some acts
kainzhong 05ea26a
[pre-commit.ci] auto fixes from pre-commit.com hooks
pre-commit-ci[bot] 6e56adf
Merge branch 'main' into cutedsl_mxfp8_common
kainzhong 1318db9
allow user to opt out from building with CuTeDSL
kainzhong 996ad08
better warning
kainzhong 5e39296
doc
kainzhong 1214762
[pre-commit.ci] auto fixes from pre-commit.com hooks
pre-commit-ci[bot] 53f7f4d
add jax test coverage for cutedsl
kainzhong 5e42c35
fix device sm query
kainzhong ae4b4e8
catch more python error
kainzhong ab73b64
[pre-commit.ci] auto fixes from pre-commit.com hooks
pre-commit-ci[bot] f67c7d9
nit
kainzhong 308702c
nit
kainzhong d9636c1
nit
kainzhong 8fa0b71
fix
kainzhong 4797fa9
make tvm-ffi required for python users
kainzhong d02acc2
optimize dbias
kainzhong 3ef8dcd
[pre-commit.ci] auto fixes from pre-commit.com hooks
pre-commit-ci[bot] 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
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,257 @@ | ||
| # 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, driven from JAX. | ||
|
|
||
| JAX companion to tests/pytorch/mxfp8/test_mxfp8_cutedsl_backend.py: the CuTeDSL dispatch | ||
| lives in TE/common, so this checks that the JAX FFI path reaches it and produces the same | ||
| bytes as the CUDA kernels. | ||
| """ | ||
|
|
||
| import ctypes | ||
| import os | ||
|
|
||
| import jax | ||
| import jax.numpy as jnp | ||
| import numpy as np | ||
| import pytest | ||
|
|
||
| from utils import assert_allclose | ||
|
|
||
| from transformer_engine.common import _get_shared_object_file | ||
| from transformer_engine.jax import cpp_extensions as tex | ||
| from transformer_engine.jax.quantize import ( | ||
| QuantizerFactory, | ||
| QuantizeLayout, | ||
| ScaledTensor1x, | ||
| ScalingMode, | ||
| helper, | ||
| ) | ||
|
|
||
| tvm_ffi = pytest.importorskip("tvm_ffi") | ||
|
|
||
| recipe_available, reason_for_no_recipe = helper.is_scaling_mode_supported( | ||
| ScalingMode.MXFP8_1D_SCALING | ||
| ) | ||
|
|
||
| # 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"))) | ||
| # We need this API to manually enable & disable the CuTeDSL backend for the tests | ||
| 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", | ||
| ) | ||
|
|
||
| # CuTeDSL's divisibility assumption strictly requires 32x32 alignment, and the JAX | ||
| # MXFP8 scale shapes additionally require 128-alignment for the fused dact paths. | ||
| MATRIX_SIZES = [ | ||
| (128, 128), | ||
| (256, 1024), | ||
| (512, 512), | ||
| (8192, 7168), | ||
| ] | ||
| # QuantizeLayout.COLWISE is absent on purpose: every quantize wrapper diverts colwise-only to a | ||
| # pure-JAX implementation before reaching the FFI (_quantize_dbias_impl, act_lu, | ||
| # quantize_dact_dbias), so it cannot exercise the kernels. The colwise-only kernels themselves are | ||
| # covered by tests/pytorch/mxfp8/test_mxfp8_cutedsl_backend.py. | ||
| Q_LAYOUTS = [QuantizeLayout.ROWWISE, QuantizeLayout.ROWWISE_COLWISE] | ||
|
|
||
| # Only GeLU activation tests are used (SiLU/ReLU/QGeLU/SReLU commented out | ||
| # in the C++ test as well). Gated variants go to a separate TE/common kernel | ||
| # that the CuTeDSL backend does not cover. | ||
| ACT_TYPE = ("gelu",) | ||
| METHODS = ["CAST_ONLY", "CAST_DBIAS", "CAST_ACT", "CAST_DACT", "CAST_DBIAS_DACT"] | ||
|
|
||
| IN_DTYPES = [jnp.float32, jnp.bfloat16, jnp.float16] | ||
| FP8_DTYPES = [jnp.float8_e4m3fn, jnp.float8_e5m2] | ||
| FP8_TO_KEY = { | ||
| jnp.float8_e4m3fn: "fp8_e4m3fn", | ||
| jnp.float8_e5m2: "fp8_e5m2", | ||
| } | ||
|
|
||
| get_shape_id = lambda s: f"{s[0]}x{s[1]}" | ||
| get_layout_id = lambda l: "rowwise" if l == QuantizeLayout.ROWWISE else "bidim" | ||
| DTYPE_TO_STR = {jnp.float32: "fp32", jnp.bfloat16: "bf16", jnp.float16: "fp16"} | ||
| get_dtype_id = DTYPE_TO_STR.get | ||
| FP8_TO_STR = {jnp.float8_e4m3fn: "e4m3", jnp.float8_e5m2: "e5m2"} | ||
| get_fp8_id = FP8_TO_STR.get | ||
|
|
||
|
|
||
| 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): | ||
| keys = jax.random.split(jax.random.PRNGKey(seed), 4) | ||
|
|
||
| def fill(k_value, k_sign): | ||
| # Mirrors InputsFillCase::uniform in fillCase_special (tests/cpp/test_common.cu) where the | ||
| # uniform range is [-2, 1] and we apply a random sign flip | ||
| v = jax.random.uniform(k_value, (M, N), jnp.float32, -2.0, 1.0) | ||
| negate = jax.random.uniform(k_sign, (M, N), jnp.float32, -1.0, 1.0) < 0.0 | ||
| return jnp.where(negate, -v, v).astype(in_dtype) | ||
|
|
||
| x = fill(keys[0], keys[1]) | ||
| # The activation input is replicated along the -2 axis, one entry per activation | ||
| # in the (possibly gated) activation type. | ||
| act_input = jnp.expand_dims(fill(keys[2], keys[3]), axis=-2) | ||
| return x, act_input | ||
|
|
||
|
|
||
| def run_quantize(method, x, act_input, q_layout, fp8_dtype): | ||
| """Quantize via the public dispatch; returns (scaled_tensor, dbias_or_None).""" | ||
| quantizer = QuantizerFactory.create( | ||
| scaling_mode=ScalingMode.MXFP8_1D_SCALING, q_dtype=fp8_dtype, q_layout=q_layout | ||
| ) | ||
| if method == "CAST_ONLY": | ||
| return tex.quantize(x, quantizer=quantizer), None | ||
| if method == "CAST_DBIAS": | ||
| return tex.quantize_dbias(x, quantizer=quantizer) | ||
| if method == "CAST_ACT": | ||
| return tex.act_lu(act_input, ACT_TYPE, quantizer=quantizer), None | ||
| if method == "CAST_DACT": | ||
| out, _ = tex.quantize_dact_dbias( | ||
| x, act_input, ACT_TYPE, is_dbias=False, quantizer=quantizer | ||
| ) | ||
| return out, None | ||
| if method == "CAST_DBIAS_DACT": | ||
| return tex.quantize_dact_dbias(x, act_input, ACT_TYPE, is_dbias=True, quantizer=quantizer) | ||
| raise ValueError(f"unknown method {method!r}") | ||
|
|
||
|
|
||
| def get_cfg_key(method, in_dtype, fp8_dtype, q_layout): | ||
| """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 implementation 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 = "gelu" | ||
| elif with_dact: | ||
| desc = "dgelu" | ||
| # MXFP8 never asks TE/common for an amax, and JAX quantize emits scales in the linear | ||
| # (non-swizzled) layout -- the GEMM swizzle happens later, in JAX (see gemm.swizzled_scale). | ||
| # trailing False is use_2d_quantization; JAX never requests 2D block scaling | ||
| flags = (True, q_layout.has_colwise, False, False, with_dbias, with_dact, with_act, False) | ||
| return ( | ||
| "cutedsl_mxfp8_" | ||
| + DTYPE_TO_STR[in_dtype] | ||
| + "_" | ||
| + FP8_TO_KEY[fp8_dtype] | ||
| + "_" | ||
| + "_".join("1" if f else "0" for f in flags) | ||
| + "_" | ||
| + desc | ||
| ) | ||
|
|
||
|
|
||
| def extract_quantized_output(out, dbias): | ||
| """Pull the values to compare between backends onto the host. | ||
|
|
||
| Materializing here is what makes the backend toggle safe: JAX dispatch is | ||
| asynchronous, so the FFI handler that reads the toggle may not have run yet when the | ||
| Python call returns. | ||
|
|
||
| ScaledTensor1x carries scale_inv already trimmed to the unpadded shape (see its | ||
| __post_init__), so there is no uninitialized scale padding to exclude. | ||
| """ | ||
| tensors = [out] if isinstance(out, ScaledTensor1x) else [out.rowwise_tensor, out.colwise_tensor] | ||
| parts = {} | ||
| for t in tensors: | ||
| name = "colwise" if t.is_colwise else "rowwise" | ||
| parts[f"{name} data"] = np.asarray(t.data.view(jnp.uint8)) | ||
| parts[f"{name} scales"] = np.asarray(t.scale_inv.view(jnp.uint8)) | ||
| return parts, None if dbias is None else np.asarray(dbias) | ||
|
|
||
|
|
||
| def run_test_case(method, shape, q_layout, in_dtype, fp8_dtype): | ||
| """Assert the CuTeDSL and CUDA backends produce bit-identical outputs for the | ||
| same input and config. | ||
| """ | ||
| M, N = shape | ||
| x, act_input = generate_inputs(M, N, in_dtype) | ||
|
|
||
| set_cutedsl_backend(False) | ||
| cuda_output, dbias_cuda = extract_quantized_output( | ||
| *run_quantize(method, x, act_input, q_layout, fp8_dtype) | ||
| ) | ||
|
|
||
| set_cutedsl_backend(True) | ||
| try: | ||
| cutedsl_output, dbias_cutedsl = extract_quantized_output( | ||
| *run_quantize(method, x, act_input, q_layout, fp8_dtype) | ||
| ) | ||
| 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, in_dtype, fp8_dtype, q_layout) | ||
| 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" | ||
| ) | ||
|
|
||
| tag = ( | ||
| f"{method}/{get_layout_id(q_layout)}/{M}x{N}/" | ||
| f"{DTYPE_TO_STR[in_dtype]}/{FP8_TO_STR[fp8_dtype]}" | ||
| ) | ||
| for name, cuda_bytes in cuda_output.items(): | ||
| assert np.array_equal( | ||
| cutedsl_output[name], cuda_bytes | ||
| ), f"{tag}: {name} differ between backends" | ||
| if dbias_cuda is not None: | ||
| # CuTeDSL kernel does dbias reduction in a slightly different order than the CUDA kernel, | ||
| # due to the non-associativity of floating-point addition, this will not be bit-identical. | ||
| assert_allclose(dbias_cutedsl, dbias_cuda, err_msg=f"{tag}: dbias differs between backends") | ||
|
|
||
|
|
||
| # 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("q_layout", Q_LAYOUTS, ids=get_layout_id) | ||
| @pytest.mark.parametrize("in_dtype", IN_DTYPES, ids=get_dtype_id) | ||
| @pytest.mark.parametrize("fp8_dtype", FP8_DTYPES, ids=get_fp8_id) | ||
| def test_cast_only(fp8_dtype, in_dtype, q_layout, shape): | ||
| run_test_case("CAST_ONLY", shape, q_layout, in_dtype, fp8_dtype) | ||
|
|
||
|
|
||
| # Test cases with varying matrix shapes and quantize layouts | ||
| # (OperatorTest_FusedCastMXFP8_Sizes). | ||
| @pytest.mark.parametrize("shape", MATRIX_SIZES, ids=get_shape_id) | ||
| @pytest.mark.parametrize("q_layout", Q_LAYOUTS, ids=get_layout_id) | ||
| @pytest.mark.parametrize("method", METHODS) | ||
| def test_sizes(method, q_layout, shape): | ||
| run_test_case(method, shape, q_layout, jnp.bfloat16, jnp.float8_e4m3fn) | ||
|
|
||
|
|
||
| # 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", METHODS) | ||
| def test_dtypes(method, fp8_dtype, in_dtype): | ||
| run_test_case(method, (256, 384), QuantizeLayout.ROWWISE_COLWISE, in_dtype, fp8_dtype) |
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.
Uh oh!
There was an error while loading. Please reload this page.