liangsu9988 commited on
Commit
6abc190
·
verified ·
1 Parent(s): daf5525

Promote latest kernel artifacts to main

Browse files
Files changed (40) hide show
  1. CARD.md +55 -0
  2. README.md +113 -6
  3. SYNC.md +49 -0
  4. VALIDATION.md +83 -0
  5. benchmarks/RESULTS.md +57 -0
  6. build.toml +74 -0
  7. build/torch211-cxx11-cu130-aarch64-linux/__init__.py +51 -0
  8. build/torch211-cxx11-cu130-aarch64-linux/{_fp4_gemm_cuda_b46a817.abi3.so → _fp4_gemm_cuda_8a66d8b.abi3.so} +2 -2
  9. build/torch211-cxx11-cu130-aarch64-linux/_ops.py +3 -3
  10. build/torch211-cxx11-cu130-aarch64-linux/metadata.json +5 -9
  11. csrc/cutlass/util/packed_stride.hpp +570 -0
  12. csrc/dequantize_fp4_sfa.cu +89 -0
  13. csrc/dequantize_fp4_sfa.cuh +21 -0
  14. csrc/gemm/fp4/cutlass_fp4_gemm_bias_bf16_sm100.cu +208 -0
  15. csrc/gemm/fp4/cutlass_fp4_gemm_bias_bf16_sm100.cuh +47 -0
  16. csrc/gemm/fp4/cutlass_nvfp4_gemm_bias_gelu_bf16out_sm120.cu +212 -0
  17. csrc/gemm/fp4/cutlass_nvfp4_gemm_bias_gelu_bf16out_sm120.cuh +37 -0
  18. csrc/gemm/fp4/cutlass_nvfp4_gemm_bias_gelu_fp4out_sm120.cu +234 -0
  19. csrc/gemm/fp4/cutlass_nvfp4_gemm_bias_gelu_fp4out_sm120.cuh +40 -0
  20. csrc/gemm/fp4/cutlass_nvfp4_gemm_dn_streamk_bias_sm120.cu +307 -0
  21. csrc/gemm/fp4/cutlass_nvfp4_gemm_dn_streamk_bias_sm120.cuh +53 -0
  22. csrc/gemm/fp4/cutlass_nvfp4_w4a16_gemm_sm100.cu +411 -0
  23. csrc/gemm/fp4/cutlass_nvfp4_w4a16_gemm_sm100.cuh +65 -0
  24. csrc/gemm/fp4/cutlass_nvfp4_w4a16_gemm_sm120.cu +690 -0
  25. csrc/gemm/fp4/cutlass_nvfp4_w4a16_gemm_sm120.cuh +112 -0
  26. csrc/gemm/fp4/fp4_w4a4_mma_warpsplit_sm120.cu +147 -0
  27. csrc/gemm/fp4/fp4_w4a4_mma_warpsplit_sm120.cuh +20 -0
  28. csrc/gemm/fp4/sm110_dispatch.cu +50 -0
  29. csrc/gemm/fp4/sm110_dispatch.cuh +45 -0
  30. csrc/quantize/quantize_fp4_sfa.cu +194 -0
  31. csrc/quantize/quantize_fp4_sfa.cuh +41 -0
  32. csrc/quantize/quantize_fp4_sfa_bf16.cu +142 -0
  33. csrc/quantize/quantize_fp4_sfa_bf16.cuh +24 -0
  34. examples/README.md +10 -0
  35. examples/fp4_gemm_linear.py +29 -0
  36. flake.nix +18 -0
  37. tests/test_fp4_gemm.py +668 -0
  38. torch-ext/fp4_gemm/__init__.py +358 -0
  39. torch-ext/torch_binding.cpp +606 -0
  40. torch-ext/torch_binding.h +31 -0
CARD.md ADDED
@@ -0,0 +1,55 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # flashrt/fp4-gemm
2
+
3
+ FlashRT native Blackwell NVFP4 A4W4 GEMM kernels. Both activations and weights
4
+ are packed FP4 inputs; this is not a BF16-activation weight-only operation.
5
+
6
+ ## Functions
7
+
8
+ - `sfa_size_bytes`
9
+ - `quantize_fp4_sfa_fp16`
10
+ - `quantize_fp4_sfa_bf16`
11
+ - `dequantize_fp4_sfa_fp16`
12
+ - `nvfp4_gemm_bf16`
13
+ - `nvfp4_gemm_bias_bf16`
14
+ - `nvfp4_gemm_bias_residual_bf16`
15
+ - `nvfp4_gemm_residual_bf16`
16
+ - `nvfp4_gemm_bias_gelu_bf16`
17
+ - `nvfp4_gemm_bias_gelu_nvfp4`
18
+ - `nvfp4_gemm_streamk_bf16`
19
+ - `nvfp4_gemm_streamk_bias_bf16`
20
+ - `fp4_w4a16_linear_bf16` (compatibility alias)
21
+
22
+ ## Example
23
+
24
+ ```python
25
+ from kernels import get_kernel
26
+ import torch
27
+
28
+ ops = get_kernel("flashrt/fp4-gemm", version=1, trust_remote_code=True)
29
+
30
+ x = torch.randn((32, 256), device="cuda", dtype=torch.float16)
31
+ w = torch.randn((512, 256), device="cuda", dtype=torch.float16)
32
+
33
+ a, sfa = ops.quantize_fp4_sfa_fp16(x, is_sfb=False)
34
+ b, sfb = ops.quantize_fp4_sfa_fp16(w, is_sfb=True)
35
+ y = ops.nvfp4_gemm_bf16(a, b, sfa, sfb)
36
+ ```
37
+
38
+ BF16 activations should use the direct producer to avoid a separate cast and
39
+ copy before every low-bit projection:
40
+
41
+ ```python
42
+ x = torch.randn((1, 5120), device="cuda", dtype=torch.bfloat16)
43
+ a, sfa = ops.quantize_fp4_sfa_bf16(x)
44
+ ```
45
+
46
+ ## Notes
47
+
48
+ - Blackwell `sm_110a` with CUDA 13+ and `sm_120a` with CUDA 12.8+.
49
+ - Inputs are packed FP4 E2M1 plus CUTLASS Sm1xx SFA/SFB scale buffers.
50
+ - Output is BF16.
51
+ - `variant=-1` is the architecture-aware production auto-dispatch;
52
+ `variant=0/1/2` expose diagnostic default, widen, and pingpong schedules.
53
+ - The canonical BF16-output GEMM and FP4 pack/unpack helpers support SM110 and
54
+ SM120. SM110 also supports bias, bias+residual, and bias+GELU-to-FP4
55
+ production epilogues used by the GROOT N1.7 Thor pipeline.
README.md CHANGED
@@ -1,9 +1,116 @@
1
- # flashrt/fp4-gemm
2
 
3
- This repository is a compatibility mirror for older `kernels` clients
4
- that resolve repositories through the default Hugging Face model repo API.
5
 
6
- Canonical Kernel Hub repo: https://huggingface.co/kernels/flashrt/fp4-gemm
 
 
 
7
 
8
- Do not edit this mirror by hand. It is generated from the Kernel Hub
9
- `vN` branches and contains the same `build/**` artifacts.
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # fp4-gemm
2
 
3
+ FlashRT native Blackwell NVFP4 A4W4 GEMM kernels.
 
4
 
5
+ This package consumes packed FP4 E2M1 tensors plus CUTLASS Sm1xx SFA/SFB scale
6
+ buffers and produces BF16 output. It is designed to pair with
7
+ `flashrt/fp4-fused-ops` and other static low-bit transformer/diffuser runtime
8
+ paths.
9
 
10
+ ## Available Functions
11
+
12
+ - `sfa_size_bytes(rows, dim)`
13
+ - `quantize_fp4_sfa_fp16(x, packed=None, sfa=None, is_sfb=False)`
14
+ - `quantize_fp4_sfa_bf16(x, packed=None, sfa=None, is_sfb=False)`
15
+ - `dequantize_fp4_sfa_fp16(packed, sfa, out=None, is_sfb=False)`
16
+ - `nvfp4_gemm_bf16(a_packed, b_packed, sfa, sfb, alpha=1.0, out=None, variant=-1)`
17
+ - `nvfp4_gemm_bias_bf16(a_packed, b_packed, sfa, sfb, bias, out=None)`
18
+ - `nvfp4_gemm_bias_residual_bf16(a_packed, b_packed, sfa, sfb, bias, residual, out=None)`
19
+ - `nvfp4_gemm_residual_bf16(a_packed, b_packed, sfa, sfb, residual, alpha=1.0, out=None)`
20
+ - `nvfp4_gemm_bias_gelu_bf16(a_packed, b_packed, sfa, sfb, bias, alpha=1.0, out=None)`
21
+ - `nvfp4_gemm_bias_gelu_nvfp4(a_packed, b_packed, sfa, sfb, bias, alpha=1.0, out_packed=None, out_sfa=None)`
22
+ - `nvfp4_gemm_streamk_bf16(a_packed, b_packed, sfa, sfb, alpha=1.0, out=None)`
23
+ - `nvfp4_gemm_streamk_bias_bf16(a_packed, b_packed, sfa, sfb, bias, alpha=1.0, out=None)`
24
+ - `fp4_w4a16_linear_bf16(...)` is retained as a compatibility alias
25
+
26
+ ## Tensor Contract
27
+
28
+ - `a_packed`: `torch.uint8`, shape `(M, K / 2)`.
29
+ - `b_packed`: `torch.uint8`, shape `(N, K / 2)`.
30
+ - `sfa`: `torch.uint8`, CUTLASS SFA layout for `(M, K)`.
31
+ - `sfb`: `torch.uint8`, CUTLASS SFB layout for `(N, K)`.
32
+ - output: `torch.bfloat16`, shape `(M, N)`.
33
+ - `K` must be divisible by 16.
34
+ - Targets: Blackwell `sm_110a` (Jetson AGX Thor, CUDA 13+) and `sm_120a`
35
+ (RTX Blackwell, CUDA 12.8+).
36
+
37
+ `variant` selects the CUTLASS schedule:
38
+
39
+ - `-1`: architecture-aware auto-dispatch (public default).
40
+ - `0`: default `<128,128,256>` cooperative schedule.
41
+ - `1`: widen `<128,256,128>` schedule, intended for very large `N`.
42
+ - `2`: pingpong schedule for A/B testing shape-specific wins.
43
+
44
+ The canonical linear API and FP4/SFA quantize/dequantize helpers are available
45
+ on both SM110 and SM120. SM110 additionally provides the GROOT N1.7 production
46
+ epilogues `nvfp4_gemm_bias_bf16`, `nvfp4_gemm_bias_residual_bf16`, and
47
+ `nvfp4_gemm_bias_gelu_nvfp4`. The latter emits packed FP4 plus CUTLASS SFA so
48
+ the following projection can consume it without a BF16 materialization and a
49
+ standalone quantization launch. Stream-K and the older BF16 GELU epilogue keep
50
+ their existing SM120 dispatch and reject unsupported architectures explicitly.
51
+
52
+ The SM110 release gate includes the production `(M,N,K)` shapes
53
+ `(41,4608,1536)`, `(41,6144,1536)`, and `(41,1536,6144)`, plus the legacy
54
+ `M=51` compatibility row. The kernels are the native sources used by FlashRT's
55
+ GROOT N1.7 Thor NVFP4 pipeline.
56
+
57
+ ## Minimal Usage
58
+
59
+ ```python
60
+ from kernels import get_kernel
61
+ import torch
62
+
63
+ ops = get_kernel("flashrt/fp4-gemm", version=1, trust_remote_code=True)
64
+
65
+ x = torch.randn((32, 256), device="cuda", dtype=torch.float16)
66
+ w = torch.randn((512, 256), device="cuda", dtype=torch.float16)
67
+
68
+ a_packed, sfa = ops.quantize_fp4_sfa_fp16(x, is_sfb=False)
69
+ b_packed, sfb = ops.quantize_fp4_sfa_fp16(w, is_sfb=True)
70
+
71
+ y = ops.nvfp4_gemm_bf16(a_packed, b_packed, sfa, sfb, alpha=1.0)
72
+ ```
73
+
74
+ For BF16 model activations, use the direct producer so the hot path does not
75
+ materialize an intermediate FP16 tensor:
76
+
77
+ ```python
78
+ x_bf16 = torch.randn((1, 5120), device="cuda", dtype=torch.bfloat16)
79
+ a_packed, sfa = ops.quantize_fp4_sfa_bf16(x_bf16)
80
+ ```
81
+
82
+ The BF16 entry writes the same E2M1 bytes and CUTLASS SFA/SFB layout as
83
+ `quantize_fp4_sfa_fp16(x_bf16.to(torch.float16))` for finite FP16-range
84
+ inputs. It is an additive API; the existing FP16 producer remains unchanged.
85
+
86
+ The quantize/dequantize helpers are included for examples and validation. A
87
+ production runtime should keep weights prepacked and should avoid quantizing in
88
+ the hot path unless that producer kernel is part of the intended low-bit block.
89
+
90
+ Use the bias/GELU and residual variants to avoid returning to BF16
91
+ elementwise code between low-bit GEMMs. Stream-K variants are selected only
92
+ for the validated large down-projection shapes; unsupported shapes reject
93
+ rather than silently selecting a losing schedule.
94
+
95
+ ## Validation
96
+
97
+ ```bash
98
+ python fp4-gemm/tests/test_fp4_gemm.py --backend source --mode full
99
+ python fp4-gemm/tests/test_fp4_gemm.py --backend installed --mode full \
100
+ --artifact fp4-gemm/build/torch211-cxx11-cu128-x86_64-linux
101
+ python fp4-gemm/benchmarks/benchmark.py --backend installed --mode headline \
102
+ --artifact fp4-gemm/build/torch211-cxx11-cu128-x86_64-linux
103
+
104
+ # Thor model-shape gate
105
+ python fp4-gemm/tests/test_fp4_gemm.py --backend installed \
106
+ --mode thor-models \
107
+ --artifact fp4-gemm/build/torch211-cxx11-cu130-aarch64-linux
108
+ ```
109
+
110
+ The correctness reference dequantizes the same FP4/SFA and FP4/SFB inputs used
111
+ by the kernel, then computes the PyTorch GEMM reference from those dequantized
112
+ low-bit values.
113
+
114
+ The producer gate also checks the BF16 direct entry byte-for-byte against the
115
+ established FP16 compatibility chain at decode widths 5120, 6144 and 17408,
116
+ plus multi-row activation and SFB layouts.
SYNC.md ADDED
@@ -0,0 +1,49 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Source Sync
2
+
3
+ - Upstream FlashRT source: `../official/FlashRT`
4
+ - Original SM110 sync commit: `132049d7c3a3534fb7d35676cd726f39408b1af6`
5
+ - GROOT N1.7 fused-epilogue sync commit:
6
+ `24df793f4fa2d50780aea03b644208c6e0cb4162`
7
+ - Initial package date: June 20, 2026
8
+
9
+ Copied source files:
10
+
11
+ - `csrc/gemm/fp4/cutlass_nvfp4_w4a16_gemm_sm120.cu/.cuh`
12
+ - `csrc/gemm/fp4/cutlass_nvfp4_w4a16_gemm_sm100.cu/.cuh`
13
+ - `csrc/gemm/fp4/cutlass_fp4_gemm_bias_bf16_sm100.cu/.cuh`
14
+ - `csrc/quantize/quantize_fp4_sfa.cu/.cuh`
15
+ - `csrc/quantize/quantize_fp4_sfa_bf16.cu/.cuh`
16
+ - `cutlass/util/packed_stride.hpp`, copied from CUTLASS tools util headers
17
+ into `csrc/cutlass/util/packed_stride.hpp` so the Hub package does not
18
+ depend on a local `third_party/cutlass/tools/util/include` path.
19
+
20
+ Packaging helper:
21
+
22
+ - `csrc/dequantize_fp4_sfa.cu/.cuh` derived from the SFA dequant validation
23
+ helper used in `fp4-fused-ops`; this package adds `is_sfb` support so tests
24
+ can dequant both A/SFA and B/SFB.
25
+
26
+ Local packaging edits:
27
+
28
+ - Added Tensor-facing PyTorch custom ops in `torch-ext/torch_binding.cpp`.
29
+ - Added Python wrappers and fake registrations in `torch-ext/fp4_gemm`.
30
+ - Added the BF16 direct SFA/SFB producer as an input-type specialization of
31
+ the existing FP16 producer. Its E2M1 encoding and CUTLASS scale layout are
32
+ unchanged; the additive entry removes a standalone activation cast.
33
+ - Public APIs accept CUDA tensors only; no raw pointers or stream arguments.
34
+ - CUTLASS SM100/SM120 block-scaled layout support is treated as package scope,
35
+ not as a test-only compiler define.
36
+ - The Tensor binding dispatches the canonical BF16-output GEMM by runtime
37
+ compute capability. CUDA 12.8 artifacts link only the SM120 implementation;
38
+ CUDA 13 artifacts link both SM110 and SM120 implementations.
39
+ - SM110 production auto-dispatch was tiled against PI0.5, GROOT, Cosmos Edge,
40
+ and LingBot VLA projection shapes. Explicit schedule IDs remain diagnostic.
41
+
42
+ Architecture limits:
43
+
44
+ - The canonical BF16-output GEMM and SFA/SFB helpers support SM110 and SM120.
45
+ - Bias, residual, and bias/GELU-to-FP4 epilogues have an independent SM110
46
+ backend copied from the production GROOT N1.7 path. Stream-K remains an
47
+ SM120-only API and rejects on SM110.
48
+ - SM110 requires CUDA 13 and the package's pinned CUTLASS 4.4 target; SM120
49
+ requires CUDA 12.8 and CUTLASS 4.0.
VALIDATION.md ADDED
@@ -0,0 +1,83 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Validation
2
+
3
+ Local source validation covers NVIDIA GeForce RTX 5090 (SM120) and NVIDIA
4
+ Jetson AGX Thor (SM110).
5
+
6
+ ```bash
7
+ python fp4-gemm/tests/test_fp4_gemm.py \
8
+ --backend source \
9
+ --mode full \
10
+ --json-out internal-tests/fp4-gemm-source-full.json
11
+ ```
12
+
13
+ Result:
14
+
15
+ - SM120 full gate: `25/25` checks passed, including all fused epilogues and
16
+ the aggregate BF16 direct-producer layout gate.
17
+ - SM110 model-shape gate: `24/24` checks passed across PI0.5, GROOT, Cosmos
18
+ Edge, and LingBot VLA projection shapes.
19
+ - Variants `0`, `1`, and `2` were checked.
20
+ - SM110 additionally checks production auto-dispatch (`variant=-1`).
21
+ - `nvfp4_gemm_bf16` is the canonical public API.
22
+ - Correctness reference dequantizes the same FP4/SFA and FP4/SFB inputs used
23
+ by the kernel, then computes PyTorch GEMM on those dequantized values.
24
+ - The direct BF16 producer is byte-exact against the established
25
+ BF16-to-FP16 plus FP16-producer contract for packed E2M1, mapped SFA/SFB
26
+ bytes, and dequantized output. Covered activation shapes are `(1,5120)`,
27
+ `(1,6144)`, `(1,17408)`, `(16,2048)`, and `(128,512)`; SFB coverage uses
28
+ `(64,1024)`.
29
+
30
+ | Shape | Variant | Max abs | Mean abs | P99 abs | Cosine |
31
+ | --- | ---: | ---: | ---: | ---: | ---: |
32
+ | M=16, N=128, K=128 | 0 | 0.0 | 0.0 | 0.0 | 1.0 |
33
+ | M=16, N=128, K=128 | 1 | 0.0 | 0.0 | 0.0 | 1.0 |
34
+ | M=16, N=128, K=128 | 2 | 0.0 | 0.0 | 0.0 | 1.0 |
35
+ | M=32, N=256, K=256 | 0 | 0.0 | 0.0 | 0.0 | 1.0 |
36
+ | M=32, N=256, K=256 | 1 | 0.0 | 0.0 | 0.0 | 1.0 |
37
+ | M=32, N=256, K=256 | 2 | 0.0 | 0.0 | 0.0 | 1.0 |
38
+ | M=64, N=512, K=512 | 0 | 0.0 | 0.0 | 0.0 | 1.0 |
39
+ | M=64, N=512, K=512 | 1 | 0.0 | 0.0 | 0.0 | 1.0 |
40
+ | M=64, N=512, K=512 | 2 | 0.0 | 0.0 | 0.0 | 1.0 |
41
+
42
+ ## Installed Artifact Validation
43
+
44
+ The local kernel-builder release candidate produced and passed ABI, manylinux,
45
+ layout, and builder `get_kernel` checks for:
46
+
47
+ - `torch211-cxx11-cu128-x86_64-linux`
48
+ - `torch211-cxx11-cu130-x86_64-linux`
49
+ - `torch212-cxx11-cu130-x86_64-linux`
50
+ - `torch212-cxx11-cu132-x86_64-linux`
51
+
52
+ The cu128/Torch 2.11 artifact passed `10/10` runtime gates: all nine
53
+ shape/variant correctness rows were exact against the staged reference, and
54
+ the public `nvfp4_gemm_bf16` wrapper was exact under
55
+ `torch.compile(fullgraph=True)`.
56
+
57
+ The SM110 release flake pins kernel-builder commit
58
+ `d720fa90fb9cd92d1bc60a9dc5c55bef2aafabb8`, which includes CUTLASS 4.5
59
+ support and the corrected CUTLASS 4.5.2 fixed-output hash. HF Jobs, the
60
+ SM110 aarch64 artifact build, and cold Hub loads must pass before the rebuilt
61
+ Hub release is considered complete.
62
+
63
+ ## BF16 Direct Producer
64
+
65
+ RTX 5090, 100 warmup iterations and 1000 measured iterations:
66
+
67
+ | Shape | BF16 direct | BF16 cast + FP16 producer | Speedup | Native BF16 producer | Hub/native |
68
+ | --- | ---: | ---: | ---: | ---: | ---: |
69
+ | M=1, K=5120 | 4.098 us | 6.404 us | 1.563x | 6.150 us | 0.666x |
70
+ | M=1, K=6144 | 4.098 us | 6.403 us | 1.562x | 8.190 us | 0.500x |
71
+ | M=1, K=17408 | 4.096 us | 6.413 us | 1.566x | 18.442 us | 0.222x |
72
+
73
+ The native BF16 producer is included as a latency comparison but uses a
74
+ different FlashRT quantization strategy. Correctness acceptance is therefore
75
+ against this package's established FP16 producer contract, where all tested
76
+ packed and mapped scale bytes are exact.
77
+
78
+ ## Thor Native Parity
79
+
80
+ The Tensor wrapper was compared against the same native FlashRT launchers on
81
+ Thor with 20 warmup and 100 measured iterations. For production auto-dispatch
82
+ across the six model shapes, wrapper/native latency ratio had median `1.019`
83
+ and maximum `1.086`. Correctness was exact (`max_abs=mean_abs=p99_abs=0`).
benchmarks/RESULTS.md ADDED
@@ -0,0 +1,57 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # fp4-gemm Benchmark Results
2
+
3
+ Installed kernel-builder artifact benchmark on NVIDIA GeForce RTX 5090,
4
+ PyTorch `2.11.0+cu128`.
5
+
6
+ Command:
7
+
8
+ ```bash
9
+ python fp4-gemm/benchmarks/benchmark.py \
10
+ --backend installed \
11
+ --artifact fp4-gemm/build/torch211-cxx11-cu128-x86_64-linux \
12
+ --mode headline \
13
+ --warmup 100 \
14
+ --iterations 500 \
15
+ --json-out internal-tests/fp4-gemm-installed-benchmark.json
16
+ ```
17
+
18
+ Reference is PyTorch GEMM over the same dequantized FP4/SFA and FP4/SFB inputs
19
+ that the FlashRT kernel consumes.
20
+
21
+ | Shape | Variant | FlashRT us | Eager us | Compile us | vs eager | vs compile | Max abs |
22
+ | --- | ---: | ---: | ---: | ---: | ---: | ---: | ---: |
23
+ | M=16, N=128, K=128 | 0 | 6.156 | 15.156 | 27.748 | 2.46x | 4.51x | 0.0 |
24
+ | M=16, N=128, K=128 | 1 | 6.152 | 15.156 | 27.748 | 2.46x | 4.51x | 0.0 |
25
+ | M=16, N=128, K=128 | 2 | 6.145 | 15.156 | 27.748 | 2.47x | 4.52x | 0.0 |
26
+ | M=32, N=256, K=256 | 0 | 6.153 | 16.685 | 35.690 | 2.71x | 5.80x | 0.0 |
27
+ | M=32, N=256, K=256 | 1 | 8.201 | 16.685 | 35.690 | 2.03x | 4.35x | 0.0 |
28
+ | M=32, N=256, K=256 | 2 | 6.147 | 16.685 | 35.690 | 2.71x | 5.81x | 0.0 |
29
+ | M=64, N=512, K=512 | 0 | 6.152 | 16.480 | 36.205 | 2.68x | 5.89x | 0.0 |
30
+ | M=64, N=512, K=512 | 1 | 10.246 | 16.480 | 36.205 | 1.61x | 3.53x | 0.0 |
31
+ | M=64, N=512, K=512 | 2 | 6.152 | 16.480 | 36.205 | 2.68x | 5.89x | 0.0 |
32
+
33
+ Variant notes:
34
+
35
+ - `variant=0` is the stable default.
36
+ - `variant=1` is the widen schedule intended for very large `N`; it is not the
37
+ best choice for these small validation shapes.
38
+ - `variant=2` is competitive on small shapes and remains exposed for explicit
39
+ A/B testing.
40
+
41
+ The PyTorch references consume the same already-dequantized FP4 tensors and do
42
+ not include quantization. The compiled reference is warmed before timing.
43
+
44
+ ## BF16 Direct Producer
45
+
46
+ Source benchmark on RTX 5090 with 100 warmup and 1000 measured iterations:
47
+
48
+ | Shape | Direct BF16 us | Cast + FP16 producer us | Speedup | Native BF16 us | Wrapper/native |
49
+ | --- | ---: | ---: | ---: | ---: | ---: |
50
+ | M=1, K=5120 | 4.098 | 6.404 | 1.563x | 6.150 | 0.666x |
51
+ | M=1, K=6144 | 4.098 | 6.403 | 1.562x | 8.190 | 0.500x |
52
+ | M=1, K=17408 | 4.096 | 6.413 | 1.566x | 18.442 | 0.222x |
53
+
54
+ The direct entry is byte-exact against the package's established
55
+ BF16-to-FP16 plus FP16-producer contract. The native timing is reported as a
56
+ performance reference only because that producer uses a distinct quantization
57
+ strategy.
build.toml ADDED
@@ -0,0 +1,74 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ [general]
2
+ name = "fp4-gemm"
3
+ license = "Apache-2.0"
4
+ version = 1
5
+ backends = ["cuda"]
6
+
7
+ [general.cuda]
8
+ minver = "12.8"
9
+
10
+ [general.hub]
11
+ repo-id = "flashrt/fp4-gemm"
12
+
13
+ [torch]
14
+ include = ["csrc"]
15
+ src = [
16
+ "torch-ext/torch_binding.cpp",
17
+ "torch-ext/torch_binding.h",
18
+ ]
19
+
20
+ [kernel.fp4_gemm_sm110]
21
+ backend = "cuda"
22
+ depends = ["torch", "cutlass_4_4"]
23
+ include = ["csrc"]
24
+ cuda-minver = "13"
25
+ cuda-capabilities = ["11.0a"]
26
+ cuda-flags = [
27
+ "--expt-relaxed-constexpr",
28
+ "-O3",
29
+ "--use_fast_math",
30
+ "-DCUTLASS_ARCH_MMA_SM100_SUPPORTED=1",
31
+ ]
32
+ src = [
33
+ "csrc/gemm/fp4/cutlass_nvfp4_w4a16_gemm_sm100.cu",
34
+ "csrc/gemm/fp4/cutlass_nvfp4_w4a16_gemm_sm100.cuh",
35
+ "csrc/gemm/fp4/cutlass_fp4_gemm_bias_bf16_sm100.cu",
36
+ "csrc/gemm/fp4/cutlass_fp4_gemm_bias_bf16_sm100.cuh",
37
+ "csrc/quantize/quantize_fp4_sfa_bf16.cu",
38
+ "csrc/quantize/quantize_fp4_sfa_bf16.cuh",
39
+ "csrc/gemm/fp4/sm110_dispatch.cu",
40
+ "csrc/gemm/fp4/sm110_dispatch.cuh",
41
+ ]
42
+
43
+ [kernel.fp4_gemm_common]
44
+ backend = "cuda"
45
+ depends = ["torch"]
46
+ include = ["csrc"]
47
+ cuda-minver = "12.8"
48
+ cuda-capabilities = ["11.0a", "12.0a"]
49
+ cuda-flags = ["--expt-relaxed-constexpr", "-O3", "--use_fast_math"]
50
+ src = [
51
+ "csrc/quantize/quantize_fp4_sfa.cu",
52
+ "csrc/quantize/quantize_fp4_sfa.cuh",
53
+ "csrc/dequantize_fp4_sfa.cu",
54
+ "csrc/dequantize_fp4_sfa.cuh",
55
+ ]
56
+
57
+ [kernel.fp4_gemm_sm120]
58
+ backend = "cuda"
59
+ depends = ["torch", "cutlass_4_0"]
60
+ include = ["csrc"]
61
+ cuda-minver = "12.8"
62
+ cuda-capabilities = ["12.0a"]
63
+ src = [
64
+ "csrc/gemm/fp4/cutlass_nvfp4_w4a16_gemm_sm120.cu",
65
+ "csrc/gemm/fp4/fp4_w4a4_mma_warpsplit_sm120.cu",
66
+ "csrc/gemm/fp4/fp4_w4a4_mma_warpsplit_sm120.cuh",
67
+ "csrc/gemm/fp4/cutlass_nvfp4_w4a16_gemm_sm120.cuh",
68
+ "csrc/gemm/fp4/cutlass_nvfp4_gemm_bias_gelu_bf16out_sm120.cu",
69
+ "csrc/gemm/fp4/cutlass_nvfp4_gemm_bias_gelu_bf16out_sm120.cuh",
70
+ "csrc/gemm/fp4/cutlass_nvfp4_gemm_bias_gelu_fp4out_sm120.cu",
71
+ "csrc/gemm/fp4/cutlass_nvfp4_gemm_bias_gelu_fp4out_sm120.cuh",
72
+ "csrc/gemm/fp4/cutlass_nvfp4_gemm_dn_streamk_bias_sm120.cu",
73
+ "csrc/gemm/fp4/cutlass_nvfp4_gemm_dn_streamk_bias_sm120.cuh",
74
+ ]
build/torch211-cxx11-cu130-aarch64-linux/__init__.py CHANGED
@@ -58,6 +58,16 @@ def _legacy_linear_fake(
58
  return None
59
 
60
 
 
 
 
 
 
 
 
 
 
 
61
  @torch.library.register_fake(add_op_namespace_prefix("quantize_fp4_sfa_fp16"))
62
  def _quant_fake(x: torch.Tensor, packed: torch.Tensor, sfa: torch.Tensor, is_sfb: bool = False) -> None:
63
  return None
@@ -191,6 +201,45 @@ def fp4_w4a16_linear_bf16(
191
  )
192
 
193
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
194
  def nvfp4_gemm_residual_bf16(
195
  a_packed: torch.Tensor,
196
  b_packed: torch.Tensor,
@@ -296,8 +345,10 @@ __all__ = [
296
  "fp4_w4a16_linear_bf16",
297
  "fp4_w4a4_gemv_warpsplit_bf16",
298
  "nvfp4_gemm_bf16",
 
299
  "nvfp4_gemm_bias_gelu_bf16",
300
  "nvfp4_gemm_bias_gelu_nvfp4",
 
301
  "nvfp4_gemm_residual_bf16",
302
  "nvfp4_gemm_streamk_bf16",
303
  "nvfp4_gemm_streamk_bias_bf16",
 
58
  return None
59
 
60
 
61
+ @torch.library.register_fake(add_op_namespace_prefix("nvfp4_gemm_bias_bf16"))
62
+ def _bias_fake(a, b, sfa, sfb, bias, out) -> None:
63
+ return None
64
+
65
+
66
+ @torch.library.register_fake(add_op_namespace_prefix("nvfp4_gemm_bias_residual_bf16"))
67
+ def _bias_residual_fake(a, b, sfa, sfb, bias, residual, out) -> None:
68
+ return None
69
+
70
+
71
  @torch.library.register_fake(add_op_namespace_prefix("quantize_fp4_sfa_fp16"))
72
  def _quant_fake(x: torch.Tensor, packed: torch.Tensor, sfa: torch.Tensor, is_sfb: bool = False) -> None:
73
  return None
 
201
  )
202
 
203
 
204
+ def nvfp4_gemm_bias_bf16(
205
+ a_packed: torch.Tensor,
206
+ b_packed: torch.Tensor,
207
+ sfa: torch.Tensor,
208
+ sfb: torch.Tensor,
209
+ bias: torch.Tensor,
210
+ *,
211
+ out: torch.Tensor | None = None,
212
+ ) -> torch.Tensor:
213
+ """SM110 NVFP4 GEMM with a fused per-column BF16 bias."""
214
+ if out is None:
215
+ out = torch.empty(
216
+ (a_packed.shape[0], b_packed.shape[0]),
217
+ device=a_packed.device,
218
+ dtype=torch.bfloat16,
219
+ )
220
+ ops.nvfp4_gemm_bias_bf16(a_packed, b_packed, sfa, sfb, bias, out)
221
+ return out
222
+
223
+
224
+ def nvfp4_gemm_bias_residual_bf16(
225
+ a_packed: torch.Tensor,
226
+ b_packed: torch.Tensor,
227
+ sfa: torch.Tensor,
228
+ sfb: torch.Tensor,
229
+ bias: torch.Tensor,
230
+ residual: torch.Tensor,
231
+ *,
232
+ out: torch.Tensor | None = None,
233
+ ) -> torch.Tensor:
234
+ """SM110 NVFP4 GEMM with fused BF16 bias and residual add."""
235
+ if out is None:
236
+ out = torch.empty_like(residual)
237
+ ops.nvfp4_gemm_bias_residual_bf16(
238
+ a_packed, b_packed, sfa, sfb, bias, residual, out
239
+ )
240
+ return out
241
+
242
+
243
  def nvfp4_gemm_residual_bf16(
244
  a_packed: torch.Tensor,
245
  b_packed: torch.Tensor,
 
345
  "fp4_w4a16_linear_bf16",
346
  "fp4_w4a4_gemv_warpsplit_bf16",
347
  "nvfp4_gemm_bf16",
348
+ "nvfp4_gemm_bias_bf16",
349
  "nvfp4_gemm_bias_gelu_bf16",
350
  "nvfp4_gemm_bias_gelu_nvfp4",
351
+ "nvfp4_gemm_bias_residual_bf16",
352
  "nvfp4_gemm_residual_bf16",
353
  "nvfp4_gemm_streamk_bf16",
354
  "nvfp4_gemm_streamk_bias_bf16",
build/torch211-cxx11-cu130-aarch64-linux/{_fp4_gemm_cuda_b46a817.abi3.so → _fp4_gemm_cuda_8a66d8b.abi3.so} RENAMED
@@ -1,3 +1,3 @@
1
  version https://git-lfs.github.com/spec/v1
2
- oid sha256:49cad7b4b7a236d14238c99a832bdc106c5e8494c8ce89e000c3bc75b45cc061
3
- size 751384
 
1
  version https://git-lfs.github.com/spec/v1
2
+ oid sha256:35bd2043a1a1880f362954fe72a22651017d30a585d38ed5d2ffe86de286b45a
3
+ size 1118808
build/torch211-cxx11-cu130-aarch64-linux/_ops.py CHANGED
@@ -1,6 +1,6 @@
1
  import torch
2
- from . import _fp4_gemm_cuda_b46a817
3
- ops = torch.ops._fp4_gemm_cuda_b46a817
4
 
5
  def add_op_namespace_prefix(op_name: str):
6
- return f"_fp4_gemm_cuda_b46a817::{op_name}"
 
1
  import torch
2
+ from . import _fp4_gemm_cuda_8a66d8b
3
+ ops = torch.ops._fp4_gemm_cuda_8a66d8b
4
 
5
  def add_op_namespace_prefix(op_name: str):
6
+ return f"_fp4_gemm_cuda_8a66d8b::{op_name}"
build/torch211-cxx11-cu130-aarch64-linux/metadata.json CHANGED
@@ -1,6 +1,6 @@
1
  {
2
  "name": "fp4-gemm",
3
- "id": "_fp4_gemm_cuda_b46a817",
4
  "version": 1,
5
  "license": "Apache-2.0",
6
  "python-depends": [],
@@ -13,20 +13,16 @@
13
  "digest": {
14
  "algorithm": "sha256",
15
  "files": {
16
- "__init__.py": "pJlGLDlOKnju+m59WyZJjpynMTvkPQTF6w+z4YfCwSY=",
17
- "_fp4_gemm_cuda_b46a817.abi3.so": "ScrXtLeiNtFCOMmagyvcEGxehJTIzongAMO8dbRcwGE=",
18
- "_ops.py": "PqcxyYzP4Co0y9NUZPOomwASQWuPPNn5QjtKrfTiZ2k=",
19
  "fp4_gemm/__init__.py": "v6p5XMfQzddhi1fLSAw4HX9CyS0rQsidvu9VsT01xi4="
20
  }
21
  },
22
  "provenance": {
23
  "kernel": {
24
- "sha": "b46a81771fdabe26f2c2494dc2ecf8492c88c45a",
25
  "dirty": false
26
- },
27
- "validation": {
28
- "torch": "2.11.0+cu130",
29
- "cuda": "13.0"
30
  }
31
  }
32
  }
 
1
  {
2
  "name": "fp4-gemm",
3
+ "id": "_fp4_gemm_cuda_8a66d8b",
4
  "version": 1,
5
  "license": "Apache-2.0",
6
  "python-depends": [],
 
13
  "digest": {
14
  "algorithm": "sha256",
15
  "files": {
16
+ "__init__.py": "+Kk/VNnIWwe9nszQYvL2qtU3DpKxbqHVfklg5WtVuB4=",
17
+ "_fp4_gemm_cuda_8a66d8b.abi3.so": "Nb0gQ6GhiA82KVT+cqImUQF9MKWF047V0v/obeKGtFo=",
18
+ "_ops.py": "/hna1MwmSjR4pewWf12X5smsDRcaiHuoR2Gri/PrBFc=",
19
  "fp4_gemm/__init__.py": "v6p5XMfQzddhi1fLSAw4HX9CyS0rQsidvu9VsT01xi4="
20
  }
21
  },
22
  "provenance": {
23
  "kernel": {
24
+ "sha": "8a66d8b",
25
  "dirty": false
 
 
 
 
26
  }
27
  }
28
  }
csrc/cutlass/util/packed_stride.hpp ADDED
@@ -0,0 +1,570 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ /***************************************************************************************************
2
+ * Copyright (c) 2023 - 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
3
+ * SPDX-License-Identifier: BSD-3-Clause
4
+ *
5
+ * Redistribution and use in source and binary forms, with or without
6
+ * modification, are permitted provided that the following conditions are met:
7
+ *
8
+ * 1. Redistributions of source code must retain the above copyright notice, this
9
+ * list of conditions and the following disclaimer.
10
+ *
11
+ * 2. Redistributions in binary form must reproduce the above copyright notice,
12
+ * this list of conditions and the following disclaimer in the documentation
13
+ * and/or other materials provided with the distribution.
14
+ *
15
+ * 3. Neither the name of the copyright holder nor the names of its
16
+ * contributors may be used to endorse or promote products derived from
17
+ * this software without specific prior written permission.
18
+ *
19
+ * THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS"
20
+ * AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE
21
+ * IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE
22
+ * DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE LIABLE
23
+ * FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL
24
+ * DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR
25
+ * SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER
26
+ * CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY,
27
+ * OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
28
+ * OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
29
+ *
30
+ **************************************************************************************************/
31
+ /*! \file
32
+ \brief Utilities for packing constructing canonical CuTe stride types for 3.x mainloop params.
33
+ */
34
+
35
+ #pragma once
36
+
37
+ #include "cute/layout.hpp"
38
+ #include "cute/container/array.hpp" // cute::array
39
+ #include "cutlass/conv/convolution.h" // cutlass::conv::Operator
40
+
41
+ /////////////////////////////////////////////////////////////////////////////////////////////////
42
+
43
+ namespace cutlass {
44
+
45
+ /////////////////////////////////////////////////////////////////////////////////////////////////
46
+
47
+ // Strides without batch mode
48
+
49
+ template <class IntT>
50
+ CUTLASS_HOST_DEVICE
51
+ cute::Stride<IntT, cute::Int<1>>
52
+ make_cute_packed_stride(cute::Stride<IntT, cute::Int<1>> s, cute::Shape<int,int,int> shape_MKL) {
53
+ static_assert(std::is_integral_v<IntT>,
54
+ "Stride must have an integral type so it can be set dynamically. Static strides not supported.");
55
+ auto s_copy = s;
56
+ cute::get<0>(s_copy) = static_cast<IntT>(cute::get<1>(shape_MKL));
57
+ return s_copy;
58
+ }
59
+
60
+ template <class IntT>
61
+ CUTLASS_HOST_DEVICE
62
+ cute::Stride<cute::Int<1>, IntT>
63
+ make_cute_packed_stride(cute::Stride<cute::Int<1>, IntT> s, cute::Shape<int,int,int> shape_MKL) {
64
+ static_assert(std::is_integral_v<IntT>,
65
+ "Stride must have an integral type so it can be set dynamically. Static strides not supported.");
66
+ auto s_copy = s;
67
+ cute::get<1>(s_copy) = static_cast<IntT>(cute::get<0>(shape_MKL));
68
+ return s_copy;
69
+ }
70
+
71
+ /////////////////////////////////////////////////////////////////////////////////////////////////
72
+
73
+ // Strides with batch mode
74
+
75
+ template <class IntT>
76
+ CUTLASS_HOST_DEVICE
77
+ cute::Stride<IntT, cute::Int<1>, int64_t>
78
+ make_cute_packed_stride(cute::Stride<IntT, cute::Int<1>, int64_t> s, cute::Shape<int,int,int> shape_MKL) {
79
+ static_assert(std::is_integral_v<IntT>,
80
+ "Stride must have an integral type so it can be set dynamically. Static strides not supported.");
81
+ auto s_copy = s;
82
+ cute::get<0>(s_copy) = static_cast<IntT>(cute::get<1>(shape_MKL));
83
+ int batch_count = cute::get<2>(shape_MKL);
84
+ if (batch_count > 1) {
85
+ cute::get<2>(s_copy) = static_cast<IntT>(cute::get<0>(shape_MKL) * cute::get<1>(shape_MKL));
86
+ }
87
+ else {
88
+ cute::get<2>(s_copy) = static_cast<IntT>(0);
89
+ }
90
+ return s_copy;
91
+ }
92
+
93
+ template <class IntT>
94
+ CUTLASS_HOST_DEVICE
95
+ cute::Stride<cute::Int<1>, IntT, int64_t>
96
+ make_cute_packed_stride(cute::Stride<cute::Int<1>, IntT, int64_t> s, cute::Shape<int,int,int> shape_MKL) {
97
+ static_assert(std::is_integral_v<IntT>,
98
+ "Stride must have an integral type so it can be set dynamically. Static strides not supported.");
99
+ auto s_copy = s;
100
+ cute::get<1>(s_copy) = static_cast<IntT>(cute::get<0>(shape_MKL));
101
+ int batch_count = cute::get<2>(shape_MKL);
102
+ if (batch_count > 1) {
103
+ cute::get<2>(s_copy) = static_cast<IntT>(cute::get<0>(shape_MKL) * cute::get<1>(shape_MKL));
104
+ }
105
+ else {
106
+ cute::get<2>(s_copy) = static_cast<IntT>(0);
107
+ }
108
+ return s_copy;
109
+ }
110
+
111
+ /////////////////////////////////////////////////////////////////////////////////////////////////
112
+
113
+ // Strides with group mode
114
+
115
+ template <class StrideIntT>
116
+ CUTLASS_HOST_DEVICE
117
+ cute::Stride<StrideIntT, cute::Int<1>, cute::Int<0>>
118
+ make_cute_packed_stride(cute::Stride<StrideIntT, cute::Int<1>, cute::Int<0>> s, cute::Shape<int,int,int> shape_MKL) {
119
+ static_assert(std::is_integral_v<StrideIntT>,
120
+ "Stride must have an integral type so it can be set dynamically. Static strides not supported.");
121
+ auto s_copy = s;
122
+ cute::get<0>(s_copy) = static_cast<StrideIntT>(cute::get<1>(shape_MKL));
123
+ return s_copy;
124
+ }
125
+
126
+ template <class StrideIntT>
127
+ CUTLASS_HOST_DEVICE
128
+ cute::Stride<cute::Int<1>, StrideIntT, cute::Int<0>>
129
+ make_cute_packed_stride(cute::Stride<cute::Int<1>, StrideIntT, cute::Int<0>> s, cute::Shape<int,int,int> shape_MKL) {
130
+ static_assert(std::is_integral_v<StrideIntT>,
131
+ "Stride must have an integral type so it can be set dynamically. Static strides not supported.");
132
+ auto s_copy = s;
133
+ cute::get<1>(s_copy) = static_cast<StrideIntT>(cute::get<0>(shape_MKL));
134
+ return s_copy;
135
+ }
136
+
137
+ /////////////////////////////////////////////////////////////////////////////////////////////////
138
+
139
+ // Strides for convolutions
140
+
141
+ // Output cutlass::layout::TensorNDHWC -> rank-3 stride (InT,_1,_0)
142
+ // Note: For fprop/dgrad kernel, strides are assumed to be layout right in NZPQK/NDHWC order
143
+ // and therefore can be coalesced to just q/w. For wgrad kernel, strides are assumed to be layout
144
+ // right in KTRSC order and can be coalesced to just k.
145
+ // We enforce this condition here with asserts.
146
+ template <class IntT, size_t RankT_>
147
+ CUTLASS_HOST_DEVICE
148
+ cute::Stride<IntT, cute::Int<1>, cute::Int<0>>
149
+ make_cute_packed_stride(
150
+ cute::Stride<IntT, cute::Int<1>, cute::Int<0>> s,
151
+ cute::array<int32_t, RankT_> shape_output,
152
+ cute::array<IntT, RankT_> stride_output,
153
+ cutlass::conv::Operator conv_op) {
154
+ static_assert(std::is_integral_v<IntT>,
155
+ "Stride must have an integral type so it can be set dynamically. Static strides not supported.");
156
+ static_assert(RankT_ >= 3u);
157
+ constexpr static int RankT = static_cast<int>(RankT_);
158
+
159
+ assert(stride_output[RankT-1] == 1);
160
+ cute::for_each(cute::make_seq<RankT-2>{}, [&](auto i) {
161
+ assert(stride_output[i] == shape_output[i+1] * stride_output[i+1]);
162
+ });
163
+
164
+ auto s_copy = s;
165
+ cute::get<0>(s_copy) = (conv_op == cutlass::conv::Operator::kWgrad) ?
166
+ stride_output[0] :
167
+ stride_output[RankT-2];
168
+ return s_copy;
169
+ }
170
+
171
+ //
172
+ // Activation tensor ((w, h, d, n), _1) for fprop kernel
173
+ //
174
+
175
+ // Activation cutlass::layout::TensorNWC -> rank-2 stride ((W,N),_1)
176
+ template <class IntT>
177
+ CUTLASS_HOST_DEVICE
178
+ cute::Stride<cute::Stride<IntT, IntT>, cute::Int<1>>
179
+ make_cute_packed_stride(
180
+ cute::Stride<cute::Stride<IntT, IntT>, cute::Int<1>> s,
181
+ cute::array<IntT, 3> stride_nwc,
182
+ conv::Operator ConvOp) {
183
+ static_assert(std::is_integral_v<IntT>,
184
+ "Stride must have an integral type so it can be set dynamically. Static strides not supported.");
185
+ assert(stride_nwc[2] == 1);
186
+ auto s_copy = s;
187
+ cute::get<0,0>(s_copy) = stride_nwc[1];
188
+ cute::get<0,1>(s_copy) = stride_nwc[0];
189
+ return s_copy;
190
+ }
191
+
192
+ // Activation cutlass::layout::TensorNHWC -> rank-2 stride ((W,H,N),_1)
193
+ template <class IntT>
194
+ CUTLASS_HOST_DEVICE
195
+ cute::Stride<cute::Stride<IntT, IntT, IntT>, cute::Int<1>>
196
+ make_cute_packed_stride(
197
+ cute::Stride<cute::Stride<IntT, IntT, IntT>, cute::Int<1>> s,
198
+ cute::array<IntT, 4> stride_nhwc,
199
+ conv::Operator ConvOp) {
200
+ static_assert(std::is_integral_v<IntT>,
201
+ "Stride must have an integral type so it can be set dynamically. Static strides not supported.");
202
+ assert(stride_nhwc[3] == 1);
203
+ auto s_copy = s;
204
+ cute::for_each(cute::make_seq<3>{}, [&](auto i) {
205
+ cute::get<0,i>(s_copy) = stride_nhwc[2-i];
206
+ });
207
+ return s_copy;
208
+ }
209
+
210
+ // Activation cutlass::layout::TensorNDHWC -> rank-2 stride ((W,H,D,N),_1)
211
+ template <class IntT>
212
+ CUTLASS_HOST_DEVICE
213
+ cute::Stride<cute::Stride<IntT, IntT, IntT, IntT>, cute::Int<1>>
214
+ make_cute_packed_stride(
215
+ cute::Stride<cute::Stride<IntT, IntT, IntT, IntT>, cute::Int<1>> s,
216
+ cute::array<IntT, 5> stride_ndhwc,
217
+ conv::Operator ConvOp) {
218
+ static_assert(std::is_integral_v<IntT>,
219
+ "Stride must have an integral type so it can be set dynamically. Static strides not supported.");
220
+
221
+ assert(stride_ndhwc[4] == 1);
222
+ auto s_copy = s;
223
+ cute::for_each(cute::make_seq<4>{}, [&](auto i) {
224
+ cute::get<0,i>(s_copy) = stride_ndhwc[3-i];
225
+ });
226
+ return s_copy;
227
+ }
228
+
229
+ //
230
+ // Filter tensor (k, (_1, s, r, t)) for fprop kernel
231
+ //
232
+
233
+ // Filter cutlass::layout::TensorNWC -> rank-2 stride (k, (_1, s))
234
+ template <class IntT>
235
+ CUTLASS_HOST_DEVICE
236
+ cute::Stride<IntT, cute::Stride<cute::Int<1>, IntT>>
237
+ make_cute_packed_stride(
238
+ cute::Stride<IntT, cute::Stride<cute::Int<1>, IntT>> s,
239
+ cute::array<IntT, 3> stride_ksc,
240
+ conv::Operator ConvOp) {
241
+ static_assert(std::is_integral_v<IntT>,
242
+ "Stride must have an integral type so it can be set dynamically. Static strides not supported.");
243
+
244
+ assert(stride_ksc[2] == 1);
245
+ auto s_copy = s;
246
+ cute::get<0,0>(s_copy) = stride_ksc[0];
247
+ cute::get<1,1>(s_copy) = stride_ksc[1];
248
+ return s_copy;
249
+ }
250
+
251
+ // Filter cutlass::layout::TensorNHWC -> rank-2 stride (k, (_1, s, r))
252
+ template <class IntT>
253
+ CUTLASS_HOST_DEVICE
254
+ cute::Stride<IntT, cute::Stride<cute::Int<1>, IntT, IntT>>
255
+ make_cute_packed_stride(
256
+ cute::Stride<IntT, cute::Stride<cute::Int<1>, IntT, IntT>> s,
257
+ cute::array<IntT, 4> stride_krsc,
258
+ conv::Operator ConvOp) {
259
+ static_assert(std::is_integral_v<IntT>,
260
+ "Stride must have an integral type so it can be set dynamically. Static strides not supported.");
261
+
262
+ assert(stride_krsc[3] == 1);
263
+ auto s_copy = s;
264
+ cute::get<0,0>(s_copy) = stride_krsc[0];
265
+ cute::for_each(cute::make_seq<2>{}, [&](auto i) {
266
+ cute::get<1,2-i>(s_copy) = stride_krsc[i+1];
267
+ });
268
+ return s_copy;
269
+ }
270
+
271
+ // Filter cutlass::layout::TensorNDHWC -> rank-2 stride (k, (_1, s, r, t))
272
+ template <class IntT>
273
+ CUTLASS_HOST_DEVICE
274
+ cute::Stride<IntT, cute::Stride<cute::Int<1>, IntT, IntT, IntT>>
275
+ make_cute_packed_stride(
276
+ cute::Stride<IntT, cute::Stride<cute::Int<1>, IntT, IntT, IntT>> s,
277
+ cute::array<IntT, 5> stride_ktrsc,
278
+ conv::Operator ConvOp) {
279
+ static_assert(std::is_integral_v<IntT>,
280
+ "Stride must have an integral type so it can be set dynamically. Static strides not supported.");
281
+
282
+ assert(stride_ktrsc[4] == 1);
283
+ auto s_copy = s;
284
+ cute::get<0,0>(s_copy) = stride_ktrsc[0];
285
+ cute::for_each(cute::make_seq<3>{}, [&](auto i) {
286
+ cute::get<1,3-i>(s_copy) = stride_ktrsc[i+1];
287
+ });
288
+ return s_copy;
289
+ }
290
+
291
+ //
292
+ // Activation tensor (_1, (w, h, d, n)) for wgrad kernel
293
+ //
294
+ // It is also Filter tensor ((_1), (k, s, r, t)) for dgrad kernel
295
+ //
296
+
297
+ // Activation cutlass::layout::TensorNWC -> rank-2 stride (_1, (W,N)) in wgrad
298
+ // Filter cutlass::layout::TensorNWC -> rank-2 stride ((_1), (k, s)) in dgrad
299
+ template <class IntT>
300
+ CUTLASS_HOST_DEVICE
301
+ cute::Stride<cute::Int<1>, cute::Stride<IntT, IntT>>
302
+ make_cute_packed_stride(
303
+ cute::Stride<cute::Int<1>, cute::Stride<IntT, IntT>> s,
304
+ cute::array<IntT, 3> stride_nwc,
305
+ conv::Operator ConvOp) {
306
+ static_assert(std::is_integral_v<IntT>,
307
+ "Stride must have an integral type so it can be set dynamically. Static strides not supported.");
308
+
309
+ assert(stride_nwc[2] == 1);
310
+ auto s_copy = s;
311
+ if (ConvOp == cutlass::conv::Operator::kWgrad) {
312
+ cute::get<1,0>(s_copy) = stride_nwc[1];
313
+ cute::get<1,1>(s_copy) = stride_nwc[0];
314
+ }
315
+ else if (ConvOp == cutlass::conv::Operator::kDgrad) {
316
+ // stride_nwc in dgrad is ksc.
317
+ cute::get<1,0>(s_copy) = stride_nwc[0];
318
+ cute::get<1,1>(s_copy) = stride_nwc[1];
319
+ }
320
+ return s_copy;
321
+ }
322
+
323
+ // Activation cutlass::layout::TensorNHWC -> rank-2 stride (_1, (W,H,N)) in wgrad
324
+ // Filter cutlass::layout::TensorNHWC -> rank-2 stride ((_1), (k, s, r)) in dgrad
325
+ template <class IntT>
326
+ CUTLASS_HOST_DEVICE
327
+ cute::Stride<cute::Int<1>, cute::Stride<IntT, IntT, IntT>>
328
+ make_cute_packed_stride(
329
+ cute::Stride<cute::Int<1>, cute::Stride<IntT, IntT, IntT>> s,
330
+ cute::array<IntT, 4> stride_nhwc,
331
+ conv::Operator ConvOp) {
332
+ static_assert(std::is_integral_v<IntT>,
333
+ "Stride must have an integral type so it can be set dynamically. Static strides not supported.");
334
+
335
+ assert(stride_nhwc[3] == 1);
336
+ auto s_copy = s;
337
+ if (ConvOp == cutlass::conv::Operator::kWgrad) {
338
+ cute::for_each(cute::make_seq<3>{}, [&](auto i) {
339
+ cute::get<1,i>(s_copy) = stride_nhwc[2-i];
340
+ });
341
+ }
342
+ else if (ConvOp == cutlass::conv::Operator::kDgrad) {
343
+ // stride_nhwc in dgrad is krsc.
344
+ cute::get<1,0>(s_copy) = stride_nhwc[0];
345
+ cute::for_each(cute::make_seq<2>{}, [&](auto i) {
346
+ cute::get<1,2-i>(s_copy) = stride_nhwc[i+1];
347
+ });
348
+ }
349
+ return s_copy;
350
+ }
351
+
352
+ // Activation cutlass::layout::TensorNDHWC -> rank-2 stride (_1, (W,H,D,N)) in wgrad
353
+ // Filter cutlass::layout::TensorNDHWC -> rank-2 stride ((_1), (k, s, r, t)) in dgrad
354
+ template <class IntT>
355
+ CUTLASS_HOST_DEVICE
356
+ cute::Stride<cute::Int<1>, cute::Stride<IntT, IntT, IntT, IntT>>
357
+ make_cute_packed_stride(
358
+ cute::Stride<cute::Int<1>, cute::Stride<IntT, IntT, IntT, IntT>> s,
359
+ cute::array<IntT, 5> stride_ndhwc,
360
+ conv::Operator ConvOp) {
361
+ static_assert(std::is_integral_v<IntT>,
362
+ "Stride must have an integral type so it can be set dynamically. Static strides not supported.");
363
+
364
+ assert(stride_ndhwc[4] == 1);
365
+ auto s_copy = s;
366
+ if (ConvOp == cutlass::conv::Operator::kWgrad) {
367
+ cute::for_each(cute::make_seq<4>{}, [&](auto i) {
368
+ cute::get<1,i>(s_copy) = stride_ndhwc[3-i];
369
+ });
370
+ }
371
+ else if (ConvOp == cutlass::conv::Operator::kDgrad) {
372
+ // stride_ndhwc in dgrad is ktrsc.
373
+ cute::get<1,0>(s_copy) = stride_ndhwc[0];
374
+ cute::for_each(cute::make_seq<3>{}, [&](auto i) {
375
+ cute::get<1,3-i>(s_copy) = stride_ndhwc[i+1];
376
+ });
377
+ }
378
+ return s_copy;
379
+ }
380
+
381
+ //
382
+ // NZPQ tensor (_1, nzpq) for wgrad kernel
383
+ //
384
+
385
+ // cutlass::layout::TensorNWC -> rank-2 stride (_1, nzpq)
386
+ template <class IntT>
387
+ CUTLASS_HOST_DEVICE
388
+ cute::Stride<cute::Int<1>, IntT>
389
+ make_cute_packed_stride(
390
+ cute::Stride<cute::Int<1>, IntT> s,
391
+ cute::array<IntT, 3> stride_nqk,
392
+ conv::Operator ConvOp) {
393
+ static_assert(std::is_integral_v<IntT>,
394
+ "Stride must have an integral type so it can be set dynamically. Static strides not supported.");
395
+
396
+ assert(stride_nqk[2] == 1);
397
+ auto s_copy = s;
398
+ cute::get<1>(s_copy) = stride_nqk[1];
399
+ return s_copy;
400
+ }
401
+
402
+ // cutlass::layout::TensorNHWC -> rank-2 stride (_1, nzpq)
403
+ template <class IntT>
404
+ CUTLASS_HOST_DEVICE
405
+ cute::Stride<cute::Int<1>, IntT>
406
+ make_cute_packed_stride(
407
+ cute::Stride<cute::Int<1>, IntT> s,
408
+ cute::array<IntT, 4> stride_npqk,
409
+ conv::Operator ConvOp) {
410
+ static_assert(std::is_integral_v<IntT>,
411
+ "Stride must have an integral type so it can be set dynamically. Static strides not supported.");
412
+
413
+ assert(stride_npqk[3] == 1);
414
+ auto s_copy = s;
415
+ cute::get<1>(s_copy) = stride_npqk[2];
416
+ return s_copy;
417
+ }
418
+
419
+ // cutlass::layout::TensorNDHWC -> rank-2 stride (_1, nzpq)
420
+ template <class IntT>
421
+ CUTLASS_HOST_DEVICE
422
+ cute::Stride<cute::Int<1>, IntT>
423
+ make_cute_packed_stride(
424
+ cute::Stride<cute::Int<1>, IntT> s,
425
+ cute::array<IntT, 5> stride_nzpqk,
426
+ conv::Operator ConvOp) {
427
+ static_assert(std::is_integral_v<IntT>,
428
+ "Stride must have an integral type so it can be set dynamically. Static strides not supported.");
429
+
430
+ assert(stride_nzpqk[4] == 1);
431
+ auto s_copy = s;
432
+ cute::get<1>(s_copy) = stride_nzpqk[3];
433
+ return s_copy;
434
+ }
435
+
436
+
437
+
438
+ //
439
+ // Wgrad output tensor (k, (_1, s, r, t), _0)
440
+ //
441
+
442
+ // Filter cutlass::layout::TensorKCS -> rank-3 stride (k, (_1, s), _0)
443
+ template <class IntT>
444
+ CUTLASS_HOST_DEVICE
445
+ cute::Stride<IntT, cute::Stride<cute::Int<1>, IntT>, cute::Int<0>>
446
+ make_cute_packed_stride(
447
+ cute::Stride<IntT, cute::Stride<cute::Int<1>, IntT>, cute::Int<0>> s,
448
+ [[maybe_unused]] cute::array<int32_t, 3> shape_output,
449
+ cute::array<IntT, 3> stride_ksc,
450
+ conv::Operator ConvOp) {
451
+ static_assert(std::is_integral_v<IntT>,
452
+ "Stride must have an integral type so it can be set dynamically. Static strides not supported.");
453
+
454
+ assert(stride_ksc[2] == 1);
455
+ auto s_copy = s;
456
+ cute::get<0,0>(s_copy) = stride_ksc[0];
457
+ cute::get<1,1>(s_copy) = stride_ksc[1];
458
+ return s_copy;
459
+ }
460
+
461
+ // Filter cutlass::layout::TensorKCSR -> rank-3 stride (k, (_1, s, r), _0)
462
+ template <class IntT>
463
+ CUTLASS_HOST_DEVICE
464
+ cute::Stride<IntT, cute::Stride<cute::Int<1>, IntT, IntT>, cute::Int<0>>
465
+ make_cute_packed_stride(
466
+ cute::Stride<IntT, cute::Stride<cute::Int<1>, IntT, IntT>, cute::Int<0>> s,
467
+ [[maybe_unused]] cute::array<int32_t, 4> shape_output,
468
+ cute::array<IntT, 4> stride_krsc,
469
+ conv::Operator ConvOp) {
470
+ static_assert(std::is_integral_v<IntT>,
471
+ "Stride must have an integral type so it can be set dynamically. Static strides not supported.");
472
+
473
+ assert(stride_krsc[3] == 1);
474
+ auto s_copy = s;
475
+ cute::get<0,0>(s_copy) = stride_krsc[0];
476
+ cute::for_each(cute::make_seq<2>{}, [&](auto i) {
477
+ cute::get<1,2-i>(s_copy) = stride_krsc[i+1];
478
+ });
479
+ return s_copy;
480
+ }
481
+
482
+ // Filter cutlass::layout::TensorKCSRT -> rank-3 stride (k, (_1, s, r, t), _0)
483
+ template <class IntT>
484
+ CUTLASS_HOST_DEVICE
485
+ cute::Stride<IntT, cute::Stride<cute::Int<1>, IntT, IntT, IntT>, cute::Int<0>>
486
+ make_cute_packed_stride(
487
+ cute::Stride<IntT, cute::Stride<cute::Int<1>, IntT, IntT, IntT>, cute::Int<0>> s,
488
+ [[maybe_unused]] cute::array<int32_t, 5> shape_output,
489
+ cute::array<IntT, 5> stride_ktrsc,
490
+ conv::Operator ConvOp) {
491
+ static_assert(std::is_integral_v<IntT>,
492
+ "Stride must have an integral type so it can be set dynamically. Static strides not supported.");
493
+
494
+ assert(stride_ktrsc[4] == 1);
495
+ auto s_copy = s;
496
+ cute::get<0,0>(s_copy) = stride_ktrsc[0];
497
+ cute::for_each(cute::make_seq<3>{}, [&](auto i) {
498
+ cute::get<1,3-i>(s_copy) = stride_ktrsc[i+1];
499
+ });
500
+ return s_copy;
501
+ }
502
+
503
+
504
+ //
505
+ // Wgrad output tensor ((_1, s, r, t), k, _0)
506
+ //
507
+
508
+ // Filter cutlass::layout::TensorCSK -> rank-3 stride ((_1, s), k, _0)
509
+ template <class IntT>
510
+ CUTLASS_HOST_DEVICE
511
+ cute::Stride<cute::Stride<cute::Int<1>, IntT>, IntT, cute::Int<0>>
512
+ make_cute_packed_stride(
513
+ cute::Stride<cute::Stride<cute::Int<1>, IntT>, IntT, cute::Int<0>> s,
514
+ [[maybe_unused]] cute::array<int32_t, 3> shape_output,
515
+ cute::array<IntT, 3> stride_ksc,
516
+ conv::Operator ConvOp) {
517
+ static_assert(std::is_integral_v<IntT>,
518
+ "Stride must have an integral type so it can be set dynamically. Static strides not supported.");
519
+
520
+ assert(stride_ksc[2] == 1);
521
+ auto s_copy = s;
522
+ cute::get<1,0>(s_copy) = stride_ksc[0];
523
+ cute::get<0,1>(s_copy) = stride_ksc[1];
524
+ return s_copy;
525
+ }
526
+
527
+ // Filter cutlass::layout::TensorCSRK -> rank-3 stride ((_1, s, r), k, _0)
528
+ template <class IntT>
529
+ CUTLASS_HOST_DEVICE
530
+ cute::Stride<cute::Stride<cute::Int<1>, IntT, IntT>, IntT, cute::Int<0>>
531
+ make_cute_packed_stride(
532
+ cute::Stride<cute::Stride<cute::Int<1>, IntT, IntT>, IntT, cute::Int<0>> s,
533
+ [[maybe_unused]] cute::array<int32_t, 4> shape_output,
534
+ cute::array<IntT, 4> stride_krsc,
535
+ conv::Operator ConvOp) {
536
+ static_assert(std::is_integral_v<IntT>,
537
+ "Stride must have an integral type so it can be set dynamically. Static strides not supported.");
538
+
539
+ assert(stride_krsc[3] == 1);
540
+ auto s_copy = s;
541
+ cute::get<1,0>(s_copy) = stride_krsc[0];
542
+ cute::for_each(cute::make_seq<2>{}, [&](auto i) {
543
+ cute::get<0,2-i>(s_copy) = stride_krsc[i+1];
544
+ });
545
+ return s_copy;
546
+ }
547
+
548
+ // Filter cutlass::layout::TensorCSRTK -> rank-3 stride ((_1, s, r, t), k, _0)
549
+ template <class IntT>
550
+ CUTLASS_HOST_DEVICE
551
+ cute::Stride<cute::Stride<cute::Int<1>, IntT, IntT, IntT>, IntT, cute::Int<0>>
552
+ make_cute_packed_stride(
553
+ cute::Stride<cute::Stride<cute::Int<1>, IntT, IntT, IntT>, IntT, cute::Int<0>> s,
554
+ [[maybe_unused]] cute::array<int32_t, 5> shape_output,
555
+ cute::array<IntT, 5> stride_ktrsc,
556
+ conv::Operator ConvOp) {
557
+ static_assert(std::is_integral_v<IntT>,
558
+ "Stride must have an integral type so it can be set dynamically. Static strides not supported.");
559
+
560
+ assert(stride_ktrsc[4] == 1);
561
+ auto s_copy = s;
562
+ cute::get<1,0>(s_copy) = stride_ktrsc[0];
563
+ cute::for_each(cute::make_seq<3>{}, [&](auto i) {
564
+ cute::get<0,3-i>(s_copy) = stride_ktrsc[i+1];
565
+ });
566
+ return s_copy;
567
+ }
568
+ /////////////////////////////////////////////////////////////////////////////////////////////////
569
+
570
+ } // namespace cutlass
csrc/dequantize_fp4_sfa.cu ADDED
@@ -0,0 +1,89 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ // SPDX-License-Identifier: Apache-2.0
2
+ #include "dequantize_fp4_sfa.cuh"
3
+
4
+ #include <cuda_fp8.h>
5
+
6
+ #ifndef CUTLASS_ARCH_MMA_SM100_SUPPORTED
7
+ # define CUTLASS_ARCH_MMA_SM100_SUPPORTED 1
8
+ #endif
9
+ #if defined(CUTLASS_ARCH_MMA_SM100_SUPPORTED) || defined(__CUDA_ARCH__)
10
+ # include "cutlass/cutlass.h"
11
+ # include "cutlass/detail/sm100_blockscaled_layout.hpp"
12
+ # include "cute/tensor.hpp"
13
+ # define FV_HAVE_CUTLASS 1
14
+ #else
15
+ # define FV_HAVE_CUTLASS 0
16
+ #endif
17
+
18
+ namespace flash_rt {
19
+ namespace fused_fp4 {
20
+
21
+ #if FV_HAVE_CUTLASS
22
+
23
+ using CfgDequant = cutlass::detail::Sm1xxBlockScaledConfig<16>;
24
+
25
+ __device__ __forceinline__ float e2m1_to_fp32_dequant(uint8_t value) {
26
+ static constexpr float mags[8] = {0.f, 0.5f, 1.f, 1.5f, 2.f, 3.f, 4.f, 6.f};
27
+ float mag = mags[value & 0x7];
28
+ return (value & 0x8) ? -mag : mag;
29
+ }
30
+
31
+ template <class LayoutSF>
32
+ __global__ void dequantize_fp4_sfa_kernel(
33
+ const uint8_t* __restrict__ packed,
34
+ const uint8_t* __restrict__ sfa,
35
+ __half* __restrict__ out,
36
+ LayoutSF layout,
37
+ int dim) {
38
+ int row = blockIdx.y;
39
+ int block_idx = blockIdx.x * blockDim.x + threadIdx.x;
40
+ int n_blocks = dim / 16;
41
+ if (block_idx >= n_blocks) return;
42
+
43
+ int col_base = block_idx * 16;
44
+ int sfa_off = layout(row, col_base, 0);
45
+ __nv_fp8_e4m3 scale_q;
46
+ *reinterpret_cast<uint8_t*>(&scale_q) = sfa[sfa_off];
47
+ float scale = static_cast<float>(scale_q);
48
+
49
+ const uint8_t* packed_block = packed + row * (dim / 2) + block_idx * 8;
50
+ __half* out_block = out + row * dim + col_base;
51
+ #pragma unroll
52
+ for (int p = 0; p < 8; ++p) {
53
+ uint8_t byte = packed_block[p];
54
+ out_block[2 * p] = __float2half(e2m1_to_fp32_dequant(byte & 0xF) * scale);
55
+ out_block[2 * p + 1] = __float2half(e2m1_to_fp32_dequant(byte >> 4) * scale);
56
+ }
57
+ }
58
+
59
+ #endif
60
+
61
+ void dequantize_fp4_sfa_fp16(
62
+ const uint8_t* packed,
63
+ const uint8_t* sfa,
64
+ __half* out,
65
+ int rows,
66
+ int dim,
67
+ bool is_sfb,
68
+ cudaStream_t stream) {
69
+ #if FV_HAVE_CUTLASS
70
+ int n_blocks = dim / 16;
71
+ dim3 block(256);
72
+ dim3 grid((n_blocks + block.x - 1) / block.x, rows);
73
+ auto shape = cute::make_shape(is_sfb ? 1 : rows, is_sfb ? rows : 1, dim, 1);
74
+ if (is_sfb) {
75
+ auto layout = CfgDequant::tile_atom_to_shape_SFB(shape);
76
+ dequantize_fp4_sfa_kernel<<<grid, block, 0, stream>>>(
77
+ packed, sfa, out, layout, dim);
78
+ } else {
79
+ auto layout = CfgDequant::tile_atom_to_shape_SFA(shape);
80
+ dequantize_fp4_sfa_kernel<<<grid, block, 0, stream>>>(
81
+ packed, sfa, out, layout, dim);
82
+ }
83
+ #else
84
+ (void)packed; (void)sfa; (void)out; (void)rows; (void)dim; (void)is_sfb; (void)stream;
85
+ #endif
86
+ }
87
+
88
+ } // namespace fused_fp4
89
+ } // namespace flash_rt
csrc/dequantize_fp4_sfa.cuh ADDED
@@ -0,0 +1,21 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ // SPDX-License-Identifier: Apache-2.0
2
+ #pragma once
3
+
4
+ #include <cstdint>
5
+ #include <cuda_runtime.h>
6
+ #include <cuda_fp16.h>
7
+
8
+ namespace flash_rt {
9
+ namespace fused_fp4 {
10
+
11
+ void dequantize_fp4_sfa_fp16(
12
+ const uint8_t* packed,
13
+ const uint8_t* sfa,
14
+ __half* out,
15
+ int rows,
16
+ int dim,
17
+ bool is_sfb,
18
+ cudaStream_t stream);
19
+
20
+ } // namespace fused_fp4
21
+ } // namespace flash_rt
csrc/gemm/fp4/cutlass_fp4_gemm_bias_bf16_sm100.cu ADDED
@@ -0,0 +1,208 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ // ============================================================================
2
+ // NVFP4 GEMMs with bf16 fused-bias epilogues. See header for the contract.
3
+ //
4
+ // Three fusions over one shared skinny-M mainloop config (Sm100,
5
+ // tile 128x64x256, cluster 1x1x1 — the narrow-N + wide-K shape that wins
6
+ // at small M where the GEMM is weight-bandwidth-bound):
7
+ // bias: LinCombPerColBias (bf16 out, beta = 0)
8
+ // bias+res: LinCombPerColBias (bf16 out, beta = 1, C = residual)
9
+ // bias+gelu: LinCombPerColBiasEltActBlockScaleFactor<GELU_taylor>
10
+ // (fp4 + SFA out for the following NVFP4 GEMM)
11
+ // ============================================================================
12
+ #include "gemm/fp4/cutlass_fp4_gemm_bias_bf16_sm100.cuh"
13
+
14
+ #include "cutlass/cutlass.h"
15
+ #include "cutlass/epilogue/thread/activation.h"
16
+ #include "cutlass/epilogue/dispatch_policy.hpp"
17
+ #include "cutlass/epilogue/fusion/operations.hpp"
18
+ #include "cutlass/gemm/dispatch_policy.hpp"
19
+ #include "cutlass/gemm/collective/collective_builder.hpp"
20
+ #include "cutlass/epilogue/collective/collective_builder.hpp"
21
+ #include "cutlass/gemm/device/gemm_universal_adapter.h"
22
+ #include "cutlass/gemm/kernel/gemm_universal.hpp"
23
+ #include "cutlass/util/packed_stride.hpp"
24
+ #include "cutlass/detail/sm100_blockscaled_layout.hpp"
25
+ #include "cute/tensor.hpp"
26
+
27
+ namespace flash_rt {
28
+ namespace fp4 {
29
+
30
+ namespace bias_bf16_gemm {
31
+
32
+ using namespace cute;
33
+
34
+ using ElementA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
35
+ using LayoutATag = cutlass::layout::RowMajor;
36
+ constexpr int AlignmentA = 32;
37
+
38
+ using ElementB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
39
+ using LayoutBTag = cutlass::layout::ColumnMajor;
40
+ constexpr int AlignmentB = 32;
41
+
42
+ using ElementAccumulator = float;
43
+ using ElementCompute = float;
44
+ using ArchTag = cutlass::arch::Sm100;
45
+ using OperatorClass = cutlass::arch::OpClassBlockScaledTensorOp;
46
+ constexpr int SFVecSize = 16;
47
+
48
+ using MmaTileShape = Shape<_128, _64, _256>;
49
+ using ClusterShape = Shape<_1, _1, _1>;
50
+
51
+ template <class FusionOp, class ElemC, class ElemD, int AlignCD>
52
+ struct BiasGemm {
53
+ using CollectiveEpilogue =
54
+ typename cutlass::epilogue::collective::CollectiveBuilder<
55
+ ArchTag, OperatorClass, MmaTileShape, ClusterShape,
56
+ cutlass::epilogue::collective::EpilogueTileAuto,
57
+ ElementAccumulator, ElementAccumulator,
58
+ ElemC, cutlass::layout::RowMajor, AlignCD,
59
+ ElemD, cutlass::layout::RowMajor, AlignCD,
60
+ cutlass::epilogue::collective::EpilogueScheduleAuto,
61
+ FusionOp>::CollectiveOp;
62
+
63
+ using CollectiveMainloop =
64
+ typename cutlass::gemm::collective::CollectiveBuilder<
65
+ ArchTag, OperatorClass,
66
+ ElementA, LayoutATag, AlignmentA,
67
+ ElementB, LayoutBTag, AlignmentB,
68
+ ElementAccumulator, MmaTileShape, ClusterShape,
69
+ cutlass::gemm::collective::StageCountAutoCarveout<static_cast<int>(
70
+ sizeof(typename CollectiveEpilogue::SharedStorage))>,
71
+ cutlass::gemm::collective::KernelScheduleAuto>::CollectiveOp;
72
+
73
+ using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
74
+ Shape<int, int, int, int>, CollectiveMainloop, CollectiveEpilogue, void>;
75
+ using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
76
+ };
77
+
78
+ // ── bias / bias+res: bf16 out ──────────────────────────────────────────────
79
+ using ElementCD = cutlass::bfloat16_t;
80
+ using FusionBias = cutlass::epilogue::fusion::LinCombPerColBias<
81
+ ElementCD, ElementCompute, ElementCD, ElementCD, ElementCompute>;
82
+ using GemmBias = BiasGemm<FusionBias, ElementCD, ElementCD, 8>::Gemm;
83
+
84
+ // ── bias+gelu: fp4 + SFA out ───────────────────────────────────────────────
85
+ using ElementDQ = cutlass::float_e2m1_t;
86
+ using ElementSFD = cutlass::float_ue4m3_t;
87
+ using FusionGelu =
88
+ cutlass::epilogue::fusion::LinCombPerColBiasEltActBlockScaleFactor<
89
+ cutlass::epilogue::thread::GELU_taylor, SFVecSize,
90
+ ElementDQ, ElementCompute, ElementSFD, cutlass::layout::RowMajor,
91
+ ElementCD, ElementDQ, ElementCompute>;
92
+ using GemmGelu = BiasGemm<FusionGelu, ElementDQ, ElementDQ, 32>::Gemm;
93
+
94
+ template <class Gemm>
95
+ static int run_gemm(typename Gemm::Arguments& args, cudaStream_t stream) {
96
+ Gemm gemm;
97
+ auto st = gemm.can_implement(args);
98
+ if (st != cutlass::Status::kSuccess) return static_cast<int>(st) | 0x10000;
99
+ size_t ws_sz = Gemm::get_workspace_size(args);
100
+ void* ws = nullptr;
101
+ if (ws_sz > 0 && cudaMalloc(&ws, ws_sz) != cudaSuccess) return -1;
102
+ st = gemm.initialize(args, ws, stream);
103
+ if (st != cutlass::Status::kSuccess) {
104
+ if (ws) cudaFree(ws);
105
+ return static_cast<int>(st) | 0x20000;
106
+ }
107
+ st = gemm.run(stream);
108
+ if (ws) cudaFree(ws);
109
+ return (st == cutlass::Status::kSuccess) ? 0
110
+ : (static_cast<int>(st) | 0x30000);
111
+ }
112
+
113
+ template <class Gemm, class ElemC, class ElemD>
114
+ static typename Gemm::Arguments make_args(
115
+ void const* A, void const* SFA, void const* B, void const* SFB,
116
+ void const* C, void* D, int M, int N, int K) {
117
+ auto stride_A = cutlass::make_cute_packed_stride(
118
+ typename Gemm::GemmKernel::StrideA{}, {M, K, 1});
119
+ auto stride_B = cutlass::make_cute_packed_stride(
120
+ typename Gemm::GemmKernel::StrideB{}, {N, K, 1});
121
+ auto stride_C = cutlass::make_cute_packed_stride(
122
+ typename Gemm::GemmKernel::StrideC{}, {M, N, 1});
123
+ auto stride_D = cutlass::make_cute_packed_stride(
124
+ typename Gemm::GemmKernel::StrideD{}, {M, N, 1});
125
+ using Cfg =
126
+ typename Gemm::GemmKernel::CollectiveMainloop::Sm1xxBlkScaledConfig;
127
+ auto layout_SFA = Cfg::tile_atom_to_shape_SFA(make_shape(M, N, K, 1));
128
+ auto layout_SFB = Cfg::tile_atom_to_shape_SFB(make_shape(M, N, K, 1));
129
+
130
+ using EA = typename ElementA::DataType;
131
+ using SA = typename ElementA::ScaleFactorType;
132
+
133
+ return typename Gemm::Arguments{
134
+ cutlass::gemm::GemmUniversalMode::kGemm, {M, N, K, 1},
135
+ {reinterpret_cast<EA const*>(A), stride_A,
136
+ reinterpret_cast<EA const*>(B), stride_B,
137
+ reinterpret_cast<SA const*>(SFA), layout_SFA,
138
+ reinterpret_cast<SA const*>(SFB), layout_SFB},
139
+ {{},
140
+ reinterpret_cast<ElemC const*>(C), stride_C,
141
+ reinterpret_cast<ElemD*>(D), stride_D}};
142
+ }
143
+
144
+ } // namespace bias_bf16_gemm
145
+
146
+ int cutlass_fp4_gemm_bias_bf16(
147
+ void const* A_packed, void const* SFA,
148
+ void const* B_packed, void const* SFB,
149
+ void const* bias_bf16,
150
+ void* D_bf16,
151
+ int M, int N, int K, cudaStream_t stream) {
152
+ using namespace bias_bf16_gemm;
153
+ auto args = make_args<GemmBias, ElementCD, ElementCD>(
154
+ A_packed, SFA, B_packed, SFB, D_bf16, D_bf16, M, N, K);
155
+ args.epilogue.thread.alpha = 1.0f;
156
+ args.epilogue.thread.beta = 0.0f;
157
+ args.epilogue.thread.bias_ptr =
158
+ reinterpret_cast<ElementCD const*>(bias_bf16);
159
+ return run_gemm<GemmBias>(args, stream);
160
+ }
161
+
162
+ int cutlass_fp4_gemm_bias_res_bf16(
163
+ void const* A_packed, void const* SFA,
164
+ void const* B_packed, void const* SFB,
165
+ void const* bias_bf16,
166
+ void const* C_bf16, void* D_bf16,
167
+ int M, int N, int K, cudaStream_t stream) {
168
+ using namespace bias_bf16_gemm;
169
+ auto args = make_args<GemmBias, ElementCD, ElementCD>(
170
+ A_packed, SFA, B_packed, SFB, C_bf16, D_bf16, M, N, K);
171
+ args.epilogue.thread.alpha = 1.0f;
172
+ args.epilogue.thread.beta = 1.0f;
173
+ args.epilogue.thread.bias_ptr =
174
+ reinterpret_cast<ElementCD const*>(bias_bf16);
175
+ return run_gemm<GemmBias>(args, stream);
176
+ }
177
+
178
+ int cutlass_fp4_gemm_bias_gelu_fp4out_bf16(
179
+ void const* A_packed, void const* SFA,
180
+ void const* B_packed, void const* SFB,
181
+ void const* bias_bf16,
182
+ void* D_packed, void* D_SFD,
183
+ int M, int N, int K, cudaStream_t stream) {
184
+ using namespace bias_bf16_gemm;
185
+ auto args = make_args<GemmGelu, ElementDQ, ElementDQ>(
186
+ A_packed, SFA, B_packed, SFB, D_packed, D_packed, M, N, K);
187
+ args.epilogue.thread.alpha = 1.0f;
188
+ args.epilogue.thread.beta = 0.0f;
189
+ args.epilogue.thread.bias_ptr =
190
+ reinterpret_cast<ElementCD const*>(bias_bf16);
191
+ // The block-scale epilogue divides by a device-resident norm constant;
192
+ // 1.0 keeps the native per-16 dynamic scale. Allocated once, first call
193
+ // must happen before any CUDA Graph capture (warmup covers this).
194
+ static float* d_norm = nullptr;
195
+ if (!d_norm) {
196
+ if (cudaMalloc(&d_norm, sizeof(float)) != cudaSuccess) return -1;
197
+ float h = 1.0f;
198
+ cudaMemcpyAsync(d_norm, &h, sizeof(float), cudaMemcpyHostToDevice,
199
+ stream);
200
+ }
201
+ args.epilogue.thread.block_scale_factor_ptr =
202
+ reinterpret_cast<bias_bf16_gemm::ElementSFD*>(D_SFD);
203
+ args.epilogue.thread.norm_constant_ptr = d_norm;
204
+ return run_gemm<GemmGelu>(args, stream);
205
+ }
206
+
207
+ } // namespace fp4
208
+ } // namespace flash_rt
csrc/gemm/fp4/cutlass_fp4_gemm_bias_bf16_sm100.cuh ADDED
@@ -0,0 +1,47 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ // ============================================================================
2
+ // FlashRT — NVFP4 GEMMs with bf16 fused-bias epilogues (SM100/SM110).
3
+ //
4
+ // bf16 companions of the fp16 fused-epilogue NVFP4 GEMMs, for pipelines
5
+ // whose activations and biases are bf16 (GR00T N1.7 DiT). All three share
6
+ // the proven skinny-M block-scaled mainloop (tile 128x64x256, cluster
7
+ // 1x1x1). A is row-major [M, K] packed e2m1, B is column-major [N, K]
8
+ // packed e2m1 (A @ B^T, nn.Linear convention); SFA/SFB use the CUTLASS
9
+ // Sm1xx tile-interleaved UE4M3 layout.
10
+ //
11
+ // Additive: new symbols only; no existing kernel is modified.
12
+ // ============================================================================
13
+ #pragma once
14
+
15
+ #include <cuda_runtime.h>
16
+
17
+ namespace flash_rt {
18
+ namespace fp4 {
19
+
20
+ // D_bf16[M,N] = A @ B^T + bias[N]
21
+ int cutlass_fp4_gemm_bias_bf16(
22
+ void const* A_packed, void const* SFA,
23
+ void const* B_packed, void const* SFB,
24
+ void const* bias_bf16,
25
+ void* D_bf16,
26
+ int M, int N, int K, cudaStream_t stream);
27
+
28
+ // D_bf16[M,N] = A @ B^T + bias[N] + C_bf16[M,N] (residual; C may alias D)
29
+ int cutlass_fp4_gemm_bias_res_bf16(
30
+ void const* A_packed, void const* SFA,
31
+ void const* B_packed, void const* SFB,
32
+ void const* bias_bf16,
33
+ void const* C_bf16, void* D_bf16,
34
+ int M, int N, int K, cudaStream_t stream);
35
+
36
+ // D_fp4[M,N], SFD = blockscale(gelu_tanh(A @ B^T + bias[N]))
37
+ // SFD is written in the SFA tile-interleaved layout over (M, N) so the
38
+ // output can feed the K side of a following NVFP4 GEMM directly.
39
+ int cutlass_fp4_gemm_bias_gelu_fp4out_bf16(
40
+ void const* A_packed, void const* SFA,
41
+ void const* B_packed, void const* SFB,
42
+ void const* bias_bf16,
43
+ void* D_packed, void* D_SFD,
44
+ int M, int N, int K, cudaStream_t stream);
45
+
46
+ } // namespace fp4
47
+ } // namespace flash_rt
csrc/gemm/fp4/cutlass_nvfp4_gemm_bias_gelu_bf16out_sm120.cu ADDED
@@ -0,0 +1,212 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ // SPDX-License-Identifier: Apache-2.0
2
+ //
3
+ // CUTLASS NVFP4 W4A16 GEMM with fused per-col-bias + GELU(tanh) epilogue,
4
+ // BF16 output, SM120a. Recipe C step 1.
5
+ //
6
+ // Replaces the (cutlass NVFP4 GEMM_up + bias_gelu_inplace_bf16) 2-launch
7
+ // chain segment in the Wan FFN forward (motus). At M=360 K=3072 N=14336
8
+ // the fused kernel ships at ~32 µs/call vs ~41 µs for the 2-launch chain
9
+ // (1.28× standalone, ~-0.5 ms E2E wall per replay in CUDA graph mode).
10
+ //
11
+ // Schedule: KernelTmaWarpSpecializedPingpong + PersistentScheduler — picked
12
+ // from the empirical sweep over {coop, pingpong} × {persistent, streamk}
13
+ // at production shape (pingpong wins by ~0.6 µs/call).
14
+ //
15
+ // TileShape <128,128,256> ClusterShape <1,1,1>: locked by cutlass v4.4
16
+ // NVFP4 sm_120 BlockScaled (all unit tests use this tile; other tiles
17
+ // fail TMA atom constraints).
18
+
19
+ #include "cutlass_nvfp4_gemm_bias_gelu_bf16out_sm120.cuh"
20
+
21
+ #include "cute/tensor.hpp"
22
+
23
+ #include "cutlass/cutlass.h"
24
+ #include "cutlass/numeric_types.h"
25
+ #include "cutlass/detail/sm100_blockscaled_layout.hpp"
26
+
27
+ #include "cutlass/epilogue/collective/collective_builder.hpp"
28
+ #include "cutlass/epilogue/thread/activation.h"
29
+ #include "cutlass/epilogue/fusion/operations.hpp"
30
+
31
+ #include "cutlass/gemm/collective/collective_builder.hpp"
32
+ #include "cutlass/gemm/dispatch_policy.hpp"
33
+ #include "cutlass/gemm/device/gemm_universal_adapter.h"
34
+ #include "cutlass/gemm/kernel/gemm_universal.hpp"
35
+
36
+ #include "cutlass/util/packed_stride.hpp"
37
+
38
+ #include <cstdio>
39
+ #include <mutex>
40
+ #include <unordered_map>
41
+
42
+ namespace flash_rt {
43
+ namespace gemm {
44
+
45
+ namespace {
46
+ using namespace cute;
47
+
48
+ using ElementA = cutlass::float_e2m1_t;
49
+ using ElementB = cutlass::float_e2m1_t;
50
+ using ElementC = cutlass::bfloat16_t;
51
+ using ElementD = cutlass::bfloat16_t;
52
+ using ElementBias = cutlass::bfloat16_t;
53
+ using ElementAccumulator = float;
54
+ using ElementCompute = float;
55
+ using ElementSF = cutlass::float_ue4m3_t;
56
+
57
+ using LayoutA = cutlass::layout::RowMajor;
58
+ using LayoutB = cutlass::layout::ColumnMajor;
59
+ using LayoutC = cutlass::layout::RowMajor;
60
+ using LayoutD = cutlass::layout::RowMajor;
61
+
62
+ using ElementPairA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
63
+ using ElementPairB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
64
+
65
+ constexpr int AlignmentA = 16 * 8 / cutlass::sizeof_bits<ElementA>::value;
66
+ constexpr int AlignmentB = 16 * 8 / cutlass::sizeof_bits<ElementB>::value;
67
+ constexpr int AlignmentC = 128 / cutlass::sizeof_bits<ElementC>::value;
68
+ constexpr int AlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
69
+
70
+ // TileShape <128,128,256>: E2E winner. Standalone bench on random inputs
71
+ // shows Tile<128,128,128>+coop is 1.06 µs faster per call (32.84 vs 33.90)
72
+ // but in CUDA graph mode E2E the 256-K tile + pingpong is 0.3 ms wall
73
+ // faster — graph scheduler reshapes the cost picture vs standalone.
74
+ using TileShape = Shape<_128, _128, _256>;
75
+ using ClusterShape = Shape<_1, _1, _1>;
76
+
77
+ using Sm1xxBlkScaledConfig = cutlass::detail::Sm1xxBlockScaledConfig<16>;
78
+
79
+ // D = GELU_tanh(alpha * acc + per_col_bias).
80
+ using FusionOperation = cutlass::epilogue::fusion::LinCombPerColBiasEltAct<
81
+ cutlass::epilogue::thread::GELU_taylor,
82
+ ElementD, ElementCompute, ElementBias, ElementC>;
83
+
84
+ using CollectiveEpilogue =
85
+ typename cutlass::epilogue::collective::CollectiveBuilder<
86
+ cutlass::arch::Sm120, cutlass::arch::OpClassTensorOp,
87
+ TileShape, ClusterShape,
88
+ cutlass::epilogue::collective::EpilogueTileAuto,
89
+ ElementAccumulator, ElementCompute,
90
+ ElementC, LayoutC, AlignmentC,
91
+ ElementD, LayoutD, AlignmentD,
92
+ cutlass::epilogue::collective::EpilogueScheduleAuto,
93
+ FusionOperation
94
+ >::CollectiveOp;
95
+
96
+ using CollectiveMainloop =
97
+ typename cutlass::gemm::collective::CollectiveBuilder<
98
+ cutlass::arch::Sm120, cutlass::arch::OpClassBlockScaledTensorOp,
99
+ ElementPairA, LayoutA, AlignmentA,
100
+ ElementPairB, LayoutB, AlignmentB,
101
+ ElementAccumulator,
102
+ TileShape, ClusterShape,
103
+ cutlass::gemm::collective::StageCountAutoCarveout<
104
+ static_cast<int>(sizeof(typename CollectiveEpilogue::SharedStorage))>,
105
+ cutlass::gemm::KernelTmaWarpSpecializedPingpong
106
+ >::CollectiveOp;
107
+
108
+ using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
109
+ Shape<int, int, int, int>, CollectiveMainloop, CollectiveEpilogue,
110
+ cutlass::gemm::PersistentScheduler>;
111
+
112
+ using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
113
+
114
+ struct ShapeKey {
115
+ int M, N, K;
116
+ bool operator==(const ShapeKey& o) const {
117
+ return M == o.M && N == o.N && K == o.K;
118
+ }
119
+ };
120
+ struct SHash {
121
+ size_t operator()(const ShapeKey& k) const noexcept {
122
+ return (size_t(k.M) * 1315423911u) ^ (size_t(k.N) * 2654435761u)
123
+ ^ size_t(k.K);
124
+ }
125
+ };
126
+ struct CachedWs { void* ptr = nullptr; size_t size = 0; };
127
+ std::unordered_map<ShapeKey, CachedWs, SHash> g_ws;
128
+ std::mutex g_mu;
129
+
130
+ void* get_ws(int M, int N, int K, size_t need) {
131
+ std::lock_guard<std::mutex> lk(g_mu);
132
+ ShapeKey k{M, N, K};
133
+ auto it = g_ws.find(k);
134
+ if (it != g_ws.end() && it->second.size >= need) return it->second.ptr;
135
+ if (it != g_ws.end()) { cudaFree(it->second.ptr); g_ws.erase(it); }
136
+ CachedWs w; w.size = need;
137
+ if (need > 0) cudaMalloc(&w.ptr, need);
138
+ g_ws[k] = w;
139
+ return w.ptr;
140
+ }
141
+
142
+ } // namespace
143
+
144
+ void fp4_w4a16_gemm_bias_gelu_bf16out_sm120(
145
+ const void* A_packed, const void* B_packed,
146
+ const void* SFA, const void* SFB,
147
+ const void* bias_bf16,
148
+ void* D_bf16,
149
+ int M, int N, int K,
150
+ float alpha,
151
+ cudaStream_t stream)
152
+ {
153
+ using StrideA = typename Gemm::GemmKernel::StrideA;
154
+ using StrideB = typename Gemm::GemmKernel::StrideB;
155
+ using StrideC = typename Gemm::GemmKernel::StrideC;
156
+ using StrideD = typename Gemm::GemmKernel::StrideD;
157
+ StrideA strA = cutlass::make_cute_packed_stride(StrideA{}, cute::make_shape(M, K, 1));
158
+ StrideB strB = cutlass::make_cute_packed_stride(StrideB{}, cute::make_shape(N, K, 1));
159
+ StrideC strC = cutlass::make_cute_packed_stride(StrideC{}, cute::make_shape(M, N, 1));
160
+ StrideD strD = cutlass::make_cute_packed_stride(StrideD{}, cute::make_shape(M, N, 1));
161
+ auto problem = cute::make_shape(M, N, K, 1);
162
+ auto layout_SFA = Sm1xxBlkScaledConfig::tile_atom_to_shape_SFA(problem);
163
+ auto layout_SFB = Sm1xxBlkScaledConfig::tile_atom_to_shape_SFB(problem);
164
+
165
+ using ArrayElementA = typename Gemm::GemmKernel::CollectiveMainloop::ArrayElementA;
166
+ using ArrayElementB = typename Gemm::GemmKernel::CollectiveMainloop::ArrayElementB;
167
+
168
+ typename Gemm::Arguments args{
169
+ cutlass::gemm::GemmUniversalMode::kGemm,
170
+ {M, N, K, 1},
171
+ {
172
+ reinterpret_cast<ArrayElementA const*>(A_packed), strA,
173
+ reinterpret_cast<ArrayElementB const*>(B_packed), strB,
174
+ reinterpret_cast<ElementSF const*>(SFA), layout_SFA,
175
+ reinterpret_cast<ElementSF const*>(SFB), layout_SFB
176
+ },
177
+ {
178
+ {alpha, 0.0f},
179
+ nullptr, strC,
180
+ reinterpret_cast<ElementD*>(D_bf16), strD
181
+ }
182
+ };
183
+ args.epilogue.thread.bias_ptr =
184
+ reinterpret_cast<ElementBias const*>(bias_bf16);
185
+
186
+ Gemm gemm;
187
+ size_t ws_size = Gemm::get_workspace_size(args);
188
+ void* ws_ptr = get_ws(M, N, K, ws_size);
189
+ auto status = gemm.can_implement(args);
190
+ if (status != cutlass::Status::kSuccess) {
191
+ std::fprintf(stderr,
192
+ "[fp4_w4a16_gemm_bias_gelu_bf16out_sm120] can_implement FAIL status=%d\n",
193
+ int(status));
194
+ return;
195
+ }
196
+ status = gemm.initialize(args, ws_ptr, stream);
197
+ if (status != cutlass::Status::kSuccess) {
198
+ std::fprintf(stderr,
199
+ "[fp4_w4a16_gemm_bias_gelu_bf16out_sm120] initialize FAIL status=%d\n",
200
+ int(status));
201
+ return;
202
+ }
203
+ status = gemm.run(stream);
204
+ if (status != cutlass::Status::kSuccess) {
205
+ std::fprintf(stderr,
206
+ "[fp4_w4a16_gemm_bias_gelu_bf16out_sm120] run FAIL status=%d\n",
207
+ int(status));
208
+ }
209
+ }
210
+
211
+ } // namespace gemm
212
+ } // namespace flash_rt
csrc/gemm/fp4/cutlass_nvfp4_gemm_bias_gelu_bf16out_sm120.cuh ADDED
@@ -0,0 +1,37 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ // SPDX-License-Identifier: Apache-2.0
2
+ //
3
+ // CUTLASS NVFP4 W4A16 GEMM with fused per-col bias + GELU(tanh) epilogue,
4
+ // BF16 output, SM120a. Header for csrc/gemm/fp4/cutlass_nvfp4_gemm_bias_
5
+ // gelu_bf16out_sm120.cu (Recipe C step 1).
6
+
7
+ #pragma once
8
+
9
+ #include <cuda_runtime.h>
10
+
11
+ namespace flash_rt {
12
+ namespace gemm {
13
+
14
+ // D = GELU_tanh(alpha * (A_fp4 @ B_fp4^T scaled by SFA/SFB) + bias_per_col)
15
+ //
16
+ // A_packed : (M, K/2) uint8 NVFP4 packed (cutlass-swizzled)
17
+ // B_packed : (N, K/2) uint8 NVFP4 packed (cutlass-swizzled)
18
+ // SFA : (M*K/16) e4m3 NVFP4 SF for A
19
+ // SFB : (N*K/16) e4m3 NVFP4 SF for B
20
+ // bias_bf16: (N,) bf16 per-col bias (added before GELU)
21
+ // D_bf16 : (M, N) bf16 output
22
+ // alpha : float32 = sf_global_a * sf_global_b
23
+ //
24
+ // Stream-safe; per-shape workspace cached internally.
25
+ void fp4_w4a16_gemm_bias_gelu_bf16out_sm120(
26
+ const void* A_packed,
27
+ const void* B_packed,
28
+ const void* SFA,
29
+ const void* SFB,
30
+ const void* bias_bf16,
31
+ void* D_bf16,
32
+ int M, int N, int K,
33
+ float alpha,
34
+ cudaStream_t stream);
35
+
36
+ } // namespace gemm
37
+ } // namespace flash_rt
csrc/gemm/fp4/cutlass_nvfp4_gemm_bias_gelu_fp4out_sm120.cu ADDED
@@ -0,0 +1,234 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ // SPDX-License-Identifier: Apache-2.0
2
+ //
3
+ // CUTLASS NVFP4 W4A16 GEMM with fused per-col bias + GELU(tanh) +
4
+ // per-block-16 NVFP4 quantization epilogue, FP4 packed output, SM120a.
5
+ //
6
+ // Replaces the 3-launch chain
7
+ // cutlass NVFP4 GEMM_up (~33 µs)
8
+ // + bias_gelu_inplace_bf16 (~8 µs)
9
+ // + quantize_bf16_to_nvfp4 (~20 µs)
10
+ // with a single cutlass-fork kernel (~33 µs).
11
+ //
12
+ // Same TileShape <128,128,256> ClusterShape <1,1,1>
13
+ // KernelTmaWarpSpecializedPingpong as the bf16-out fork; the
14
+ // FusionOperation is swapped to LinCombPerColBiasEltActBlockScaleFactor
15
+ // which produces packed NVFP4 + UE4M3 SF in cutlass-swizzled layout
16
+ // (consumable directly by the downstream NVFP4 W4A16 GEMM_dn that reads
17
+ // the same SF layout).
18
+
19
+ #include "cutlass_nvfp4_gemm_bias_gelu_fp4out_sm120.cuh"
20
+
21
+ #include "cute/tensor.hpp"
22
+
23
+ #include "cutlass/cutlass.h"
24
+ #include "cutlass/numeric_types.h"
25
+ #include "cutlass/detail/sm100_blockscaled_layout.hpp"
26
+
27
+ #include "cutlass/epilogue/collective/collective_builder.hpp"
28
+ #include "cutlass/epilogue/thread/activation.h"
29
+ #include "cutlass/epilogue/fusion/operations.hpp"
30
+
31
+ #include "cutlass/gemm/collective/collective_builder.hpp"
32
+ #include "cutlass/gemm/dispatch_policy.hpp"
33
+ #include "cutlass/gemm/device/gemm_universal_adapter.h"
34
+ #include "cutlass/gemm/kernel/gemm_universal.hpp"
35
+
36
+ #include "cutlass/util/packed_stride.hpp"
37
+
38
+ #include <cstdio>
39
+ #include <mutex>
40
+ #include <unordered_map>
41
+
42
+ namespace flash_rt {
43
+ namespace gemm {
44
+
45
+ namespace {
46
+ using namespace cute;
47
+
48
+ using ElementA = cutlass::float_e2m1_t;
49
+ using ElementB = cutlass::float_e2m1_t;
50
+ using ElementC = cutlass::bfloat16_t;
51
+ using ElementD = cutlass::float_e2m1_t;
52
+ using ElementSFD = cutlass::float_ue4m3_t;
53
+ using ElementBias = cutlass::bfloat16_t;
54
+ using ElementAccumulator = float;
55
+ using ElementCompute = float;
56
+ using ElementSF = cutlass::float_ue4m3_t;
57
+
58
+ using LayoutA = cutlass::layout::RowMajor;
59
+ using LayoutB = cutlass::layout::ColumnMajor;
60
+ using LayoutC = cutlass::layout::ColumnMajor;
61
+ using LayoutD = cutlass::layout::RowMajor;
62
+ using LayoutSFDTag = cutlass::layout::RowMajor;
63
+
64
+ using ElementPairA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
65
+ using ElementPairB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
66
+
67
+ constexpr int AlignmentA = 16 * 8 / cutlass::sizeof_bits<ElementA>::value;
68
+ constexpr int AlignmentB = 16 * 8 / cutlass::sizeof_bits<ElementB>::value;
69
+ constexpr int AlignmentC = 128 / cutlass::sizeof_bits<ElementC>::value;
70
+ constexpr int AlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
71
+
72
+ using TileShape = Shape<_128, _128, _256>;
73
+ using ClusterShape = Shape<_1, _1, _1>;
74
+ using Sm1xxBlkScaledConfig = cutlass::detail::Sm1xxBlockScaledConfig<16>;
75
+
76
+ constexpr int OutputSFVectorSize = 16;
77
+
78
+ using FusionOperation = cutlass::epilogue::fusion::LinCombPerColBiasEltActBlockScaleFactor<
79
+ cutlass::epilogue::thread::GELU_taylor,
80
+ OutputSFVectorSize,
81
+ ElementD,
82
+ ElementCompute,
83
+ ElementSFD,
84
+ LayoutSFDTag,
85
+ ElementBias>;
86
+
87
+ using CollectiveEpilogue =
88
+ typename cutlass::epilogue::collective::CollectiveBuilder<
89
+ cutlass::arch::Sm120, cutlass::arch::OpClassTensorOp,
90
+ TileShape, ClusterShape,
91
+ cutlass::epilogue::collective::EpilogueTileAuto,
92
+ ElementAccumulator, ElementCompute,
93
+ ElementC, LayoutC, AlignmentC,
94
+ ElementD, LayoutD, AlignmentD,
95
+ cutlass::epilogue::collective::EpilogueScheduleAuto,
96
+ FusionOperation
97
+ >::CollectiveOp;
98
+
99
+ using CollectiveMainloop =
100
+ typename cutlass::gemm::collective::CollectiveBuilder<
101
+ cutlass::arch::Sm120, cutlass::arch::OpClassBlockScaledTensorOp,
102
+ ElementPairA, LayoutA, AlignmentA,
103
+ ElementPairB, LayoutB, AlignmentB,
104
+ ElementAccumulator,
105
+ TileShape, ClusterShape,
106
+ cutlass::gemm::collective::StageCountAutoCarveout<
107
+ static_cast<int>(sizeof(typename CollectiveEpilogue::SharedStorage))>,
108
+ cutlass::gemm::KernelTmaWarpSpecializedPingpong
109
+ >::CollectiveOp;
110
+
111
+ using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
112
+ Shape<int, int, int, int>, CollectiveMainloop, CollectiveEpilogue,
113
+ cutlass::gemm::PersistentScheduler>;
114
+
115
+ using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
116
+
117
+ using SfdOutputCfg = cutlass::detail::Sm1xxBlockScaledOutputConfig<OutputSFVectorSize>;
118
+
119
+ struct ShapeKey {
120
+ int M, N, K;
121
+ bool operator==(const ShapeKey& o) const {
122
+ return M == o.M && N == o.N && K == o.K;
123
+ }
124
+ };
125
+ struct SHash {
126
+ size_t operator()(const ShapeKey& k) const noexcept {
127
+ return (size_t(k.M) * 1315423911u) ^ (size_t(k.N) * 2654435761u)
128
+ ^ size_t(k.K);
129
+ }
130
+ };
131
+ struct CachedWs { void* ptr = nullptr; size_t size = 0; };
132
+ std::unordered_map<ShapeKey, CachedWs, SHash> g_ws;
133
+ std::mutex g_mu;
134
+
135
+ void* get_ws(int M, int N, int K, size_t need) {
136
+ std::lock_guard<std::mutex> lk(g_mu);
137
+ ShapeKey k{M, N, K};
138
+ auto it = g_ws.find(k);
139
+ if (it != g_ws.end() && it->second.size >= need) return it->second.ptr;
140
+ if (it != g_ws.end()) { cudaFree(it->second.ptr); g_ws.erase(it); }
141
+ CachedWs w; w.size = need;
142
+ if (need > 0) cudaMalloc(&w.ptr, need);
143
+ g_ws[k] = w;
144
+ return w.ptr;
145
+ }
146
+
147
+ float* get_norm_const_one() {
148
+ static float* p = nullptr;
149
+ if (p == nullptr) {
150
+ cudaMalloc(&p, sizeof(float));
151
+ float one = 1.0f;
152
+ cudaMemcpy(p, &one, sizeof(float), cudaMemcpyHostToDevice);
153
+ }
154
+ return p;
155
+ }
156
+
157
+ } // namespace
158
+
159
+ void fp4_w4a16_gemm_bias_gelu_fp4out_sm120(
160
+ const void* A_packed, const void* B_packed,
161
+ const void* SFA, const void* SFB,
162
+ const void* bias_bf16,
163
+ void* D_packed,
164
+ void* SFD,
165
+ int M, int N, int K,
166
+ float alpha,
167
+ cudaStream_t stream)
168
+ {
169
+ using StrideA = typename Gemm::GemmKernel::StrideA;
170
+ using StrideB = typename Gemm::GemmKernel::StrideB;
171
+ using StrideC = typename Gemm::GemmKernel::StrideC;
172
+ using StrideD = typename Gemm::GemmKernel::StrideD;
173
+ StrideA strA = cutlass::make_cute_packed_stride(StrideA{}, cute::make_shape(M, K, 1));
174
+ StrideB strB = cutlass::make_cute_packed_stride(StrideB{}, cute::make_shape(N, K, 1));
175
+ StrideC strC = cutlass::make_cute_packed_stride(StrideC{}, cute::make_shape(M, N, 1));
176
+ StrideD strD = cutlass::make_cute_packed_stride(StrideD{}, cute::make_shape(M, N, 1));
177
+
178
+ auto problem_MNKL = cute::make_shape(M, N, K, 1);
179
+ auto layout_SFA = Sm1xxBlkScaledConfig::tile_atom_to_shape_SFA(problem_MNKL);
180
+ auto layout_SFB = Sm1xxBlkScaledConfig::tile_atom_to_shape_SFB(problem_MNKL);
181
+
182
+ using ArrayElementA = typename Gemm::GemmKernel::CollectiveMainloop::ArrayElementA;
183
+ using ArrayElementB = typename Gemm::GemmKernel::CollectiveMainloop::ArrayElementB;
184
+
185
+ float* norm_const_dev = get_norm_const_one();
186
+
187
+ typename Gemm::Arguments args{
188
+ cutlass::gemm::GemmUniversalMode::kGemm,
189
+ {M, N, K, 1},
190
+ {
191
+ reinterpret_cast<ArrayElementA const*>(A_packed), strA,
192
+ reinterpret_cast<ArrayElementB const*>(B_packed), strB,
193
+ reinterpret_cast<ElementSF const*>(SFA), layout_SFA,
194
+ reinterpret_cast<ElementSF const*>(SFB), layout_SFB
195
+ },
196
+ {
197
+ {alpha, 0.0f},
198
+ nullptr, strC,
199
+ reinterpret_cast<ElementD*>(D_packed), strD
200
+ }
201
+ };
202
+ args.epilogue.thread.bias_ptr =
203
+ reinterpret_cast<ElementBias const*>(bias_bf16);
204
+ args.epilogue.thread.block_scale_factor_ptr =
205
+ reinterpret_cast<ElementSFD*>(SFD);
206
+ args.epilogue.thread.norm_constant_ptr = norm_const_dev;
207
+
208
+ Gemm gemm;
209
+ size_t ws_size = Gemm::get_workspace_size(args);
210
+ void* ws_ptr = get_ws(M, N, K, ws_size);
211
+ auto status = gemm.can_implement(args);
212
+ if (status != cutlass::Status::kSuccess) {
213
+ std::fprintf(stderr,
214
+ "[fp4_w4a16_gemm_bias_gelu_fp4out_sm120] can_implement FAIL "
215
+ "M=%d N=%d K=%d status=%d\n", M, N, K, int(status));
216
+ return;
217
+ }
218
+ status = gemm.initialize(args, ws_ptr, stream);
219
+ if (status != cutlass::Status::kSuccess) {
220
+ std::fprintf(stderr,
221
+ "[fp4_w4a16_gemm_bias_gelu_fp4out_sm120] initialize FAIL "
222
+ "M=%d N=%d K=%d status=%d\n", M, N, K, int(status));
223
+ return;
224
+ }
225
+ status = gemm.run(stream);
226
+ if (status != cutlass::Status::kSuccess) {
227
+ std::fprintf(stderr,
228
+ "[fp4_w4a16_gemm_bias_gelu_fp4out_sm120] run FAIL status=%d\n",
229
+ int(status));
230
+ }
231
+ }
232
+
233
+ } // namespace gemm
234
+ } // namespace flash_rt
csrc/gemm/fp4/cutlass_nvfp4_gemm_bias_gelu_fp4out_sm120.cuh ADDED
@@ -0,0 +1,40 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ // SPDX-License-Identifier: Apache-2.0
2
+ //
3
+ // CUTLASS NVFP4 W4A16 GEMM with fused per-col bias + GELU(tanh) +
4
+ // per-block-16 NVFP4 quantization epilogue, FP4 packed output, SM120a.
5
+
6
+ #pragma once
7
+
8
+ #include <cuda_runtime.h>
9
+
10
+ namespace flash_rt {
11
+ namespace gemm {
12
+
13
+ // D_packed[m, n/2] = pack_FP4(
14
+ // GELU_tanh(alpha * (A_fp4 @ B_fp4^T scaled by SFA/SFB) + bias_per_col)
15
+ // ) with per-16-block NVFP4 SFD in cutlass-swizzled UE4M3 layout.
16
+ //
17
+ // A_packed : (M, K/2) uint8 NVFP4 packed (cutlass-swizzled SF)
18
+ // B_packed : (N, K/2) uint8 NVFP4 packed (cutlass-swizzled SF)
19
+ // SFA : (M*K/16) e4m3
20
+ // SFB : (N*K/16) e4m3
21
+ // bias_bf16: (N,) bf16 per-col bias
22
+ // D_packed : (M, N/2) uint8 NVFP4 packed
23
+ // SFD : (M*N/16) e4m3 output SF, cutlass-swizzled layout
24
+ // alpha : float32 = sf_global_a * sf_global_b
25
+ //
26
+ // Stream-safe; per-shape workspace cached internally.
27
+ void fp4_w4a16_gemm_bias_gelu_fp4out_sm120(
28
+ const void* A_packed,
29
+ const void* B_packed,
30
+ const void* SFA,
31
+ const void* SFB,
32
+ const void* bias_bf16,
33
+ void* D_packed,
34
+ void* SFD,
35
+ int M, int N, int K,
36
+ float alpha,
37
+ cudaStream_t stream);
38
+
39
+ } // namespace gemm
40
+ } // namespace flash_rt
csrc/gemm/fp4/cutlass_nvfp4_gemm_dn_streamk_bias_sm120.cu ADDED
@@ -0,0 +1,307 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ // SPDX-License-Identifier: Apache-2.0
2
+ //
3
+ // CUTLASS NVFP4 W4A16 GEMM with fused per-col bias epilogue, BF16 output,
4
+ // **StreamK scheduler**, SM120a.
5
+ //
6
+ // fvk's default fp4_w4a16_gemm_sm120_bf16out uses
7
+ // KernelTmaWarpSpecializedCooperative + PersistentScheduler. At motus
8
+ // Wan FFN GEMM_dn shape (M=360, N=K_motus=3072, K=F=14336), TileShape
9
+ // <128,128,256> yields 3 × 24 = 72 CTAs on 170 SMs = 0.42 wave, leaving
10
+ // the GPU under-utilized. StreamK partitions the K-axis to issue more
11
+ // CTAs and converges to ~1.7 waves, recovering 1.277× speedup standalone
12
+ // (47 µs → 37 µs).
13
+ //
14
+ // Per-col bias is absorbed into the epilogue via LinCombPerColBiasEltAct
15
+ // with Identity activation — eliminates the separate add_bias_bf16 launch.
16
+
17
+ #include "cutlass_nvfp4_gemm_dn_streamk_bias_sm120.cuh"
18
+
19
+ #include "cute/tensor.hpp"
20
+
21
+ #include "cutlass/cutlass.h"
22
+ #include "cutlass/numeric_types.h"
23
+ #include "cutlass/detail/sm100_blockscaled_layout.hpp"
24
+
25
+ #include "cutlass/epilogue/collective/collective_builder.hpp"
26
+ #include "cutlass/epilogue/thread/activation.h"
27
+ #include "cutlass/epilogue/thread/linear_combination.h"
28
+ #include "cutlass/epilogue/fusion/operations.hpp"
29
+
30
+ #include "cutlass/gemm/collective/collective_builder.hpp"
31
+ #include "cutlass/gemm/dispatch_policy.hpp"
32
+ #include "cutlass/gemm/device/gemm_universal_adapter.h"
33
+ #include "cutlass/gemm/kernel/gemm_universal.hpp"
34
+
35
+ #include "cutlass/util/packed_stride.hpp"
36
+
37
+ #include <cstdio>
38
+ #include <mutex>
39
+ #include <unordered_map>
40
+
41
+ namespace flash_rt {
42
+ namespace gemm {
43
+
44
+ namespace {
45
+ using namespace cute;
46
+
47
+ using ElementA = cutlass::float_e2m1_t;
48
+ using ElementB = cutlass::float_e2m1_t;
49
+ using ElementC = cutlass::bfloat16_t;
50
+ using ElementD = cutlass::bfloat16_t;
51
+ using ElementBias = cutlass::bfloat16_t;
52
+ using ElementAccumulator = float;
53
+ using ElementCompute = float;
54
+ using ElementSF = cutlass::float_ue4m3_t;
55
+
56
+ using LayoutA = cutlass::layout::RowMajor;
57
+ using LayoutB = cutlass::layout::ColumnMajor;
58
+ using LayoutC = cutlass::layout::RowMajor;
59
+ using LayoutD = cutlass::layout::RowMajor;
60
+
61
+ using ElementPairA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
62
+ using ElementPairB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
63
+
64
+ constexpr int AlignmentA = 16 * 8 / cutlass::sizeof_bits<ElementA>::value;
65
+ constexpr int AlignmentB = 16 * 8 / cutlass::sizeof_bits<ElementB>::value;
66
+ constexpr int AlignmentC = 128 / cutlass::sizeof_bits<ElementC>::value;
67
+ constexpr int AlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
68
+
69
+ using TileShape = Shape<_128, _128, _256>;
70
+ using ClusterShape = Shape<_1, _1, _1>;
71
+ using Sm1xxBlkScaledConfig = cutlass::detail::Sm1xxBlockScaledConfig<16>;
72
+
73
+ using FusionOperation = cutlass::epilogue::fusion::LinCombPerColBiasEltAct<
74
+ cutlass::epilogue::thread::Identity,
75
+ ElementD, ElementCompute, ElementBias, ElementC>;
76
+
77
+ using CollectiveEpilogue =
78
+ typename cutlass::epilogue::collective::CollectiveBuilder<
79
+ cutlass::arch::Sm120, cutlass::arch::OpClassTensorOp,
80
+ TileShape, ClusterShape,
81
+ cutlass::epilogue::collective::EpilogueTileAuto,
82
+ ElementAccumulator, ElementCompute,
83
+ ElementC, LayoutC, AlignmentC,
84
+ ElementD, LayoutD, AlignmentD,
85
+ cutlass::epilogue::collective::EpilogueScheduleAuto,
86
+ FusionOperation
87
+ >::CollectiveOp;
88
+
89
+ using CollectiveMainloop =
90
+ typename cutlass::gemm::collective::CollectiveBuilder<
91
+ cutlass::arch::Sm120, cutlass::arch::OpClassBlockScaledTensorOp,
92
+ ElementPairA, LayoutA, AlignmentA,
93
+ ElementPairB, LayoutB, AlignmentB,
94
+ ElementAccumulator,
95
+ TileShape, ClusterShape,
96
+ cutlass::gemm::collective::StageCountAutoCarveout<
97
+ static_cast<int>(sizeof(typename CollectiveEpilogue::SharedStorage))>,
98
+ cutlass::gemm::KernelTmaWarpSpecializedCooperative
99
+ >::CollectiveOp;
100
+
101
+ using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
102
+ Shape<int, int, int, int>, CollectiveMainloop, CollectiveEpilogue,
103
+ cutlass::gemm::StreamKScheduler>;
104
+
105
+ using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
106
+
107
+ using NoBiasFusionOperation = cutlass::epilogue::fusion::LinearCombination<
108
+ ElementD, ElementCompute, ElementC, ElementCompute>;
109
+
110
+ using NoBiasCollectiveEpilogue =
111
+ typename cutlass::epilogue::collective::CollectiveBuilder<
112
+ cutlass::arch::Sm120, cutlass::arch::OpClassTensorOp,
113
+ TileShape, ClusterShape,
114
+ cutlass::epilogue::collective::EpilogueTileAuto,
115
+ ElementAccumulator, ElementCompute,
116
+ ElementC, LayoutC, AlignmentC,
117
+ ElementD, LayoutD, AlignmentD,
118
+ cutlass::epilogue::collective::EpilogueScheduleAuto,
119
+ NoBiasFusionOperation
120
+ >::CollectiveOp;
121
+
122
+ using NoBiasCollectiveMainloop =
123
+ typename cutlass::gemm::collective::CollectiveBuilder<
124
+ cutlass::arch::Sm120, cutlass::arch::OpClassBlockScaledTensorOp,
125
+ ElementPairA, LayoutA, AlignmentA,
126
+ ElementPairB, LayoutB, AlignmentB,
127
+ ElementAccumulator,
128
+ TileShape, ClusterShape,
129
+ cutlass::gemm::collective::StageCountAutoCarveout<
130
+ static_cast<int>(sizeof(typename NoBiasCollectiveEpilogue::SharedStorage))>,
131
+ cutlass::gemm::KernelTmaWarpSpecializedCooperative
132
+ >::CollectiveOp;
133
+
134
+ using NoBiasGemmKernel = cutlass::gemm::kernel::GemmUniversal<
135
+ Shape<int, int, int, int>, NoBiasCollectiveMainloop,
136
+ NoBiasCollectiveEpilogue, cutlass::gemm::StreamKScheduler>;
137
+
138
+ using NoBiasGemm =
139
+ cutlass::gemm::device::GemmUniversalAdapter<NoBiasGemmKernel>;
140
+
141
+ struct ShapeKey {
142
+ int M, N, K;
143
+ bool operator==(const ShapeKey& o) const {
144
+ return M == o.M && N == o.N && K == o.K;
145
+ }
146
+ };
147
+ struct SHash {
148
+ size_t operator()(const ShapeKey& k) const noexcept {
149
+ return (size_t(k.M) * 1315423911u) ^ (size_t(k.N) * 2654435761u)
150
+ ^ size_t(k.K);
151
+ }
152
+ };
153
+ struct CachedWs { void* ptr = nullptr; size_t size = 0; };
154
+ std::unordered_map<ShapeKey, CachedWs, SHash> g_ws;
155
+ std::mutex g_mu;
156
+
157
+ void* get_ws(int M, int N, int K, size_t need) {
158
+ std::lock_guard<std::mutex> lk(g_mu);
159
+ ShapeKey k{M, N, K};
160
+ auto it = g_ws.find(k);
161
+ if (it != g_ws.end() && it->second.size >= need) return it->second.ptr;
162
+ if (it != g_ws.end()) { cudaFree(it->second.ptr); g_ws.erase(it); }
163
+ CachedWs w; w.size = need;
164
+ if (need > 0) cudaMalloc(&w.ptr, need);
165
+ g_ws[k] = w;
166
+ return w.ptr;
167
+ }
168
+
169
+ } // namespace
170
+
171
+ void fp4_w4a16_gemm_dn_streamk_bf16out_sm120(
172
+ const void* A_packed, const void* B_packed,
173
+ const void* SFA, const void* SFB,
174
+ void* D_bf16,
175
+ int M, int N, int K,
176
+ float alpha,
177
+ cudaStream_t stream)
178
+ {
179
+ using StrideA = typename NoBiasGemm::GemmKernel::StrideA;
180
+ using StrideB = typename NoBiasGemm::GemmKernel::StrideB;
181
+ using StrideC = typename NoBiasGemm::GemmKernel::StrideC;
182
+ using StrideD = typename NoBiasGemm::GemmKernel::StrideD;
183
+ StrideA strA = cutlass::make_cute_packed_stride(StrideA{}, cute::make_shape(M, K, 1));
184
+ StrideB strB = cutlass::make_cute_packed_stride(StrideB{}, cute::make_shape(N, K, 1));
185
+ StrideC strC = cutlass::make_cute_packed_stride(StrideC{}, cute::make_shape(M, N, 1));
186
+ StrideD strD = cutlass::make_cute_packed_stride(StrideD{}, cute::make_shape(M, N, 1));
187
+
188
+ auto problem = cute::make_shape(M, N, K, 1);
189
+ auto layout_SFA = Sm1xxBlkScaledConfig::tile_atom_to_shape_SFA(problem);
190
+ auto layout_SFB = Sm1xxBlkScaledConfig::tile_atom_to_shape_SFB(problem);
191
+
192
+ using ArrayElementA =
193
+ typename NoBiasGemm::GemmKernel::CollectiveMainloop::ArrayElementA;
194
+ using ArrayElementB =
195
+ typename NoBiasGemm::GemmKernel::CollectiveMainloop::ArrayElementB;
196
+
197
+ typename NoBiasGemm::Arguments args{
198
+ cutlass::gemm::GemmUniversalMode::kGemm,
199
+ {M, N, K, 1},
200
+ {
201
+ reinterpret_cast<ArrayElementA const*>(A_packed), strA,
202
+ reinterpret_cast<ArrayElementB const*>(B_packed), strB,
203
+ reinterpret_cast<ElementSF const*>(SFA), layout_SFA,
204
+ reinterpret_cast<ElementSF const*>(SFB), layout_SFB
205
+ },
206
+ {
207
+ {alpha, 0.0f},
208
+ nullptr, strC,
209
+ reinterpret_cast<ElementD*>(D_bf16), strD
210
+ }
211
+ };
212
+
213
+ NoBiasGemm gemm;
214
+ size_t ws_size = NoBiasGemm::get_workspace_size(args);
215
+ void* ws_ptr = get_ws(M, N, K, ws_size);
216
+ auto status = gemm.can_implement(args);
217
+ if (status != cutlass::Status::kSuccess) {
218
+ std::fprintf(stderr,
219
+ "[fp4_w4a16_gemm_dn_streamk_bf16out_sm120] can_implement FAIL "
220
+ "M=%d N=%d K=%d status=%d\n", M, N, K, int(status));
221
+ return;
222
+ }
223
+ status = gemm.initialize(args, ws_ptr, stream);
224
+ if (status != cutlass::Status::kSuccess) {
225
+ std::fprintf(stderr,
226
+ "[fp4_w4a16_gemm_dn_streamk_bf16out_sm120] initialize FAIL "
227
+ "M=%d N=%d K=%d status=%d\n", M, N, K, int(status));
228
+ return;
229
+ }
230
+ status = gemm.run(stream);
231
+ if (status != cutlass::Status::kSuccess) {
232
+ std::fprintf(stderr,
233
+ "[fp4_w4a16_gemm_dn_streamk_bf16out_sm120] run FAIL status=%d\n",
234
+ int(status));
235
+ }
236
+ }
237
+
238
+ void fp4_w4a16_gemm_dn_streamk_bias_bf16out_sm120(
239
+ const void* A_packed, const void* B_packed,
240
+ const void* SFA, const void* SFB,
241
+ const void* bias_bf16,
242
+ void* D_bf16,
243
+ int M, int N, int K,
244
+ float alpha,
245
+ cudaStream_t stream)
246
+ {
247
+ using StrideA = typename Gemm::GemmKernel::StrideA;
248
+ using StrideB = typename Gemm::GemmKernel::StrideB;
249
+ using StrideC = typename Gemm::GemmKernel::StrideC;
250
+ using StrideD = typename Gemm::GemmKernel::StrideD;
251
+ StrideA strA = cutlass::make_cute_packed_stride(StrideA{}, cute::make_shape(M, K, 1));
252
+ StrideB strB = cutlass::make_cute_packed_stride(StrideB{}, cute::make_shape(N, K, 1));
253
+ StrideC strC = cutlass::make_cute_packed_stride(StrideC{}, cute::make_shape(M, N, 1));
254
+ StrideD strD = cutlass::make_cute_packed_stride(StrideD{}, cute::make_shape(M, N, 1));
255
+
256
+ auto problem = cute::make_shape(M, N, K, 1);
257
+ auto layout_SFA = Sm1xxBlkScaledConfig::tile_atom_to_shape_SFA(problem);
258
+ auto layout_SFB = Sm1xxBlkScaledConfig::tile_atom_to_shape_SFB(problem);
259
+
260
+ using ArrayElementA = typename Gemm::GemmKernel::CollectiveMainloop::ArrayElementA;
261
+ using ArrayElementB = typename Gemm::GemmKernel::CollectiveMainloop::ArrayElementB;
262
+
263
+ typename Gemm::Arguments args{
264
+ cutlass::gemm::GemmUniversalMode::kGemm,
265
+ {M, N, K, 1},
266
+ {
267
+ reinterpret_cast<ArrayElementA const*>(A_packed), strA,
268
+ reinterpret_cast<ArrayElementB const*>(B_packed), strB,
269
+ reinterpret_cast<ElementSF const*>(SFA), layout_SFA,
270
+ reinterpret_cast<ElementSF const*>(SFB), layout_SFB
271
+ },
272
+ {
273
+ {alpha, 0.0f},
274
+ nullptr, strC,
275
+ reinterpret_cast<ElementD*>(D_bf16), strD
276
+ }
277
+ };
278
+ args.epilogue.thread.bias_ptr =
279
+ reinterpret_cast<ElementBias const*>(bias_bf16);
280
+
281
+ Gemm gemm;
282
+ size_t ws_size = Gemm::get_workspace_size(args);
283
+ void* ws_ptr = get_ws(M, N, K, ws_size);
284
+ auto status = gemm.can_implement(args);
285
+ if (status != cutlass::Status::kSuccess) {
286
+ std::fprintf(stderr,
287
+ "[fp4_w4a16_gemm_dn_streamk_bias_bf16out_sm120] can_implement FAIL "
288
+ "M=%d N=%d K=%d status=%d\n", M, N, K, int(status));
289
+ return;
290
+ }
291
+ status = gemm.initialize(args, ws_ptr, stream);
292
+ if (status != cutlass::Status::kSuccess) {
293
+ std::fprintf(stderr,
294
+ "[fp4_w4a16_gemm_dn_streamk_bias_bf16out_sm120] initialize FAIL "
295
+ "M=%d N=%d K=%d status=%d\n", M, N, K, int(status));
296
+ return;
297
+ }
298
+ status = gemm.run(stream);
299
+ if (status != cutlass::Status::kSuccess) {
300
+ std::fprintf(stderr,
301
+ "[fp4_w4a16_gemm_dn_streamk_bias_bf16out_sm120] run FAIL status=%d\n",
302
+ int(status));
303
+ }
304
+ }
305
+
306
+ } // namespace gemm
307
+ } // namespace flash_rt
csrc/gemm/fp4/cutlass_nvfp4_gemm_dn_streamk_bias_sm120.cuh ADDED
@@ -0,0 +1,53 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ // SPDX-License-Identifier: Apache-2.0
2
+ //
3
+ // CUTLASS NVFP4 W4A16 GEMM_dn with fused per-col bias epilogue, BF16
4
+ // output, **StreamK scheduler**, SM120a.
5
+
6
+ #pragma once
7
+
8
+ #include <cuda_runtime.h>
9
+
10
+ namespace flash_rt {
11
+ namespace gemm {
12
+
13
+ // D = (alpha * (A_fp4 @ B_fp4^T scaled by SFA/SFB) + per_col_bias) → bf16
14
+ //
15
+ // A_packed : (M, K/2) uint8 NVFP4 packed (cutlass-swizzled SF)
16
+ // B_packed : (N, K/2) uint8 NVFP4 packed
17
+ // SFA : (M*K/16) e4m3
18
+ // SFB : (N*K/16) e4m3
19
+ // bias_bf16: (N,) bf16 per-col bias added in epilogue
20
+ // D_bf16 : (M, N) bf16 output
21
+ // alpha : float32 = sf_global_a * sf_global_b
22
+ //
23
+ // Stream-safe; per-shape workspace cached internally. Uses
24
+ // StreamKScheduler to recover SM utilization at the motus Wan FFN
25
+ // GEMM_dn shape (M=360, N=3072, K=14336): 1.277× over default
26
+ // PersistentScheduler standalone.
27
+ void fp4_w4a16_gemm_dn_streamk_bias_bf16out_sm120(
28
+ const void* A_packed,
29
+ const void* B_packed,
30
+ const void* SFA,
31
+ const void* SFB,
32
+ const void* bias_bf16,
33
+ void* D_bf16,
34
+ int M, int N, int K,
35
+ float alpha,
36
+ cudaStream_t stream);
37
+
38
+ // D = alpha * (A_fp4 @ B_fp4^T scaled by SFA/SFB) -> bf16
39
+ //
40
+ // Same StreamK schedule as the bias variant, but with a pure linear-combine
41
+ // epilogue. This matches Motus down-only sites whose down bias is skipped.
42
+ void fp4_w4a16_gemm_dn_streamk_bf16out_sm120(
43
+ const void* A_packed,
44
+ const void* B_packed,
45
+ const void* SFA,
46
+ const void* SFB,
47
+ void* D_bf16,
48
+ int M, int N, int K,
49
+ float alpha,
50
+ cudaStream_t stream);
51
+
52
+ } // namespace gemm
53
+ } // namespace flash_rt
csrc/gemm/fp4/cutlass_nvfp4_w4a16_gemm_sm100.cu ADDED
@@ -0,0 +1,411 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ // SPDX-License-Identifier: Apache-2.0
2
+ //
3
+ // Sm100 NVFP4 W4A16 block-scaled GEMM. BF16 output.
4
+ //
5
+ // Header: cutlass_nvfp4_w4a16_gemm_sm100.cuh.
6
+ //
7
+ // Template structure is a translation of the verified Sm120 path
8
+ // (cutlass_nvfp4_w4a16_gemm_sm120.cu) onto the SM100 dispatch:
9
+ // - arch::Sm120 -> arch::Sm100
10
+ // - KernelTmaWarpSpecializedCooperative -> KernelScheduleAuto
11
+ // - KernelTmaWarpSpecializedPingpong -> KernelScheduleAuto
12
+ // All other types and layouts (FP4 e2m1 A/B, ue4m3 group scales,
13
+ // row-major D in bf16, group_size=16) match the Sm120 variant byte
14
+ // for byte. The wire-format contract (activation quantizer SFA layout,
15
+ // loader SFB layout, alpha = sf_global_a * sf_global_b) is identical
16
+ // so the Qwen3.6 frontend re-uses the same calls.
17
+ //
18
+ // Built only when GPU_ARCH==110 (Thor). The Sm100 dispatch reaches the
19
+ // correct sm_110a tcgen05 mainloop without any per-arch macro.
20
+
21
+ #include "cutlass_nvfp4_w4a16_gemm_sm100.cuh"
22
+
23
+ #include "cute/tensor.hpp"
24
+ #include "cute/atom/mma_atom.hpp"
25
+
26
+ #include "cutlass/cutlass.h"
27
+ #include "cutlass/numeric_types.h"
28
+ #include "cutlass/detail/sm100_blockscaled_layout.hpp"
29
+
30
+ #include "cutlass/epilogue/collective/collective_builder.hpp"
31
+ #include "cutlass/epilogue/collective/default_epilogue.hpp"
32
+ #include "cutlass/epilogue/thread/linear_combination.h"
33
+
34
+ #include "cutlass/gemm/collective/collective_builder.hpp"
35
+ #include "cutlass/gemm/dispatch_policy.hpp"
36
+ #include "cutlass/gemm/device/gemm_universal_adapter.h"
37
+ #include "cutlass/gemm/kernel/gemm_universal.hpp"
38
+
39
+ #include "cutlass/util/packed_stride.hpp"
40
+
41
+ #include <cstdio>
42
+ #include <mutex>
43
+ #include <unordered_map>
44
+
45
+ namespace flash_rt {
46
+ namespace gemm {
47
+
48
+ // ─────────────────────────────────────────────────────────────────
49
+ // Default tile <128,128,256>, cluster <1,1,1>, schedule Auto.
50
+ // ─────────────────────────────────────────────────────────────────
51
+ namespace sm100_default {
52
+
53
+ using namespace cute;
54
+
55
+ using ElementA = cutlass::float_e2m1_t;
56
+ using ElementB = cutlass::float_e2m1_t;
57
+ using ElementC = cutlass::bfloat16_t;
58
+ using ElementD = cutlass::bfloat16_t;
59
+ using ElementAccumulator = float;
60
+ using ElementCompute = float;
61
+ using ElementSF = cutlass::float_ue4m3_t;
62
+
63
+ using LayoutA = cutlass::layout::RowMajor;
64
+ using LayoutB = cutlass::layout::ColumnMajor;
65
+ using LayoutC = cutlass::layout::RowMajor;
66
+ using LayoutD = cutlass::layout::RowMajor;
67
+
68
+ using ElementPairA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
69
+ using ElementPairB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
70
+
71
+ constexpr int AlignmentA = 16 * 8 / cutlass::sizeof_bits<ElementA>::value; // 32
72
+ constexpr int AlignmentB = 16 * 8 / cutlass::sizeof_bits<ElementB>::value; // 32
73
+ constexpr int AlignmentC = 128 / cutlass::sizeof_bits<ElementC>::value; // 8
74
+ constexpr int AlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value; // 8
75
+
76
+ using TileShape = Shape<_128, _128, _256>;
77
+ using ClusterShape = Shape<_1, _1, _1>;
78
+
79
+ using Sm1xxBlkScaledConfig = cutlass::detail::Sm1xxBlockScaledConfig<16>;
80
+
81
+ using CollectiveEpilogue =
82
+ typename cutlass::epilogue::collective::CollectiveBuilder<
83
+ cutlass::arch::Sm100, cutlass::arch::OpClassBlockScaledTensorOp,
84
+ TileShape, ClusterShape,
85
+ cutlass::epilogue::collective::EpilogueTileAuto,
86
+ ElementAccumulator, ElementCompute,
87
+ ElementC, LayoutC, AlignmentC,
88
+ ElementD, LayoutD, AlignmentD,
89
+ cutlass::epilogue::collective::EpilogueScheduleAuto
90
+ >::CollectiveOp;
91
+
92
+ using CollectiveMainloop =
93
+ typename cutlass::gemm::collective::CollectiveBuilder<
94
+ cutlass::arch::Sm100, cutlass::arch::OpClassBlockScaledTensorOp,
95
+ ElementPairA, LayoutA, AlignmentA,
96
+ ElementPairB, LayoutB, AlignmentB,
97
+ ElementAccumulator,
98
+ TileShape, ClusterShape,
99
+ cutlass::gemm::collective::StageCountAutoCarveout<
100
+ static_cast<int>(sizeof(typename CollectiveEpilogue::SharedStorage))>,
101
+ cutlass::gemm::collective::KernelScheduleAuto
102
+ >::CollectiveOp;
103
+
104
+ using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
105
+ Shape<int, int, int, int>,
106
+ CollectiveMainloop,
107
+ CollectiveEpilogue>;
108
+
109
+ using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
110
+
111
+ struct ShapeKey {
112
+ int M, N, K;
113
+ bool operator==(const ShapeKey& o) const {
114
+ return M == o.M && N == o.N && K == o.K;
115
+ }
116
+ };
117
+ struct ShapeKeyHash {
118
+ size_t operator()(const ShapeKey& k) const noexcept {
119
+ return (static_cast<size_t>(k.M) * 1315423911u)
120
+ ^ (static_cast<size_t>(k.N) * 2654435761u)
121
+ ^ static_cast<size_t>(k.K);
122
+ }
123
+ };
124
+ struct CachedWorkspace { void* ptr = nullptr; size_t size = 0; };
125
+
126
+ std::unordered_map<ShapeKey, CachedWorkspace, ShapeKeyHash> g_ws_cache;
127
+ std::mutex g_ws_mu;
128
+
129
+ void* get_workspace(int M, int N, int K, size_t needed) {
130
+ std::lock_guard<std::mutex> lk(g_ws_mu);
131
+ ShapeKey key{M, N, K};
132
+ auto it = g_ws_cache.find(key);
133
+ if (it != g_ws_cache.end() && it->second.size >= needed) return it->second.ptr;
134
+ if (it != g_ws_cache.end()) { cudaFree(it->second.ptr); g_ws_cache.erase(it); }
135
+ CachedWorkspace w; w.size = needed;
136
+ if (needed > 0) cudaMalloc(&w.ptr, needed);
137
+ g_ws_cache[key] = w;
138
+ return w.ptr;
139
+ }
140
+
141
+ cutlass::Status run_gemm(
142
+ const void* A_packed, const void* B_packed, void* D_bf16,
143
+ int M, int N, int K,
144
+ const void* SFA, const void* SFB,
145
+ float alpha,
146
+ cudaStream_t stream)
147
+ {
148
+ using StrideA = typename Gemm::GemmKernel::StrideA;
149
+ using StrideB = typename Gemm::GemmKernel::StrideB;
150
+ using StrideC = typename Gemm::GemmKernel::StrideC;
151
+ using StrideD = typename Gemm::GemmKernel::StrideD;
152
+
153
+ StrideA stride_A = cutlass::make_cute_packed_stride(StrideA{}, cute::make_shape(M, K, 1));
154
+ StrideB stride_B = cutlass::make_cute_packed_stride(StrideB{}, cute::make_shape(N, K, 1));
155
+ StrideC stride_C = cutlass::make_cute_packed_stride(StrideC{}, cute::make_shape(M, N, 1));
156
+ StrideD stride_D = cutlass::make_cute_packed_stride(StrideD{}, cute::make_shape(M, N, 1));
157
+
158
+ auto problem_shape_MNKL = cute::make_shape(M, N, K, 1);
159
+ auto layout_SFA = Sm1xxBlkScaledConfig::tile_atom_to_shape_SFA(problem_shape_MNKL);
160
+ auto layout_SFB = Sm1xxBlkScaledConfig::tile_atom_to_shape_SFB(problem_shape_MNKL);
161
+
162
+ using ArrayElementA = typename Gemm::GemmKernel::CollectiveMainloop::ArrayElementA;
163
+ using ArrayElementB = typename Gemm::GemmKernel::CollectiveMainloop::ArrayElementB;
164
+
165
+ typename Gemm::Arguments args{
166
+ cutlass::gemm::GemmUniversalMode::kGemm,
167
+ {M, N, K, 1},
168
+ {
169
+ reinterpret_cast<ArrayElementA const*>(A_packed), stride_A,
170
+ reinterpret_cast<ArrayElementB const*>(B_packed), stride_B,
171
+ reinterpret_cast<ElementSF const*>(SFA), layout_SFA,
172
+ reinterpret_cast<ElementSF const*>(SFB), layout_SFB
173
+ },
174
+ {
175
+ {alpha, 0.0f},
176
+ nullptr, stride_C,
177
+ reinterpret_cast<ElementD*>(D_bf16), stride_D
178
+ }
179
+ };
180
+
181
+ Gemm gemm;
182
+ size_t ws_size = Gemm::get_workspace_size(args);
183
+ void* ws_ptr = get_workspace(M, N, K, ws_size);
184
+
185
+ auto status = gemm.can_implement(args);
186
+ if (status != cutlass::Status::kSuccess) {
187
+ std::fprintf(stderr,
188
+ "[fp4_w4a16_gemm_sm100_bf16out] can_implement FAIL M=%d N=%d K=%d (status=%d)\n",
189
+ M, N, K, static_cast<int>(status));
190
+ return status;
191
+ }
192
+ status = gemm.initialize(args, ws_ptr, stream);
193
+ if (status != cutlass::Status::kSuccess) {
194
+ std::fprintf(stderr,
195
+ "[fp4_w4a16_gemm_sm100_bf16out] initialize FAIL M=%d N=%d K=%d (status=%d)\n",
196
+ M, N, K, static_cast<int>(status));
197
+ return status;
198
+ }
199
+ return gemm.run(stream);
200
+ }
201
+
202
+ } // namespace sm100_default
203
+
204
+ void fp4_w4a16_gemm_sm100_bf16out(
205
+ const void* A_packed, const void* B_packed, void* D_bf16,
206
+ int M, int N, int K,
207
+ const void* SFA, const void* SFB,
208
+ float alpha, cudaStream_t stream)
209
+ {
210
+ cutlass::Status status = sm100_default::run_gemm(
211
+ A_packed, B_packed, D_bf16, M, N, K, SFA, SFB, alpha, stream);
212
+ if (status != cutlass::Status::kSuccess) {
213
+ std::fprintf(stderr,
214
+ "[fp4_w4a16_gemm_sm100_bf16out] run FAIL M=%d N=%d K=%d (status=%d); D output undefined\n",
215
+ M, N, K, static_cast<int>(status));
216
+ }
217
+ }
218
+
219
+ // ─────────────────────────────────────────────────────────────────
220
+ // Wide-N tile <128,256,128>, cluster <1,1,1>, schedule Auto.
221
+ // ─────────────────────────────────────────────────────────────────
222
+ namespace sm100_widen {
223
+
224
+ using namespace cute;
225
+
226
+ using ElementA = cutlass::float_e2m1_t;
227
+ using ElementB = cutlass::float_e2m1_t;
228
+ using ElementC = cutlass::bfloat16_t;
229
+ using ElementD = cutlass::bfloat16_t;
230
+ using ElementAccumulator = float;
231
+ using ElementCompute = float;
232
+ using ElementSF = cutlass::float_ue4m3_t;
233
+ using LayoutA = cutlass::layout::RowMajor;
234
+ using LayoutB = cutlass::layout::ColumnMajor;
235
+ using LayoutC = cutlass::layout::RowMajor;
236
+ using LayoutD = cutlass::layout::RowMajor;
237
+ using ElementPairA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
238
+ using ElementPairB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
239
+ constexpr int AlignmentA = 16 * 8 / cutlass::sizeof_bits<ElementA>::value;
240
+ constexpr int AlignmentB = 16 * 8 / cutlass::sizeof_bits<ElementB>::value;
241
+ constexpr int AlignmentC = 128 / cutlass::sizeof_bits<ElementC>::value;
242
+ constexpr int AlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
243
+
244
+ using TileShape = Shape<_128, _256, _128>;
245
+ using ClusterShape = Shape<_1, _1, _1>;
246
+ using Sm1xxBlkScaledConfig = cutlass::detail::Sm1xxBlockScaledConfig<16>;
247
+
248
+ using CollectiveEpilogue =
249
+ typename cutlass::epilogue::collective::CollectiveBuilder<
250
+ cutlass::arch::Sm100, cutlass::arch::OpClassBlockScaledTensorOp,
251
+ TileShape, ClusterShape,
252
+ cutlass::epilogue::collective::EpilogueTileAuto,
253
+ ElementAccumulator, ElementCompute,
254
+ ElementC, LayoutC, AlignmentC,
255
+ ElementD, LayoutD, AlignmentD,
256
+ cutlass::epilogue::collective::EpilogueScheduleAuto
257
+ >::CollectiveOp;
258
+
259
+ using CollectiveMainloop =
260
+ typename cutlass::gemm::collective::CollectiveBuilder<
261
+ cutlass::arch::Sm100, cutlass::arch::OpClassBlockScaledTensorOp,
262
+ ElementPairA, LayoutA, AlignmentA,
263
+ ElementPairB, LayoutB, AlignmentB,
264
+ ElementAccumulator,
265
+ TileShape, ClusterShape,
266
+ cutlass::gemm::collective::StageCountAutoCarveout<
267
+ static_cast<int>(sizeof(typename CollectiveEpilogue::SharedStorage))>,
268
+ cutlass::gemm::collective::KernelScheduleAuto
269
+ >::CollectiveOp;
270
+
271
+ using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
272
+ Shape<int, int, int, int>,
273
+ CollectiveMainloop,
274
+ CollectiveEpilogue>;
275
+
276
+ using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
277
+
278
+ struct ShapeKey {
279
+ int M, N, K;
280
+ bool operator==(const ShapeKey& o) const {
281
+ return M == o.M && N == o.N && K == o.K;
282
+ }
283
+ };
284
+ struct ShapeKeyHash {
285
+ size_t operator()(const ShapeKey& k) const noexcept {
286
+ return (static_cast<size_t>(k.M) * 1315423911u)
287
+ ^ (static_cast<size_t>(k.N) * 2654435761u)
288
+ ^ static_cast<size_t>(k.K);
289
+ }
290
+ };
291
+ struct CachedWorkspace { void* ptr = nullptr; size_t size = 0; };
292
+
293
+ std::unordered_map<ShapeKey, CachedWorkspace, ShapeKeyHash> g_ws_cache_widen;
294
+ std::mutex g_ws_mu_widen;
295
+
296
+ void* get_workspace_widen(int M, int N, int K, size_t needed) {
297
+ std::lock_guard<std::mutex> lk(g_ws_mu_widen);
298
+ ShapeKey key{M, N, K};
299
+ auto it = g_ws_cache_widen.find(key);
300
+ if (it != g_ws_cache_widen.end() && it->second.size >= needed) return it->second.ptr;
301
+ if (it != g_ws_cache_widen.end()) { cudaFree(it->second.ptr); g_ws_cache_widen.erase(it); }
302
+ CachedWorkspace w; w.size = needed;
303
+ if (needed > 0) cudaMalloc(&w.ptr, needed);
304
+ g_ws_cache_widen[key] = w;
305
+ return w.ptr;
306
+ }
307
+
308
+ cutlass::Status run_gemm(
309
+ const void* A_packed, const void* B_packed, void* D_bf16,
310
+ int M, int N, int K,
311
+ const void* SFA, const void* SFB,
312
+ float alpha,
313
+ cudaStream_t stream)
314
+ {
315
+ using StrideA = typename Gemm::GemmKernel::StrideA;
316
+ using StrideB = typename Gemm::GemmKernel::StrideB;
317
+ using StrideC = typename Gemm::GemmKernel::StrideC;
318
+ using StrideD = typename Gemm::GemmKernel::StrideD;
319
+
320
+ StrideA stride_A = cutlass::make_cute_packed_stride(StrideA{}, cute::make_shape(M, K, 1));
321
+ StrideB stride_B = cutlass::make_cute_packed_stride(StrideB{}, cute::make_shape(N, K, 1));
322
+ StrideC stride_C = cutlass::make_cute_packed_stride(StrideC{}, cute::make_shape(M, N, 1));
323
+ StrideD stride_D = cutlass::make_cute_packed_stride(StrideD{}, cute::make_shape(M, N, 1));
324
+
325
+ auto problem_shape_MNKL = cute::make_shape(M, N, K, 1);
326
+ auto layout_SFA = Sm1xxBlkScaledConfig::tile_atom_to_shape_SFA(problem_shape_MNKL);
327
+ auto layout_SFB = Sm1xxBlkScaledConfig::tile_atom_to_shape_SFB(problem_shape_MNKL);
328
+
329
+ using ArrayElementA = typename Gemm::GemmKernel::CollectiveMainloop::ArrayElementA;
330
+ using ArrayElementB = typename Gemm::GemmKernel::CollectiveMainloop::ArrayElementB;
331
+
332
+ typename Gemm::Arguments args{
333
+ cutlass::gemm::GemmUniversalMode::kGemm,
334
+ {M, N, K, 1},
335
+ {
336
+ reinterpret_cast<ArrayElementA const*>(A_packed), stride_A,
337
+ reinterpret_cast<ArrayElementB const*>(B_packed), stride_B,
338
+ reinterpret_cast<ElementSF const*>(SFA), layout_SFA,
339
+ reinterpret_cast<ElementSF const*>(SFB), layout_SFB
340
+ },
341
+ {
342
+ {alpha, 0.0f},
343
+ nullptr, stride_C,
344
+ reinterpret_cast<ElementD*>(D_bf16), stride_D
345
+ }
346
+ };
347
+
348
+ Gemm gemm;
349
+ size_t ws_size = Gemm::get_workspace_size(args);
350
+ void* ws_ptr = get_workspace_widen(M, N, K, ws_size);
351
+
352
+ auto status = gemm.can_implement(args);
353
+ if (status != cutlass::Status::kSuccess) {
354
+ std::fprintf(stderr,
355
+ "[fp4_w4a16_gemm_sm100_bf16out_widen] can_implement FAIL M=%d N=%d K=%d (status=%d)\n",
356
+ M, N, K, static_cast<int>(status));
357
+ return status;
358
+ }
359
+ status = gemm.initialize(args, ws_ptr, stream);
360
+ if (status != cutlass::Status::kSuccess) {
361
+ std::fprintf(stderr,
362
+ "[fp4_w4a16_gemm_sm100_bf16out_widen] initialize FAIL M=%d N=%d K=%d (status=%d)\n",
363
+ M, N, K, static_cast<int>(status));
364
+ return status;
365
+ }
366
+ return gemm.run(stream);
367
+ }
368
+
369
+ } // namespace sm100_widen
370
+
371
+ void fp4_w4a16_gemm_sm100_bf16out_widen(
372
+ const void* A_packed, const void* B_packed, void* D_bf16,
373
+ int M, int N, int K,
374
+ const void* SFA, const void* SFB,
375
+ float alpha, cudaStream_t stream)
376
+ {
377
+ cutlass::Status status = sm100_widen::run_gemm(
378
+ A_packed, B_packed, D_bf16, M, N, K, SFA, SFB, alpha, stream);
379
+ if (status != cutlass::Status::kSuccess) {
380
+ std::fprintf(stderr,
381
+ "[fp4_w4a16_gemm_sm100_bf16out_widen] run FAIL M=%d N=%d K=%d (status=%d); D output undefined\n",
382
+ M, N, K, static_cast<int>(status));
383
+ }
384
+ }
385
+
386
+ // ─────────────────────────────────────────────────────────────────
387
+ // Pingpong placeholder: same default tile + Auto schedule. The Sm100
388
+ // dispatch under KernelScheduleAuto already exercises a 2SM pingpong-
389
+ // style schedule, so this entry exists for binding parity with the
390
+ // Sm120 surface. A dedicated alternate schedule may replace this
391
+ // after the Thor tile sweep.
392
+ // ─────────────────────────────────────────────────────────────────
393
+ void fp4_w4a16_gemm_sm100_bf16out_pingpong(
394
+ const void* A_packed, const void* B_packed, void* D_bf16,
395
+ int M, int N, int K,
396
+ const void* SFA, const void* SFB,
397
+ float alpha, cudaStream_t stream)
398
+ {
399
+ // Routes through the default tile until the Thor sweep adds a
400
+ // distinct pingpong-equivalent schedule.
401
+ cutlass::Status status = sm100_default::run_gemm(
402
+ A_packed, B_packed, D_bf16, M, N, K, SFA, SFB, alpha, stream);
403
+ if (status != cutlass::Status::kSuccess) {
404
+ std::fprintf(stderr,
405
+ "[fp4_w4a16_gemm_sm100_bf16out_pingpong] run FAIL M=%d N=%d K=%d (status=%d); D output undefined\n",
406
+ M, N, K, static_cast<int>(status));
407
+ }
408
+ }
409
+
410
+ } // namespace gemm
411
+ } // namespace flash_rt
csrc/gemm/fp4/cutlass_nvfp4_w4a16_gemm_sm100.cuh ADDED
@@ -0,0 +1,65 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ // SPDX-License-Identifier: Apache-2.0
2
+ //
3
+ // CUTLASS-based NVFP4 W4A16 GEMM for sm_100-class Blackwell (datacenter
4
+ // SM100 / Jetson AGX Thor SM110). Block-scaled FP4 GEMM matching the
5
+ // Qwen3.6 NVFP4 ckpt schema (compressed-tensors `nvfp4-pack-quantized`).
6
+ //
7
+ // Sibling of cutlass_nvfp4_w4a16_gemm_sm120.cuh. The two differ only in
8
+ // CUTLASS arch dispatch and kernel-schedule policy:
9
+ // - sm120: arch::Sm120 + KernelTmaWarpSpecializedCooperative /
10
+ // KernelTmaWarpSpecializedPingpong
11
+ // - sm100: arch::Sm100 + KernelScheduleAuto
12
+ // On Thor (sm_110a) the Sm100 dispatch path produces the correct
13
+ // blockscaled tcgen05 mainloop.
14
+ //
15
+ // Wire-format contract is identical to the sm120 variant, so the
16
+ // Python-side weight/scale layout and the activation quantizer are
17
+ // reused unchanged. The pybind layer binds these Thor symbols under
18
+ // the existing public names ``fp4_w4a16_gemm_sm120_bf16out*`` so the
19
+ // Qwen3.6 frontend code path does not need any hardware fork.
20
+
21
+ #pragma once
22
+
23
+ #include <cuda_runtime.h>
24
+
25
+ namespace flash_rt {
26
+ namespace gemm {
27
+
28
+ // Default tile <128,128,256>, cluster <1,1,1>, KernelScheduleAuto.
29
+ void fp4_w4a16_gemm_sm100_bf16out(
30
+ const void* A_packed, // (M, K/2) u8 row-major
31
+ const void* B_packed, // (N, K/2) u8 row-major (read as ColMajor (K,N))
32
+ void* D_bf16, // (M, N) bf16 row-major
33
+ int M, int N, int K,
34
+ const void* SFA, // (M, K/16) e4m3 (Sm1xx blockscaled atom layout)
35
+ const void* SFB, // (N, K/16) e4m3 (Sm1xx blockscaled atom layout)
36
+ float alpha, // = sf_global_a * sf_global_b
37
+ cudaStream_t stream);
38
+
39
+ // Wide-N tile <128,256,128>, cluster <1,1,1>, KernelScheduleAuto.
40
+ // For shapes with very large N (lm_head, MLP gate/up).
41
+ void fp4_w4a16_gemm_sm100_bf16out_widen(
42
+ const void* A_packed,
43
+ const void* B_packed,
44
+ void* D_bf16,
45
+ int M, int N, int K,
46
+ const void* SFA,
47
+ const void* SFB,
48
+ float alpha,
49
+ cudaStream_t stream);
50
+
51
+ // Default tile <128,128,256>, cluster <1,1,1>, KernelScheduleAuto.
52
+ // Kept as a separate symbol so callers can A/B against the default
53
+ // variant after the tile sweep produces a Thor-tuned schedule.
54
+ void fp4_w4a16_gemm_sm100_bf16out_pingpong(
55
+ const void* A_packed,
56
+ const void* B_packed,
57
+ void* D_bf16,
58
+ int M, int N, int K,
59
+ const void* SFA,
60
+ const void* SFB,
61
+ float alpha,
62
+ cudaStream_t stream);
63
+
64
+ } // namespace gemm
65
+ } // namespace flash_rt
csrc/gemm/fp4/cutlass_nvfp4_w4a16_gemm_sm120.cu ADDED
@@ -0,0 +1,690 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ // SPDX-License-Identifier: Apache-2.0
2
+ //
3
+ // CUTLASS NVFP4 W4A16 block-scaled GEMM, SM120a, BF16 output.
4
+ // Header: cutlass_nvfp4_w4a16_gemm_sm120.cuh
5
+ //
6
+ // Template is a direct port of NVIDIA's verified unit test at
7
+ // third_party/cutlass/test/unit/gemm/device/
8
+ // sm120_blockscaled_tensorop_gemm/sm120_bs_gemm_nvf4_nvf4_f32_bf16.cu
9
+ // (one config: TileShape <128,128,256>, ClusterShape <1,1,1>,
10
+ // KernelTmaWarpSpecializedCooperative, OpClassBlockScaledTensorOp).
11
+ //
12
+ // Why this is the "vendor-best" path (not hand-written):
13
+ // * Uses CUTLASS 4.x's `CollectiveBuilder` for SM120 BlockScaled
14
+ // mainloop + epilogue. NVIDIA tunes the mainloop schedule.
15
+ // * Same family as the existing FP8 SM120 GEMM (cutlass_sm120_block128_
16
+ // fp8_gemm.cu) — only the element types and `OpClass` differ.
17
+ // * No handwritten PTX or custom layout swizzling.
18
+ //
19
+ // Per-shape argument cache + workspace match the FP8 path so the
20
+ // hot-path is launch-only.
21
+
22
+ #include "cutlass_nvfp4_w4a16_gemm_sm120.cuh"
23
+
24
+ #include "cute/tensor.hpp"
25
+ #include "cute/atom/mma_atom.hpp"
26
+
27
+ #include "cutlass/cutlass.h"
28
+ #include "cutlass/numeric_types.h"
29
+ #include "cutlass/detail/sm100_blockscaled_layout.hpp"
30
+
31
+ #include "cutlass/epilogue/collective/collective_builder.hpp"
32
+ #include "cutlass/epilogue/collective/default_epilogue.hpp"
33
+ #include "cutlass/epilogue/thread/linear_combination.h"
34
+
35
+ #include "cutlass/gemm/collective/collective_builder.hpp"
36
+ #include "cutlass/gemm/dispatch_policy.hpp"
37
+ #include "cutlass/gemm/device/gemm_universal_adapter.h"
38
+ #include "cutlass/gemm/kernel/gemm_universal.hpp"
39
+
40
+ #include "cutlass/util/packed_stride.hpp"
41
+
42
+ #include <cstdio>
43
+ #include <mutex>
44
+ #include <unordered_map>
45
+
46
+ namespace flash_rt {
47
+ namespace gemm {
48
+
49
+ namespace {
50
+
51
+ using namespace cute;
52
+
53
+ // ── Element / layout types (copy of unit test "kernel_1") ────────
54
+ using ElementA = cutlass::float_e2m1_t;
55
+ using ElementB = cutlass::float_e2m1_t;
56
+ using ElementC = cutlass::bfloat16_t;
57
+ using ElementD = cutlass::bfloat16_t;
58
+ using ElementAccumulator = float;
59
+ using ElementCompute = float;
60
+ using ElementSF = cutlass::float_ue4m3_t;
61
+
62
+ using LayoutA = cutlass::layout::RowMajor;
63
+ using LayoutB = cutlass::layout::ColumnMajor;
64
+ // Public API hands us D in row-major (matches the FP8 sm120 kernel
65
+ // and what HF / our pipeline downstream expects for (M, N) tensors).
66
+ using LayoutC = cutlass::layout::RowMajor;
67
+ using LayoutD = cutlass::layout::RowMajor;
68
+
69
+ using ElementPairA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
70
+ using ElementPairB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
71
+
72
+ // 16-byte alignments in bits = 16 * 8 = 128. Convert to element count.
73
+ constexpr int AlignmentA = 16 * 8 / cutlass::sizeof_bits<ElementA>::value; // 32
74
+ constexpr int AlignmentB = 16 * 8 / cutlass::sizeof_bits<ElementB>::value; // 32
75
+ constexpr int AlignmentC = 128 / cutlass::sizeof_bits<ElementC>::value; // 8
76
+ constexpr int AlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value; // 8
77
+
78
+ using TileShape = Shape<_128, _128, _256>;
79
+ using ClusterShape = Shape<_1, _1, _1>;
80
+
81
+ // SF tensor layout helper. Vector size = 16 (NVFP4 group) per ckpt.
82
+ // Sm1xxBlockScaledConfig generates the (M-blk, K-blk) atom layout
83
+ // CUTLASS expects on-device. Our weight loader and act quantizer
84
+ // must produce SF in this exact layout (transformed once).
85
+ using Sm1xxBlkScaledConfig = cutlass::detail::Sm1xxBlockScaledConfig<16>;
86
+
87
+ using CollectiveEpilogue =
88
+ typename cutlass::epilogue::collective::CollectiveBuilder<
89
+ cutlass::arch::Sm120, cutlass::arch::OpClassTensorOp,
90
+ TileShape, ClusterShape,
91
+ cutlass::epilogue::collective::EpilogueTileAuto,
92
+ ElementAccumulator, ElementCompute,
93
+ ElementC, LayoutC, AlignmentC,
94
+ ElementD, LayoutD, AlignmentD,
95
+ cutlass::epilogue::collective::EpilogueScheduleAuto
96
+ >::CollectiveOp;
97
+
98
+ using CollectiveMainloop =
99
+ typename cutlass::gemm::collective::CollectiveBuilder<
100
+ cutlass::arch::Sm120, cutlass::arch::OpClassBlockScaledTensorOp,
101
+ ElementPairA, LayoutA, AlignmentA,
102
+ ElementPairB, LayoutB, AlignmentB,
103
+ ElementAccumulator,
104
+ TileShape, ClusterShape,
105
+ cutlass::gemm::collective::StageCountAutoCarveout<
106
+ static_cast<int>(sizeof(typename CollectiveEpilogue::SharedStorage))>,
107
+ cutlass::gemm::KernelTmaWarpSpecializedCooperative
108
+ >::CollectiveOp;
109
+
110
+ using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
111
+ Shape<int, int, int, int>,
112
+ CollectiveMainloop,
113
+ CollectiveEpilogue,
114
+ cutlass::gemm::PersistentScheduler>;
115
+
116
+ using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
117
+
118
+ // ── Per-shape workspace cache (mirrors FP8 path) ─────────────────
119
+ struct ShapeKey {
120
+ int M, N, K;
121
+ bool operator==(const ShapeKey& o) const {
122
+ return M == o.M && N == o.N && K == o.K;
123
+ }
124
+ };
125
+ struct ShapeKeyHash {
126
+ size_t operator()(const ShapeKey& k) const noexcept {
127
+ return (static_cast<size_t>(k.M) * 1315423911u)
128
+ ^ (static_cast<size_t>(k.N) * 2654435761u)
129
+ ^ static_cast<size_t>(k.K);
130
+ }
131
+ };
132
+
133
+ struct CachedWorkspace {
134
+ void* ptr = nullptr;
135
+ size_t size = 0;
136
+ };
137
+
138
+ std::unordered_map<ShapeKey, CachedWorkspace, ShapeKeyHash> g_ws_cache;
139
+ std::mutex g_ws_mu;
140
+
141
+ void* get_workspace(int M, int N, int K, size_t needed) {
142
+ std::lock_guard<std::mutex> lk(g_ws_mu);
143
+ ShapeKey key{M, N, K};
144
+ auto it = g_ws_cache.find(key);
145
+ if (it != g_ws_cache.end() && it->second.size >= needed) {
146
+ return it->second.ptr;
147
+ }
148
+ if (it != g_ws_cache.end()) {
149
+ cudaFree(it->second.ptr);
150
+ g_ws_cache.erase(it);
151
+ }
152
+ CachedWorkspace w;
153
+ w.size = needed;
154
+ if (needed > 0) {
155
+ cudaMalloc(&w.ptr, needed);
156
+ }
157
+ g_ws_cache[key] = w;
158
+ return w.ptr;
159
+ }
160
+
161
+ cutlass::Status run_gemm(
162
+ const void* A_packed, const void* B_packed, void* D_bf16,
163
+ int M, int N, int K,
164
+ const void* SFA, const void* SFB,
165
+ float alpha,
166
+ cudaStream_t stream)
167
+ {
168
+ using StrideA = typename Gemm::GemmKernel::StrideA;
169
+ using StrideB = typename Gemm::GemmKernel::StrideB;
170
+ using StrideC = typename Gemm::GemmKernel::StrideC;
171
+ using StrideD = typename Gemm::GemmKernel::StrideD;
172
+
173
+ StrideA stride_A = cutlass::make_cute_packed_stride(
174
+ StrideA{}, cute::make_shape(M, K, 1));
175
+ StrideB stride_B = cutlass::make_cute_packed_stride(
176
+ StrideB{}, cute::make_shape(N, K, 1));
177
+ StrideC stride_C = cutlass::make_cute_packed_stride(
178
+ StrideC{}, cute::make_shape(M, N, 1));
179
+ StrideD stride_D = cutlass::make_cute_packed_stride(
180
+ StrideD{}, cute::make_shape(M, N, 1));
181
+
182
+ auto problem_shape_MNKL = cute::make_shape(M, N, K, 1);
183
+ auto layout_SFA = Sm1xxBlkScaledConfig::tile_atom_to_shape_SFA(problem_shape_MNKL);
184
+ auto layout_SFB = Sm1xxBlkScaledConfig::tile_atom_to_shape_SFB(problem_shape_MNKL);
185
+
186
+ using ArrayElementA = typename Gemm::GemmKernel::CollectiveMainloop::ArrayElementA;
187
+ using ArrayElementB = typename Gemm::GemmKernel::CollectiveMainloop::ArrayElementB;
188
+
189
+ typename Gemm::Arguments args{
190
+ cutlass::gemm::GemmUniversalMode::kGemm,
191
+ {M, N, K, 1},
192
+ {
193
+ reinterpret_cast<ArrayElementA const*>(A_packed), stride_A,
194
+ reinterpret_cast<ArrayElementB const*>(B_packed), stride_B,
195
+ reinterpret_cast<ElementSF const*>(SFA), layout_SFA,
196
+ reinterpret_cast<ElementSF const*>(SFB), layout_SFB
197
+ },
198
+ {
199
+ {alpha, 0.0f}, // (alpha, beta)
200
+ nullptr, stride_C, // C unused (beta=0)
201
+ reinterpret_cast<ElementD*>(D_bf16), stride_D
202
+ }
203
+ };
204
+
205
+ Gemm gemm;
206
+ size_t ws_size = Gemm::get_workspace_size(args);
207
+ void* ws_ptr = get_workspace(M, N, K, ws_size);
208
+
209
+ auto status = gemm.can_implement(args);
210
+ if (status != cutlass::Status::kSuccess) {
211
+ std::fprintf(stderr,
212
+ "[fp4_w4a16_gemm_sm120_bf16out] can_implement FAIL "
213
+ "M=%d N=%d K=%d (status=%d)\n",
214
+ M, N, K, static_cast<int>(status));
215
+ return status;
216
+ }
217
+ status = gemm.initialize(args, ws_ptr, stream);
218
+ if (status != cutlass::Status::kSuccess) {
219
+ std::fprintf(stderr,
220
+ "[fp4_w4a16_gemm_sm120_bf16out] initialize FAIL "
221
+ "M=%d N=%d K=%d (status=%d)\n",
222
+ M, N, K, static_cast<int>(status));
223
+ return status;
224
+ }
225
+ return gemm.run(stream);
226
+ }
227
+
228
+ // Same default tile, but with a per-element residual C addend folded into the
229
+ // epilogue: D = alpha*(A*B) + C. The epilogue already carries the C operand
230
+ // (ElementC=bf16) — here we just feed C (beta=1) instead of nullptr (beta=0).
231
+ // Lets o_proj/down fuse their residual add, so the following rms_norm reads ONE
232
+ // tensor (D) instead of two (gemm_out + residual). C must be bf16 (M,N) row-major.
233
+ cutlass::Status run_gemm_residual(
234
+ const void* A_packed, const void* B_packed,
235
+ const void* C_residual, void* D_bf16,
236
+ int M, int N, int K,
237
+ const void* SFA, const void* SFB,
238
+ float alpha,
239
+ cudaStream_t stream)
240
+ {
241
+ using StrideA = typename Gemm::GemmKernel::StrideA;
242
+ using StrideB = typename Gemm::GemmKernel::StrideB;
243
+ using StrideC = typename Gemm::GemmKernel::StrideC;
244
+ using StrideD = typename Gemm::GemmKernel::StrideD;
245
+
246
+ StrideA stride_A = cutlass::make_cute_packed_stride(StrideA{}, cute::make_shape(M, K, 1));
247
+ StrideB stride_B = cutlass::make_cute_packed_stride(StrideB{}, cute::make_shape(N, K, 1));
248
+ StrideC stride_C = cutlass::make_cute_packed_stride(StrideC{}, cute::make_shape(M, N, 1));
249
+ StrideD stride_D = cutlass::make_cute_packed_stride(StrideD{}, cute::make_shape(M, N, 1));
250
+
251
+ auto problem_shape_MNKL = cute::make_shape(M, N, K, 1);
252
+ auto layout_SFA = Sm1xxBlkScaledConfig::tile_atom_to_shape_SFA(problem_shape_MNKL);
253
+ auto layout_SFB = Sm1xxBlkScaledConfig::tile_atom_to_shape_SFB(problem_shape_MNKL);
254
+
255
+ using ArrayElementA = typename Gemm::GemmKernel::CollectiveMainloop::ArrayElementA;
256
+ using ArrayElementB = typename Gemm::GemmKernel::CollectiveMainloop::ArrayElementB;
257
+
258
+ typename Gemm::Arguments args{
259
+ cutlass::gemm::GemmUniversalMode::kGemm,
260
+ {M, N, K, 1},
261
+ {
262
+ reinterpret_cast<ArrayElementA const*>(A_packed), stride_A,
263
+ reinterpret_cast<ArrayElementB const*>(B_packed), stride_B,
264
+ reinterpret_cast<ElementSF const*>(SFA), layout_SFA,
265
+ reinterpret_cast<ElementSF const*>(SFB), layout_SFB
266
+ },
267
+ {
268
+ {alpha, 1.0f}, // (alpha, beta=1)
269
+ reinterpret_cast<ElementC const*>(C_residual), stride_C,
270
+ reinterpret_cast<ElementD*>(D_bf16), stride_D
271
+ }
272
+ };
273
+
274
+ Gemm gemm;
275
+ size_t ws_size = Gemm::get_workspace_size(args);
276
+ void* ws_ptr = get_workspace(M, N, K, ws_size);
277
+
278
+ auto status = gemm.can_implement(args);
279
+ if (status != cutlass::Status::kSuccess) {
280
+ std::fprintf(stderr, "[fp4_w4a16_gemm_residual] can_implement FAIL "
281
+ "M=%d N=%d K=%d (status=%d)\n", M, N, K, static_cast<int>(status));
282
+ return status;
283
+ }
284
+ status = gemm.initialize(args, ws_ptr, stream);
285
+ if (status != cutlass::Status::kSuccess) return status;
286
+ return gemm.run(stream);
287
+ }
288
+
289
+ } // namespace
290
+
291
+ void fp4_w4a16_gemm_residual_sm120_bf16out(
292
+ const void* A_packed, const void* B_packed,
293
+ const void* C_residual, void* D_bf16,
294
+ int M, int N, int K,
295
+ const void* SFA, const void* SFB,
296
+ float alpha, cudaStream_t stream)
297
+ {
298
+ cutlass::Status status = run_gemm_residual(
299
+ A_packed, B_packed, C_residual, D_bf16, M, N, K, SFA, SFB, alpha, stream);
300
+ if (status != cutlass::Status::kSuccess) {
301
+ std::fprintf(stderr, "[fp4_w4a16_gemm_residual_sm120_bf16out] run FAIL "
302
+ "M=%d N=%d K=%d (status=%d)\n", M, N, K, static_cast<int>(status));
303
+ }
304
+ }
305
+
306
+ void fp4_w4a16_gemm_sm120_bf16out(
307
+ const void* A_packed,
308
+ const void* B_packed,
309
+ void* D_bf16,
310
+ int M, int N, int K,
311
+ const void* SFA,
312
+ const void* SFB,
313
+ float alpha,
314
+ cudaStream_t stream)
315
+ {
316
+ cutlass::Status status = run_gemm(
317
+ A_packed, B_packed, D_bf16, M, N, K, SFA, SFB, alpha, stream);
318
+ if (status != cutlass::Status::kSuccess) {
319
+ std::fprintf(stderr,
320
+ "[fp4_w4a16_gemm_sm120_bf16out] run FAIL "
321
+ "M=%d N=%d K=%d (status=%d); D output undefined\n",
322
+ M, N, K, static_cast<int>(status));
323
+ }
324
+ }
325
+
326
+ // ============================================================
327
+ // WIDEN variant: TileShape <128, 256, 128>. Same kernel template
328
+ // machinery as above but a wider N tile + narrower K. Profiled
329
+ // faster on shapes with very large N (lm_head N=248320: 88% BW vs
330
+ // 64% baseline; MLP gate/up N=17408: 66% vs 56%). Slower on small/
331
+ // medium-N shapes (k/v_proj N=1024, lin in_proj_z N=6144, etc.) so
332
+ // callers dispatch by shape.
333
+ // ============================================================
334
+ namespace widen {
335
+
336
+ using namespace cute;
337
+
338
+ using ElementA = cutlass::float_e2m1_t;
339
+ using ElementB = cutlass::float_e2m1_t;
340
+ using ElementC = cutlass::bfloat16_t;
341
+ using ElementD = cutlass::bfloat16_t;
342
+ using ElementAccumulator = float;
343
+ using ElementCompute = float;
344
+ using ElementSF = cutlass::float_ue4m3_t;
345
+ using LayoutA = cutlass::layout::RowMajor;
346
+ using LayoutB = cutlass::layout::ColumnMajor;
347
+ using LayoutC = cutlass::layout::RowMajor;
348
+ using LayoutD = cutlass::layout::RowMajor;
349
+ using ElementPairA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
350
+ using ElementPairB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
351
+ constexpr int AlignmentA = 16 * 8 / cutlass::sizeof_bits<ElementA>::value;
352
+ constexpr int AlignmentB = 16 * 8 / cutlass::sizeof_bits<ElementB>::value;
353
+ constexpr int AlignmentC = 128 / cutlass::sizeof_bits<ElementC>::value;
354
+ constexpr int AlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
355
+
356
+ using TileShape = Shape<_128, _256, _128>;
357
+ using ClusterShape = Shape<_1, _1, _1>;
358
+
359
+ using Sm1xxBlkScaledConfig = cutlass::detail::Sm1xxBlockScaledConfig<16>;
360
+
361
+ using CollectiveEpilogue =
362
+ typename cutlass::epilogue::collective::CollectiveBuilder<
363
+ cutlass::arch::Sm120, cutlass::arch::OpClassTensorOp,
364
+ TileShape, ClusterShape,
365
+ cutlass::epilogue::collective::EpilogueTileAuto,
366
+ ElementAccumulator, ElementCompute,
367
+ ElementC, LayoutC, AlignmentC,
368
+ ElementD, LayoutD, AlignmentD,
369
+ cutlass::epilogue::collective::EpilogueScheduleAuto
370
+ >::CollectiveOp;
371
+
372
+ using CollectiveMainloop =
373
+ typename cutlass::gemm::collective::CollectiveBuilder<
374
+ cutlass::arch::Sm120, cutlass::arch::OpClassBlockScaledTensorOp,
375
+ ElementPairA, LayoutA, AlignmentA,
376
+ ElementPairB, LayoutB, AlignmentB,
377
+ ElementAccumulator,
378
+ TileShape, ClusterShape,
379
+ cutlass::gemm::collective::StageCountAutoCarveout<
380
+ static_cast<int>(sizeof(typename CollectiveEpilogue::SharedStorage))>,
381
+ cutlass::gemm::KernelTmaWarpSpecializedCooperative
382
+ >::CollectiveOp;
383
+
384
+ using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
385
+ Shape<int, int, int, int>,
386
+ CollectiveMainloop,
387
+ CollectiveEpilogue,
388
+ cutlass::gemm::PersistentScheduler>;
389
+
390
+ using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
391
+
392
+ struct ShapeKey {
393
+ int M, N, K;
394
+ bool operator==(const ShapeKey& o) const {
395
+ return M == o.M && N == o.N && K == o.K;
396
+ }
397
+ };
398
+ struct ShapeKeyHash {
399
+ size_t operator()(const ShapeKey& k) const noexcept {
400
+ return (static_cast<size_t>(k.M) * 1315423911u)
401
+ ^ (static_cast<size_t>(k.N) * 2654435761u)
402
+ ^ static_cast<size_t>(k.K);
403
+ }
404
+ };
405
+ struct CachedWorkspace { void* ptr = nullptr; size_t size = 0; };
406
+ std::unordered_map<ShapeKey, CachedWorkspace, ShapeKeyHash> g_ws_cache_widen;
407
+ std::mutex g_ws_mu_widen;
408
+
409
+ void* get_workspace_widen(int M, int N, int K, size_t needed) {
410
+ std::lock_guard<std::mutex> lk(g_ws_mu_widen);
411
+ ShapeKey key{M, N, K};
412
+ auto it = g_ws_cache_widen.find(key);
413
+ if (it != g_ws_cache_widen.end() && it->second.size >= needed) {
414
+ return it->second.ptr;
415
+ }
416
+ if (it != g_ws_cache_widen.end()) {
417
+ cudaFree(it->second.ptr);
418
+ g_ws_cache_widen.erase(it);
419
+ }
420
+ CachedWorkspace w;
421
+ w.size = needed;
422
+ if (needed > 0) cudaMalloc(&w.ptr, needed);
423
+ g_ws_cache_widen[key] = w;
424
+ return w.ptr;
425
+ }
426
+
427
+ cutlass::Status run_gemm_widen(
428
+ const void* A_packed, const void* B_packed, void* D_bf16,
429
+ int M, int N, int K,
430
+ const void* SFA, const void* SFB,
431
+ float alpha,
432
+ cudaStream_t stream)
433
+ {
434
+ using StrideA = typename Gemm::GemmKernel::StrideA;
435
+ using StrideB = typename Gemm::GemmKernel::StrideB;
436
+ using StrideC = typename Gemm::GemmKernel::StrideC;
437
+ using StrideD = typename Gemm::GemmKernel::StrideD;
438
+
439
+ StrideA stride_A = cutlass::make_cute_packed_stride(
440
+ StrideA{}, cute::make_shape(M, K, 1));
441
+ StrideB stride_B = cutlass::make_cute_packed_stride(
442
+ StrideB{}, cute::make_shape(N, K, 1));
443
+ StrideC stride_C = cutlass::make_cute_packed_stride(
444
+ StrideC{}, cute::make_shape(M, N, 1));
445
+ StrideD stride_D = cutlass::make_cute_packed_stride(
446
+ StrideD{}, cute::make_shape(M, N, 1));
447
+
448
+ auto problem_shape_MNKL = cute::make_shape(M, N, K, 1);
449
+ auto layout_SFA = Sm1xxBlkScaledConfig::tile_atom_to_shape_SFA(problem_shape_MNKL);
450
+ auto layout_SFB = Sm1xxBlkScaledConfig::tile_atom_to_shape_SFB(problem_shape_MNKL);
451
+
452
+ using ArrayElementA = typename Gemm::GemmKernel::CollectiveMainloop::ArrayElementA;
453
+ using ArrayElementB = typename Gemm::GemmKernel::CollectiveMainloop::ArrayElementB;
454
+
455
+ typename Gemm::Arguments args{
456
+ cutlass::gemm::GemmUniversalMode::kGemm,
457
+ {M, N, K, 1},
458
+ {
459
+ reinterpret_cast<ArrayElementA const*>(A_packed), stride_A,
460
+ reinterpret_cast<ArrayElementB const*>(B_packed), stride_B,
461
+ reinterpret_cast<ElementSF const*>(SFA), layout_SFA,
462
+ reinterpret_cast<ElementSF const*>(SFB), layout_SFB
463
+ },
464
+ {
465
+ {alpha, 0.0f},
466
+ nullptr, stride_C,
467
+ reinterpret_cast<ElementD*>(D_bf16), stride_D
468
+ }
469
+ };
470
+
471
+ Gemm gemm;
472
+ size_t ws_size = Gemm::get_workspace_size(args);
473
+ void* ws_ptr = get_workspace_widen(M, N, K, ws_size);
474
+
475
+ auto status = gemm.can_implement(args);
476
+ if (status != cutlass::Status::kSuccess) {
477
+ std::fprintf(stderr,
478
+ "[fp4_w4a16_gemm_sm120_bf16out_widen] can_implement FAIL "
479
+ "M=%d N=%d K=%d (status=%d)\n",
480
+ M, N, K, static_cast<int>(status));
481
+ return status;
482
+ }
483
+ status = gemm.initialize(args, ws_ptr, stream);
484
+ if (status != cutlass::Status::kSuccess) {
485
+ std::fprintf(stderr,
486
+ "[fp4_w4a16_gemm_sm120_bf16out_widen] initialize FAIL "
487
+ "M=%d N=%d K=%d (status=%d)\n",
488
+ M, N, K, static_cast<int>(status));
489
+ return status;
490
+ }
491
+ return gemm.run(stream);
492
+ }
493
+
494
+ } // namespace widen
495
+
496
+ void fp4_w4a16_gemm_sm120_bf16out_widen(
497
+ const void* A_packed,
498
+ const void* B_packed,
499
+ void* D_bf16,
500
+ int M, int N, int K,
501
+ const void* SFA,
502
+ const void* SFB,
503
+ float alpha,
504
+ cudaStream_t stream)
505
+ {
506
+ cutlass::Status status = widen::run_gemm_widen(
507
+ A_packed, B_packed, D_bf16, M, N, K, SFA, SFB, alpha, stream);
508
+ if (status != cutlass::Status::kSuccess) {
509
+ std::fprintf(stderr,
510
+ "[fp4_w4a16_gemm_sm120_bf16out_widen] run FAIL "
511
+ "M=%d N=%d K=%d (status=%d); D output undefined\n",
512
+ M, N, K, static_cast<int>(status));
513
+ }
514
+ }
515
+
516
+ // ============================================================
517
+ // PINGPONG variant: default <128,128,256> tile with
518
+ // KernelTmaWarpSpecializedPingpong. This is intentionally separate from
519
+ // the production entrypoint so Qwen can A/B per shape before any dispatch
520
+ // policy is changed.
521
+ // ============================================================
522
+ namespace pingpong {
523
+
524
+ using namespace cute;
525
+
526
+ using ElementA = cutlass::float_e2m1_t;
527
+ using ElementB = cutlass::float_e2m1_t;
528
+ using ElementC = cutlass::bfloat16_t;
529
+ using ElementD = cutlass::bfloat16_t;
530
+ using ElementAccumulator = float;
531
+ using ElementCompute = float;
532
+ using ElementSF = cutlass::float_ue4m3_t;
533
+ using LayoutA = cutlass::layout::RowMajor;
534
+ using LayoutB = cutlass::layout::ColumnMajor;
535
+ using LayoutC = cutlass::layout::RowMajor;
536
+ using LayoutD = cutlass::layout::RowMajor;
537
+ using ElementPairA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
538
+ using ElementPairB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
539
+ constexpr int AlignmentA = 16 * 8 / cutlass::sizeof_bits<ElementA>::value;
540
+ constexpr int AlignmentB = 16 * 8 / cutlass::sizeof_bits<ElementB>::value;
541
+ constexpr int AlignmentC = 128 / cutlass::sizeof_bits<ElementC>::value;
542
+ constexpr int AlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
543
+
544
+ using TileShape = Shape<_128, _128, _256>;
545
+ using ClusterShape = Shape<_1, _1, _1>;
546
+ using Sm1xxBlkScaledConfig = cutlass::detail::Sm1xxBlockScaledConfig<16>;
547
+
548
+ using CollectiveEpilogue =
549
+ typename cutlass::epilogue::collective::CollectiveBuilder<
550
+ cutlass::arch::Sm120, cutlass::arch::OpClassTensorOp,
551
+ TileShape, ClusterShape,
552
+ cutlass::epilogue::collective::EpilogueTileAuto,
553
+ ElementAccumulator, ElementCompute,
554
+ ElementC, LayoutC, AlignmentC,
555
+ ElementD, LayoutD, AlignmentD,
556
+ cutlass::epilogue::collective::EpilogueScheduleAuto
557
+ >::CollectiveOp;
558
+
559
+ using CollectiveMainloop =
560
+ typename cutlass::gemm::collective::CollectiveBuilder<
561
+ cutlass::arch::Sm120, cutlass::arch::OpClassBlockScaledTensorOp,
562
+ ElementPairA, LayoutA, AlignmentA,
563
+ ElementPairB, LayoutB, AlignmentB,
564
+ ElementAccumulator,
565
+ TileShape, ClusterShape,
566
+ cutlass::gemm::collective::StageCountAutoCarveout<
567
+ static_cast<int>(sizeof(typename CollectiveEpilogue::SharedStorage))>,
568
+ cutlass::gemm::KernelTmaWarpSpecializedPingpong
569
+ >::CollectiveOp;
570
+
571
+ using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
572
+ Shape<int, int, int, int>,
573
+ CollectiveMainloop,
574
+ CollectiveEpilogue,
575
+ cutlass::gemm::PersistentScheduler>;
576
+
577
+ using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
578
+
579
+ std::unordered_map<ShapeKey, CachedWorkspace, ShapeKeyHash> g_ws_cache_pingpong;
580
+ std::mutex g_ws_mu_pingpong;
581
+
582
+ void* get_workspace_pingpong(int M, int N, int K, size_t needed) {
583
+ std::lock_guard<std::mutex> lk(g_ws_mu_pingpong);
584
+ ShapeKey key{M, N, K};
585
+ auto it = g_ws_cache_pingpong.find(key);
586
+ if (it != g_ws_cache_pingpong.end() && it->second.size >= needed) {
587
+ return it->second.ptr;
588
+ }
589
+ if (it != g_ws_cache_pingpong.end()) {
590
+ cudaFree(it->second.ptr);
591
+ g_ws_cache_pingpong.erase(it);
592
+ }
593
+ CachedWorkspace w;
594
+ w.size = needed;
595
+ if (needed > 0) cudaMalloc(&w.ptr, needed);
596
+ g_ws_cache_pingpong[key] = w;
597
+ return w.ptr;
598
+ }
599
+
600
+ cutlass::Status run_gemm_pingpong(
601
+ const void* A_packed, const void* B_packed, void* D_bf16,
602
+ int M, int N, int K,
603
+ const void* SFA, const void* SFB,
604
+ float alpha,
605
+ cudaStream_t stream)
606
+ {
607
+ using StrideA = typename Gemm::GemmKernel::StrideA;
608
+ using StrideB = typename Gemm::GemmKernel::StrideB;
609
+ using StrideC = typename Gemm::GemmKernel::StrideC;
610
+ using StrideD = typename Gemm::GemmKernel::StrideD;
611
+
612
+ StrideA stride_A = cutlass::make_cute_packed_stride(
613
+ StrideA{}, cute::make_shape(M, K, 1));
614
+ StrideB stride_B = cutlass::make_cute_packed_stride(
615
+ StrideB{}, cute::make_shape(N, K, 1));
616
+ StrideC stride_C = cutlass::make_cute_packed_stride(
617
+ StrideC{}, cute::make_shape(M, N, 1));
618
+ StrideD stride_D = cutlass::make_cute_packed_stride(
619
+ StrideD{}, cute::make_shape(M, N, 1));
620
+
621
+ auto problem_shape_MNKL = cute::make_shape(M, N, K, 1);
622
+ auto layout_SFA = Sm1xxBlkScaledConfig::tile_atom_to_shape_SFA(problem_shape_MNKL);
623
+ auto layout_SFB = Sm1xxBlkScaledConfig::tile_atom_to_shape_SFB(problem_shape_MNKL);
624
+
625
+ using ArrayElementA = typename Gemm::GemmKernel::CollectiveMainloop::ArrayElementA;
626
+ using ArrayElementB = typename Gemm::GemmKernel::CollectiveMainloop::ArrayElementB;
627
+
628
+ typename Gemm::Arguments args{
629
+ cutlass::gemm::GemmUniversalMode::kGemm,
630
+ {M, N, K, 1},
631
+ {
632
+ reinterpret_cast<ArrayElementA const*>(A_packed), stride_A,
633
+ reinterpret_cast<ArrayElementB const*>(B_packed), stride_B,
634
+ reinterpret_cast<ElementSF const*>(SFA), layout_SFA,
635
+ reinterpret_cast<ElementSF const*>(SFB), layout_SFB
636
+ },
637
+ {
638
+ {alpha, 0.0f},
639
+ nullptr, stride_C,
640
+ reinterpret_cast<ElementD*>(D_bf16), stride_D
641
+ }
642
+ };
643
+
644
+ Gemm gemm;
645
+ size_t ws_size = Gemm::get_workspace_size(args);
646
+ void* ws_ptr = get_workspace_pingpong(M, N, K, ws_size);
647
+
648
+ auto status = gemm.can_implement(args);
649
+ if (status != cutlass::Status::kSuccess) {
650
+ std::fprintf(stderr,
651
+ "[fp4_w4a16_gemm_sm120_bf16out_pingpong] can_implement FAIL "
652
+ "M=%d N=%d K=%d (status=%d)\n",
653
+ M, N, K, static_cast<int>(status));
654
+ return status;
655
+ }
656
+ status = gemm.initialize(args, ws_ptr, stream);
657
+ if (status != cutlass::Status::kSuccess) {
658
+ std::fprintf(stderr,
659
+ "[fp4_w4a16_gemm_sm120_bf16out_pingpong] initialize FAIL "
660
+ "M=%d N=%d K=%d (status=%d)\n",
661
+ M, N, K, static_cast<int>(status));
662
+ return status;
663
+ }
664
+ return gemm.run(stream);
665
+ }
666
+
667
+ } // namespace pingpong
668
+
669
+ void fp4_w4a16_gemm_sm120_bf16out_pingpong(
670
+ const void* A_packed,
671
+ const void* B_packed,
672
+ void* D_bf16,
673
+ int M, int N, int K,
674
+ const void* SFA,
675
+ const void* SFB,
676
+ float alpha,
677
+ cudaStream_t stream)
678
+ {
679
+ cutlass::Status status = pingpong::run_gemm_pingpong(
680
+ A_packed, B_packed, D_bf16, M, N, K, SFA, SFB, alpha, stream);
681
+ if (status != cutlass::Status::kSuccess) {
682
+ std::fprintf(stderr,
683
+ "[fp4_w4a16_gemm_sm120_bf16out_pingpong] run FAIL "
684
+ "M=%d N=%d K=%d (status=%d); D output undefined\n",
685
+ M, N, K, static_cast<int>(status));
686
+ }
687
+ }
688
+
689
+ } // namespace gemm
690
+ } // namespace flash_rt
csrc/gemm/fp4/cutlass_nvfp4_w4a16_gemm_sm120.cuh ADDED
@@ -0,0 +1,112 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ // SPDX-License-Identifier: Apache-2.0
2
+ //
3
+ // CUTLASS-based NVFP4 W4A16 GEMM for SM120a (RTX 5090 / Blackwell
4
+ // consumer GeForce). Native block-scaled FP4 GEMM matching the Qwen3.6
5
+ // NVFP4 ckpt schema (compressed-tensors `nvfp4-pack-quantized` format).
6
+ //
7
+ // Wraps NVIDIA's verified template from
8
+ // third_party/cutlass/test/unit/gemm/device/sm120_blockscaled_tensorop_gemm/
9
+ // sm120_bs_gemm_nvf4_nvf4_f32_bf16.cu — using `OpClassBlockScaledTensorOp`
10
+ // + `nv_float4_t<float_e2m1_t>` + `float_ue4m3_t` group scales, BF16
11
+ // output. SM_120 + SM_121 (RTX 5090 / 5080) gated.
12
+ //
13
+ // Schema (matches both A=act and B=weight after per-token NVFP4 quant):
14
+ // * elements : 4-bit FP e2m1, packed two per byte
15
+ // * group scale : FP8 ue4m3, one scale per 16-element group
16
+ // (`group_size = 16` per the ckpt config.json)
17
+ // * global scale: a single FP32 per tensor (fed via the epilogue's
18
+ // alpha so we get D = sf_global_a * sf_global_b *
19
+ // (A * B) with a single multiply instead of a
20
+ // per-tile rescale)
21
+ //
22
+ // Caller responsibilities:
23
+ // * A_packed, B_packed are u8 arrays viewing FP4 e2m1 (2x packed).
24
+ // A is row-major (M, K/2 byte-pairs). B is column-major weight
25
+ // view; we accept the natural HF row-major (N, K/2) layout and
26
+ // reinterpret as ColumnMajor (K, N) — same memory.
27
+ // * SFA, SFB are FP8 ue4m3 with the CUTLASS Sm1xx blockscaled tile
28
+ // atom layout. The activation quantizer (`quantize_fp4_dynamic_*`)
29
+ // is responsible for emitting SFA in this layout; the weight
30
+ // loader does the same SFB transform once at load.
31
+ //
32
+ // Constraints (verified by `can_implement` at runtime):
33
+ // * K must be a multiple of 16 (group size).
34
+ // * Pointer alignments: A/B 16 bytes (32 FP4 elements), C/D 16 bytes
35
+ // (8 BF16 elements).
36
+ // * M unrestricted. We pick a tile shape by M (small-M variant
37
+ // coming once profiled — first cut uses the unit test's
38
+ // <128,128,256> for all M and lets CUTLASS handle padding).
39
+
40
+ #pragma once
41
+
42
+ #include <cuda_runtime.h>
43
+
44
+ namespace flash_rt {
45
+ namespace gemm {
46
+
47
+ // NVFP4 W4A16 GEMM, BF16 output, SM120a (RTX 5090).
48
+ //
49
+ // A_packed : (M, K/2) u8 row-major (FP4 e2m1, 2 per byte)
50
+ // B_packed : (N, K/2) u8 row-major (FP4 e2m1, 2 per byte)
51
+ // — read as ColumnMajor (K, N)
52
+ // D_bf16 : (M, N) bf16 row-major
53
+ // SFA : (M, K/16) e4m3 (CUTLASS blockscaled atom layout)
54
+ // SFB : (N, K/16) e4m3 (CUTLASS blockscaled atom layout)
55
+ // alpha : fp32 scalar = act_global_scale * w_global_scale
56
+ //
57
+ // Stream-safe; per-shape arguments + workspace cached internally
58
+ // (mirrors the FP8 sm_120 kernel).
59
+ void fp4_w4a16_gemm_sm120_bf16out(
60
+ const void* A_packed, // (M, K/2) u8
61
+ const void* B_packed, // (N, K/2) u8
62
+ void* D_bf16, // (M, N) bf16
63
+ int M, int N, int K,
64
+ const void* SFA, // (M, K/16) e4m3 (Sm1xx blockscaled layout)
65
+ const void* SFB, // (N, K/16) e4m3 (Sm1xx blockscaled layout)
66
+ float alpha, // = sf_global_a * sf_global_b
67
+ cudaStream_t stream);
68
+
69
+ // Residual variant: D = alpha*(A*B) + C, C a per-element bf16 (M,N) addend.
70
+ // Folds the post-GEMM residual add (o_proj/down) into the epilogue so the
71
+ // following rms_norm reads one tensor (D) not two. Default tile (same as above).
72
+ void fp4_w4a16_gemm_residual_sm120_bf16out(
73
+ const void* A_packed,
74
+ const void* B_packed,
75
+ const void* C_residual, // (M, N) bf16 row-major
76
+ void* D_bf16,
77
+ int M, int N, int K,
78
+ const void* SFA,
79
+ const void* SFB,
80
+ float alpha,
81
+ cudaStream_t stream);
82
+
83
+ // Wide-N variant: TileShape <128, 256, 128>. For shapes with very
84
+ // large N (lm_head N=248320, MLP gate/up N=17408) where the wider N
85
+ // tile uses fewer waves and hits ~88%/66% peak BW vs ~64%/56% for
86
+ // the default <128,128,256> tile. For small/medium N (<= 6144) the
87
+ // default kernel is faster — caller dispatches by shape.
88
+ void fp4_w4a16_gemm_sm120_bf16out_widen(
89
+ const void* A_packed,
90
+ const void* B_packed,
91
+ void* D_bf16,
92
+ int M, int N, int K,
93
+ const void* SFA,
94
+ const void* SFB,
95
+ float alpha,
96
+ cudaStream_t stream);
97
+
98
+ // Same tile shape as the default kernel, but with
99
+ // KernelTmaWarpSpecializedPingpong. Kept as an explicit opt-in variant so
100
+ // callers can A/B schedule effects per shape without perturbing the default.
101
+ void fp4_w4a16_gemm_sm120_bf16out_pingpong(
102
+ const void* A_packed,
103
+ const void* B_packed,
104
+ void* D_bf16,
105
+ int M, int N, int K,
106
+ const void* SFA,
107
+ const void* SFB,
108
+ float alpha,
109
+ cudaStream_t stream);
110
+
111
+ } // namespace gemm
112
+ } // namespace flash_rt
csrc/gemm/fp4/fp4_w4a4_mma_warpsplit_sm120.cu ADDED
@@ -0,0 +1,147 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ // SPDX-License-Identifier: Apache-2.0
2
+ //
3
+ // Warp-split-K NVFP4 W4A4 M=1 GEMV for sm_120 — for the long-K / small-N
4
+ // decode shapes (mlp_down K=17408, out_proj) where the single-warp full_n
5
+ // kernel underfills the SMs. Instead of splitting K across BLOCKS (which
6
+ // needs a cross-block fp32 reduce that is fragile under CUDA-graph replay),
7
+ // this splits K across WARPS WITHIN one block: 8 N-cols/block, WARPS warps,
8
+ // each warp streams K/WARPS, and the warp partials are summed in SHARED
9
+ // MEMORY (intra-block) before the bf16 write. Single kernel, direct output,
10
+ // no cross-kernel intermediate -> graph-replay safe. More warps/SM (occupancy)
11
+ // + shorter per-warp streams give the same fill-the-SM win as block split-K.
12
+ // Additive: new file + new entry point.
13
+ //
14
+ // Header: fp4_w4a4_mma_warpsplit_sm120.cuh.
15
+ #include "gemm/fp4/fp4_w4a4_mma_warpsplit_sm120.cuh"
16
+
17
+ #include <cuda_bf16.h>
18
+ #include <cuda_runtime.h>
19
+ #include <cstdint>
20
+
21
+ #include "cute/arch/mma_sm120.hpp"
22
+ #include "cutlass/numeric_types.h"
23
+
24
+ namespace flash_rt {
25
+ namespace gemm {
26
+ namespace {
27
+
28
+ using AtomType = cute::SM120::BLOCKSCALED::SM120_16x8x64_TN_VS<
29
+ cutlass::float_e2m1_t, cutlass::float_e2m1_t, float,
30
+ cutlass::float_ue4m3_t, 16>;
31
+
32
+ __device__ __forceinline__ uint32_t fa(const uint8_t* s, int t0, int t1, int r) {
33
+ int ro = ((r & 1) ? (t1 + 8) : t1) * 32;
34
+ return *reinterpret_cast<const uint32_t*>(s + ro + t0 * 4 + ((r >> 1) & 1) * 16);
35
+ }
36
+ __device__ __forceinline__ uint32_t fb(const uint8_t* s, int t0, int t1, int r) {
37
+ return *reinterpret_cast<const uint32_t*>(s + t1 * 32 + t0 * 4 + r * 16);
38
+ }
39
+ __device__ __forceinline__ uint32_t fsa(const uint8_t* p, int u) {
40
+ return *reinterpret_cast<const uint32_t*>(p + u * 4);
41
+ }
42
+ __device__ __forceinline__ void cpa(uint8_t* d, const uint8_t* s) {
43
+ uint32_t i = __cvta_generic_to_shared(d);
44
+ asm volatile("cp.async.ca.shared.global.L2::128B [%0], [%1], 4;\n" :: "r"(i), "l"(s));
45
+ }
46
+ __device__ __forceinline__ void commit() { asm volatile("cp.async.commit_group;\n" ::); }
47
+ template <int N> __device__ __forceinline__ void waitg() {
48
+ asm volatile("cp.async.wait_group %0;\n" :: "n"(N));
49
+ }
50
+
51
+ template <int STAGES, int WARPS>
52
+ __global__ void warpsplit_kernel(
53
+ const uint8_t* __restrict__ A, const uint8_t* __restrict__ B,
54
+ const uint8_t* __restrict__ SFA, const uint8_t* __restrict__ SFB,
55
+ __nv_bfloat16* __restrict__ D, float alpha, int N, int K) {
56
+ // per-warp pipeline buffers
57
+ __shared__ uint8_t sA[WARPS][STAGES][16 * 32];
58
+ __shared__ uint8_t sSFA[WARPS][STAGES][16 * 4];
59
+ __shared__ uint8_t sB[WARPS][STAGES][8 * 32];
60
+ __shared__ uint8_t sSFB[WARPS][STAGES][8 * 4];
61
+ __shared__ float s_red[WARPS][8]; // each warp's 8 col partials
62
+
63
+ int tid = threadIdx.x, warp = tid >> 5, lane = tid & 31;
64
+ int my_n = blockIdx.x * 8;
65
+ const int KI = K / 64, KIw = KI / WARPS; // K-tiles per warp
66
+ const int kt0 = warp * KIw;
67
+ const int KH = K / 2, ncs = (K / 16 + 3) / 4;
68
+ int t0 = lane & 3, t1 = lane >> 2, sau = (lane & 1) * 8 + (lane >> 2), sbu = lane >> 2;
69
+ float c0 = 0, c1 = 0, c2 = 0, c3 = 0;
70
+
71
+ uint8_t (*mA)[16 * 32] = sA[warp];
72
+ uint8_t (*mSFA)[16 * 4] = sSFA[warp];
73
+ uint8_t (*mB)[8 * 32] = sB[warp];
74
+ uint8_t (*mSFB)[8 * 4] = sSFB[warp];
75
+
76
+ if (lane >= 1 && lane < 16) {
77
+ #pragma unroll
78
+ for (int st = 0; st < STAGES; ++st) {
79
+ int4* av = reinterpret_cast<int4*>(mA[st]); int4 z{0, 0, 0, 0};
80
+ av[lane * 2] = z; av[lane * 2 + 1] = z;
81
+ }
82
+ if (lane < 4) for (int st = 0; st < STAGES; ++st)
83
+ for (int i = 4 + lane; i < 64; i += 4) mSFA[st][i] = 0;
84
+ }
85
+ auto ld = [&](int bf, int kt) {
86
+ int bo = kt * 32;
87
+ if (lane < 8) cpa(mA[bf] + lane * 4, A + bo + lane * 4);
88
+ if (lane == 0) cpa(mSFA[bf], SFA + kt * 512);
89
+ for (int c = 0; c < 2; ++c) { int ch = lane + c * 32, col = ch >> 3, off = ch & 7;
90
+ cpa(mB[bf] + ch * 4, B + (my_n + col) * KH + bo + off * 4); }
91
+ if (lane < 8) { int col = my_n + lane, rb = col >> 7, ri = col & 127;
92
+ int si = rb * ncs + kt, ib = (ri & 31) * 16 + ((ri >> 5) & 3) * 4;
93
+ cpa(mSFB[bf] + lane * 4, SFB + si * 512 + ib); }
94
+ };
95
+ #pragma unroll
96
+ for (int st = 0; st < STAGES - 1; ++st) { if (st < KIw) ld(st, kt0 + st); commit(); }
97
+ for (int j = 0; j < KIw; ++j) {
98
+ int cb = j % STAGES, jp = j + STAGES - 1;
99
+ if (jp < KIw) ld(jp % STAGES, kt0 + jp);
100
+ commit(); waitg<STAGES - 1>(); __syncwarp();
101
+ uint32_t a0 = fa(mA[cb], t0, t1, 0), a1 = fa(mA[cb], t0, t1, 1);
102
+ uint32_t a2 = fa(mA[cb], t0, t1, 2), a3 = fa(mA[cb], t0, t1, 3);
103
+ uint32_t b0 = fb(mB[cb], t0, t1, 0), b1 = fb(mB[cb], t0, t1, 1);
104
+ uint32_t sfa = fsa(mSFA[cb], sau), sfb = fsa(mSFB[cb], sbu);
105
+ float d0, d1, d2, d3;
106
+ AtomType::fma(d0, d1, d2, d3, a0, a1, a2, a3, b0, b1, c0, c1, c2, c3, sfa, sfb);
107
+ c0 = d0; c1 = d1; c2 = d2; c3 = d3;
108
+ }
109
+ // each warp: lanes 0..3 hold row-0 partials c0 (col 2r) / c1 (col 2r+1)
110
+ int q = lane >> 2, r = lane & 3;
111
+ if (q == 0) { s_red[warp][r * 2] = c0; s_red[warp][r * 2 + 1] = c1; }
112
+ __syncthreads();
113
+ // warp 0 sums the WARPS partials per col and writes the bf16 output
114
+ if (warp == 0 && lane < 8) {
115
+ float acc = 0.f;
116
+ #pragma unroll
117
+ for (int w = 0; w < WARPS; ++w) acc += s_red[w][lane];
118
+ int col = my_n + lane;
119
+ if (col < N) D[col] = __float2bfloat16(acc * alpha);
120
+ }
121
+ }
122
+
123
+ } // namespace
124
+
125
+ int fp4_w4a4_mma_sm120_warpsplit_bf16out(
126
+ const void* A_packed, const void* B_packed, void* D_bf16, int N, int K,
127
+ const void* SFA, const void* SFB, float alpha, int warps, int stages,
128
+ cudaStream_t stream) {
129
+ if (!A_packed || !B_packed || !D_bf16 || !SFA || !SFB) return 1;
130
+ if (K <= 0 || (K % 64) != 0 || ((K / 64) % warps) != 0) return 2;
131
+ if (N <= 0 || (N % 8) != 0) return 3;
132
+ dim3 grid(N / 8);
133
+ auto a = reinterpret_cast<const uint8_t*>(A_packed);
134
+ auto b = reinterpret_cast<const uint8_t*>(B_packed);
135
+ auto sa = reinterpret_cast<const uint8_t*>(SFA);
136
+ auto sb = reinterpret_cast<const uint8_t*>(SFB);
137
+ auto d = reinterpret_cast<__nv_bfloat16*>(D_bf16);
138
+ #define WS_L(ST, WP) warpsplit_kernel<ST, WP><<<grid, WP * 32, 0, stream>>>(a, b, sa, sb, d, alpha, N, K)
139
+ if (warps == 2) { if (stages == 3) WS_L(3, 2); else if (stages == 4) WS_L(4, 2); else if (stages == 6) WS_L(6, 2); else return 5; }
140
+ else if (warps == 4) { if (stages == 3) WS_L(3, 4); else if (stages == 4) WS_L(4, 4); else if (stages == 6) WS_L(6, 4); else return 5; }
141
+ else if (warps == 8) { if (stages == 3) WS_L(3, 8); else if (stages == 4) WS_L(4, 8); else return 5; }
142
+ else return 6;
143
+ return 0;
144
+ }
145
+
146
+ } // namespace gemm
147
+ } // namespace flash_rt
csrc/gemm/fp4/fp4_w4a4_mma_warpsplit_sm120.cuh ADDED
@@ -0,0 +1,20 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ // SPDX-License-Identifier: Apache-2.0
2
+ //
3
+ // Warp-split-K NVFP4 W4A4 M=1 GEMV for sm_120: 8 N-cols/block, `warps` warps
4
+ // each streaming K/warps, partials summed in shared memory (intra-block) then
5
+ // written bf16. Graph-replay safe (no cross-block/cross-kernel intermediate).
6
+ // For long-K/small-N decode shapes (mlp_down, out_proj) the full_n kernel
7
+ // underfills. Additive.
8
+ #pragma once
9
+ #include <cuda_runtime.h>
10
+ namespace flash_rt {
11
+ namespace gemm {
12
+ // A_packed (K/2,), B_packed (N,K/2), D_bf16 (N,). SFA (K/16,), SFB (N,K/16)
13
+ // swizzled. warps in {2,4,8}, stages in {3,4,6}. N%8==0, K%64==0,
14
+ // (K/64)%warps==0. Returns 0 on success.
15
+ int fp4_w4a4_mma_sm120_warpsplit_bf16out(
16
+ const void* A_packed, const void* B_packed, void* D_bf16, int N, int K,
17
+ const void* SFA, const void* SFB, float alpha, int warps, int stages,
18
+ cudaStream_t stream);
19
+ } // namespace gemm
20
+ } // namespace flash_rt
csrc/gemm/fp4/sm110_dispatch.cu ADDED
@@ -0,0 +1,50 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ #include "gemm/fp4/sm110_dispatch.cuh"
2
+
3
+ #include "gemm/fp4/cutlass_nvfp4_w4a16_gemm_sm100.cuh"
4
+ #include "gemm/fp4/cutlass_fp4_gemm_bias_bf16_sm100.cuh"
5
+ #include "quantize/quantize_fp4_sfa_bf16.cuh"
6
+
7
+ namespace flash_rt::hub {
8
+ namespace {
9
+
10
+ void launch_sm110(
11
+ const void* a,
12
+ const void* b,
13
+ void* out,
14
+ int m,
15
+ int n,
16
+ int k,
17
+ const void* sfa,
18
+ const void* sfb,
19
+ float alpha,
20
+ int variant,
21
+ cudaStream_t stream) {
22
+ if (variant == 1) {
23
+ gemm::fp4_w4a16_gemm_sm100_bf16out_widen(
24
+ a, b, out, m, n, k, sfa, sfb, alpha, stream);
25
+ } else if (variant == 2) {
26
+ gemm::fp4_w4a16_gemm_sm100_bf16out_pingpong(
27
+ a, b, out, m, n, k, sfa, sfb, alpha, stream);
28
+ } else {
29
+ gemm::fp4_w4a16_gemm_sm100_bf16out(
30
+ a, b, out, m, n, k, sfa, sfb, alpha, stream);
31
+ }
32
+ }
33
+
34
+ struct Sm110DispatchRegistration {
35
+ Sm110DispatchRegistration() {
36
+ sm110_gemm_dispatch = &launch_sm110;
37
+ sm110_gemm_bias_dispatch = &fp4::cutlass_fp4_gemm_bias_bf16;
38
+ sm110_gemm_bias_residual_dispatch =
39
+ &fp4::cutlass_fp4_gemm_bias_res_bf16;
40
+ sm110_gemm_bias_gelu_fp4_dispatch =
41
+ &fp4::cutlass_fp4_gemm_bias_gelu_fp4out_bf16;
42
+ sm110_quantize_bf16_dispatch =
43
+ &fp4::quantize_fp4_dynamic_sfa_bf16_vec;
44
+ }
45
+ };
46
+
47
+ Sm110DispatchRegistration registration;
48
+
49
+ } // namespace
50
+ } // namespace flash_rt::hub
csrc/gemm/fp4/sm110_dispatch.cuh ADDED
@@ -0,0 +1,45 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ #pragma once
2
+
3
+ #include <cuda_runtime_api.h>
4
+
5
+ namespace flash_rt::hub {
6
+
7
+ using Sm110GemmDispatch = void (*)(
8
+ const void* a,
9
+ const void* b,
10
+ void* out,
11
+ int m,
12
+ int n,
13
+ int k,
14
+ const void* sfa,
15
+ const void* sfb,
16
+ float alpha,
17
+ int variant,
18
+ cudaStream_t stream);
19
+
20
+ using Sm110GemmBiasDispatch = int (*)(
21
+ const void* a, const void* sfa, const void* b, const void* sfb,
22
+ const void* bias, void* out, int m, int n, int k,
23
+ cudaStream_t stream);
24
+
25
+ using Sm110GemmBiasResidualDispatch = int (*)(
26
+ const void* a, const void* sfa, const void* b, const void* sfb,
27
+ const void* bias, const void* residual, void* out,
28
+ int m, int n, int k, cudaStream_t stream);
29
+
30
+ using Sm110GemmBiasGeluFp4Dispatch = int (*)(
31
+ const void* a, const void* sfa, const void* b, const void* sfb,
32
+ const void* bias, void* out_packed, void* out_sfa,
33
+ int m, int n, int k, cudaStream_t stream);
34
+
35
+ using Sm110QuantizeBf16Dispatch = int (*)(
36
+ const void* x, void* packed, void* sfa, int rows, int dim,
37
+ bool is_sfb, cudaStream_t stream);
38
+
39
+ extern Sm110GemmDispatch sm110_gemm_dispatch;
40
+ extern Sm110GemmBiasDispatch sm110_gemm_bias_dispatch;
41
+ extern Sm110GemmBiasResidualDispatch sm110_gemm_bias_residual_dispatch;
42
+ extern Sm110GemmBiasGeluFp4Dispatch sm110_gemm_bias_gelu_fp4_dispatch;
43
+ extern Sm110QuantizeBf16Dispatch sm110_quantize_bf16_dispatch;
44
+
45
+ } // namespace flash_rt::hub
csrc/quantize/quantize_fp4_sfa.cu ADDED
@@ -0,0 +1,194 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ // ============================================================================
2
+ // Fused FP4 quantize + CUTLASS SFA/SFB tile-interleaved scale write.
3
+ //
4
+ // Implementation = kernel_quantize_fp4 (quantize_fp4_dynamic.cu) with the
5
+ // scale-store address replaced by the CUTLASS layout functor. Packed fp4
6
+ // elements layout is UNCHANGED (still linear [N, D/2]), only the scale
7
+ // byte goes to a different location.
8
+ // ============================================================================
9
+ #include "quantize_fp4_sfa.cuh"
10
+
11
+ #include <cuda_bf16.h>
12
+ #include <cuda_fp16.h>
13
+ #include <cuda_fp8.h>
14
+
15
+ #ifndef CUTLASS_ARCH_MMA_SM100_SUPPORTED
16
+ # define CUTLASS_ARCH_MMA_SM100_SUPPORTED 1
17
+ #endif
18
+ #if defined(CUTLASS_ARCH_MMA_SM100_SUPPORTED) || defined(__CUDA_ARCH__)
19
+ # include "cutlass/cutlass.h"
20
+ # include "cutlass/detail/sm100_blockscaled_layout.hpp"
21
+ # include "cute/tensor.hpp"
22
+ # define FV_HAVE_CUTLASS 1
23
+ #else
24
+ # define FV_HAVE_CUTLASS 0
25
+ #endif
26
+
27
+ namespace flash_rt {
28
+ namespace fp4 {
29
+
30
+ #if FV_HAVE_CUTLASS
31
+
32
+ using Cfg = cutlass::detail::Sm1xxBlockScaledConfig<16>;
33
+
34
+ // ── Device helpers (duplicated locally to stay additive — not linking against
35
+ // quantize_fp4_dynamic.cu so we don't risk ODR issues). Identical logic. ──
36
+ __device__ __forceinline__ uint8_t fp32_to_e2m1_sfa(float x) {
37
+ uint8_t sign = (x < 0.f) ? 0x8u : 0x0u;
38
+ float ax = fabsf(x);
39
+ uint8_t mant;
40
+ if (ax <= 0.25f) mant = 0u;
41
+ else if (ax <= 0.75f) mant = 1u;
42
+ else if (ax <= 1.25f) mant = 2u;
43
+ else if (ax <= 1.75f) mant = 3u;
44
+ else if (ax <= 2.5f) mant = 4u;
45
+ else if (ax <= 3.5f) mant = 5u;
46
+ else if (ax <= 5.0f) mant = 6u;
47
+ else mant = 7u;
48
+ return sign | mant;
49
+ }
50
+
51
+ __device__ __forceinline__ __nv_fp8_e4m3 quantize_ue4m3_sfa(float x) {
52
+ float v = fmaxf(x, 0.f);
53
+ return __nv_fp8_e4m3(v);
54
+ }
55
+
56
+ __device__ __forceinline__ float dequantize_ue4m3_sfa(__nv_fp8_e4m3 s) {
57
+ return static_cast<float>(s);
58
+ }
59
+
60
+ __device__ __forceinline__ float input_to_float(__half value) {
61
+ return __half2float(value);
62
+ }
63
+
64
+ __device__ __forceinline__ float input_to_float(__nv_bfloat16 value) {
65
+ return __bfloat162float(value);
66
+ }
67
+
68
+ // ── Fused kernel ──
69
+ // One thread per (row, 16-element block). Scale byte goes to
70
+ // dst_sfa[layout(row, block_idx*16, 0)].
71
+ template <typename Input, class LayoutSF>
72
+ __global__ void kernel_quantize_fp4_sfa(
73
+ const Input* __restrict__ src,
74
+ uint8_t* __restrict__ dst_packed,
75
+ uint8_t* __restrict__ dst_sfa, // raw byte view of the CUTLASS SFA/SFB buffer
76
+ LayoutSF layout,
77
+ int N, int D) {
78
+ const int block_idx = blockIdx.x * blockDim.x + threadIdx.x;
79
+ const int row = blockIdx.y;
80
+ const int n_blocks = D / 16;
81
+ if (row >= N || block_idx >= n_blocks) return;
82
+
83
+ const int base = row * D + block_idx * 16;
84
+ float vals[16];
85
+ float amax = 0.f;
86
+ #pragma unroll
87
+ for (int i = 0; i < 16; ++i) {
88
+ vals[i] = input_to_float(src[base + i]);
89
+ float a = fabsf(vals[i]);
90
+ if (a > amax) amax = a;
91
+ }
92
+
93
+ float desired = amax / 6.f;
94
+ if (desired < 1e-12f) desired = 1e-12f;
95
+ __nv_fp8_e4m3 bs_q = quantize_ue4m3_sfa(desired);
96
+ float bs_dq = dequantize_ue4m3_sfa(bs_q);
97
+
98
+ // ── CORE FUSION: direct SFA tile-layout write ──
99
+ // LayoutSF maps (row, k, L=0) → byte offset. k is the full-K coordinate;
100
+ // SFVecSize=16 is baked in so any k in [block*16, block*16+15] hits the
101
+ // same offset. Use block_idx*16 (same convention as reshape_scales_sfa.cu).
102
+ int sfa_off = layout(row, block_idx * 16, 0);
103
+ dst_sfa[sfa_off] = *reinterpret_cast<uint8_t*>(&bs_q);
104
+
105
+ // Packed fp4 elements: layout unchanged.
106
+ const int out_base = row * (D / 2) + block_idx * 8;
107
+ const float inv_bs = 1.f / bs_dq;
108
+ #pragma unroll
109
+ for (int p = 0; p < 8; ++p) {
110
+ float v_lo = vals[2 * p ] * inv_bs;
111
+ float v_hi = vals[2 * p + 1] * inv_bs;
112
+ uint8_t lo = fp32_to_e2m1_sfa(v_lo);
113
+ uint8_t hi = fp32_to_e2m1_sfa(v_hi);
114
+ dst_packed[out_base + p] = lo | (hi << 4);
115
+ }
116
+ }
117
+
118
+ #endif // FV_HAVE_CUTLASS
119
+
120
+ int quantize_fp4_dynamic_sfa_fp16(
121
+ const void* src_fp16, void* dst_packed, void* dst_sfa,
122
+ int N, int D, bool is_sfb, cudaStream_t stream) {
123
+ #if FV_HAVE_CUTLASS
124
+ if (D % 16 != 0) return -1;
125
+ const int n_blocks = D / 16;
126
+ const int threads = 128;
127
+ dim3 grid((n_blocks + threads - 1) / threads, N);
128
+ dim3 block(threads);
129
+
130
+ // Shape: SFA uses (M=N, 1, K=D, L=1); SFB uses (1, N=N, K=D, L=1).
131
+ auto shape = cute::make_shape(
132
+ is_sfb ? 1 : N,
133
+ is_sfb ? N : 1,
134
+ D, 1);
135
+
136
+ if (is_sfb) {
137
+ auto layout = Cfg::tile_atom_to_shape_SFB(shape);
138
+ kernel_quantize_fp4_sfa<__half><<<grid, block, 0, stream>>>(
139
+ reinterpret_cast<const __half*>(src_fp16),
140
+ reinterpret_cast<uint8_t*>(dst_packed),
141
+ reinterpret_cast<uint8_t*>(dst_sfa),
142
+ layout, N, D);
143
+ } else {
144
+ auto layout = Cfg::tile_atom_to_shape_SFA(shape);
145
+ kernel_quantize_fp4_sfa<__half><<<grid, block, 0, stream>>>(
146
+ reinterpret_cast<const __half*>(src_fp16),
147
+ reinterpret_cast<uint8_t*>(dst_packed),
148
+ reinterpret_cast<uint8_t*>(dst_sfa),
149
+ layout, N, D);
150
+ }
151
+ cudaError_t e = cudaGetLastError();
152
+ return (e == cudaSuccess) ? 0 : -static_cast<int>(e);
153
+ #else
154
+ (void)src_fp16; (void)dst_packed; (void)dst_sfa;
155
+ (void)N; (void)D; (void)is_sfb; (void)stream;
156
+ return -2;
157
+ #endif
158
+ }
159
+
160
+ int quantize_fp4_dynamic_sfa_bf16(
161
+ const void* src_bf16, void* dst_packed, void* dst_sfa,
162
+ int N, int D, bool is_sfb, cudaStream_t stream) {
163
+ #if FV_HAVE_CUTLASS
164
+ if (D % 16 != 0) return -1;
165
+ const int n_blocks = D / 16;
166
+ const int threads = 128;
167
+ dim3 grid((n_blocks + threads - 1) / threads, N);
168
+ dim3 block(threads);
169
+ auto shape = cute::make_shape(is_sfb ? 1 : N, is_sfb ? N : 1, D, 1);
170
+
171
+ if (is_sfb) {
172
+ auto layout = Cfg::tile_atom_to_shape_SFB(shape);
173
+ kernel_quantize_fp4_sfa<__nv_bfloat16><<<grid, block, 0, stream>>>(
174
+ reinterpret_cast<const __nv_bfloat16*>(src_bf16),
175
+ reinterpret_cast<uint8_t*>(dst_packed),
176
+ reinterpret_cast<uint8_t*>(dst_sfa), layout, N, D);
177
+ } else {
178
+ auto layout = Cfg::tile_atom_to_shape_SFA(shape);
179
+ kernel_quantize_fp4_sfa<__nv_bfloat16><<<grid, block, 0, stream>>>(
180
+ reinterpret_cast<const __nv_bfloat16*>(src_bf16),
181
+ reinterpret_cast<uint8_t*>(dst_packed),
182
+ reinterpret_cast<uint8_t*>(dst_sfa), layout, N, D);
183
+ }
184
+ cudaError_t e = cudaGetLastError();
185
+ return (e == cudaSuccess) ? 0 : -static_cast<int>(e);
186
+ #else
187
+ (void)src_bf16; (void)dst_packed; (void)dst_sfa;
188
+ (void)N; (void)D; (void)is_sfb; (void)stream;
189
+ return -2;
190
+ #endif
191
+ }
192
+
193
+ } // namespace fp4
194
+ } // namespace flash_rt
csrc/quantize/quantize_fp4_sfa.cuh ADDED
@@ -0,0 +1,41 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ // ============================================================================
2
+ // FlashRT — fused (FP4 quantize + CUTLASS SFA/SFB tile-interleave) kernel.
3
+ //
4
+ // Equivalent to:
5
+ // quantize_fp4_dynamic_fp16(src, packed, linear_scales, N, D)
6
+ // reshape_linear_scales_to_sfa(linear_scales, sfa, N, D, is_sfb)
7
+ // in a SINGLE kernel launch. Scale byte is written directly to the CUTLASS
8
+ // tile-interleaved offset — linear_scales intermediate buffer is gone.
9
+ //
10
+ // Additive: does NOT modify quantize_fp4_dynamic.* or reshape_scales_sfa.*.
11
+ // Both remain callable for existing paths.
12
+ // ============================================================================
13
+ #pragma once
14
+ #include <cuda_runtime.h>
15
+
16
+ namespace flash_rt {
17
+ namespace fp4 {
18
+
19
+ // fp16 [N, D] → packed [N, D/2] (e2m1) + SFA/SFB tile-interleaved UE4M3 scales.
20
+ // is_sfb = false → SFA layout (use for A = activation, shape [M=N, K=D])
21
+ // is_sfb = true → SFB layout (use for B = weight, shape [N=N, K=D])
22
+ // Returns 0 on success.
23
+ int quantize_fp4_dynamic_sfa_fp16(
24
+ const void* src_fp16,
25
+ void* dst_packed,
26
+ void* dst_sfa,
27
+ int N, int D, bool is_sfb,
28
+ cudaStream_t stream);
29
+
30
+ // BF16 [N, D] -> the exact same packed E2M1 + CUTLASS SFA/SFB layout as the
31
+ // FP16 entry. This avoids a standalone BF16-to-FP16 conversion in decode
32
+ // pipelines whose activations are already BF16.
33
+ int quantize_fp4_dynamic_sfa_bf16(
34
+ const void* src_bf16,
35
+ void* dst_packed,
36
+ void* dst_sfa,
37
+ int N, int D, bool is_sfb,
38
+ cudaStream_t stream);
39
+
40
+ } // namespace fp4
41
+ } // namespace flash_rt
csrc/quantize/quantize_fp4_sfa_bf16.cu ADDED
@@ -0,0 +1,142 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ // ============================================================================
2
+ // bf16-input vectorized fused FP4 quantize + CUTLASS SFA/SFB scales.
3
+ //
4
+ // Same per-block scale selection and e2m1 rounding as
5
+ // quantize_fp4_dynamic_sfa_fp16_vec, with bf16 source elements. Each
6
+ // thread quantizes one 16-element block: two 16-byte loads, one 8-byte
7
+ // packed store, one SFA byte at the tile-interleaved offset.
8
+ // ============================================================================
9
+ #include "quantize_fp4_sfa_bf16.cuh"
10
+
11
+ #include <cuda_bf16.h>
12
+ #include <cuda_fp8.h>
13
+
14
+ #if defined(CUTLASS_ARCH_MMA_SM100_SUPPORTED) || defined(__CUDA_ARCH__)
15
+ # include "cutlass/cutlass.h"
16
+ # include "cutlass/detail/sm100_blockscaled_layout.hpp"
17
+ # include "cute/tensor.hpp"
18
+ # define FV_HAVE_CUTLASS 1
19
+ #else
20
+ # define FV_HAVE_CUTLASS 0
21
+ #endif
22
+
23
+ namespace flash_rt {
24
+ namespace fp4 {
25
+
26
+ #if FV_HAVE_CUTLASS
27
+
28
+ namespace {
29
+
30
+ using CfgVecB = cutlass::detail::Sm1xxBlockScaledConfig<16>;
31
+
32
+ __device__ __forceinline__ uint8_t fp32_to_e2m1_bvec(float x) {
33
+ uint8_t sign = (x < 0.f) ? 0x8u : 0x0u;
34
+ float ax = fabsf(x);
35
+ uint8_t mant;
36
+ if (ax <= 0.25f) mant = 0u;
37
+ else if (ax <= 0.75f) mant = 1u;
38
+ else if (ax <= 1.25f) mant = 2u;
39
+ else if (ax <= 1.75f) mant = 3u;
40
+ else if (ax <= 2.5f) mant = 4u;
41
+ else if (ax <= 3.5f) mant = 5u;
42
+ else if (ax <= 5.0f) mant = 6u;
43
+ else mant = 7u;
44
+ return sign | mant;
45
+ }
46
+
47
+ template <class LayoutSF>
48
+ __global__ void kernel_quantize_fp4_sfa_bf16_vec(
49
+ const int4* __restrict__ src, // bf16 [N, D] as int4 (8 elements)
50
+ uint2* __restrict__ dst_packed, // [N, D/2] bytes as uint2 (1 block)
51
+ uint8_t* __restrict__ dst_sfa,
52
+ LayoutSF layout,
53
+ int N, int D8) { // D8 = D / 8 int4 chunks per row
54
+ const int block_idx = blockIdx.x * blockDim.x + threadIdx.x;
55
+ const int row = blockIdx.y;
56
+ const int n_blocks = D8 >> 1; // 16 elements per block
57
+ if (row >= N || block_idx >= n_blocks) return;
58
+
59
+ const int4 raw0 = src[row * D8 + 2 * block_idx];
60
+ const int4 raw1 = src[row * D8 + 2 * block_idx + 1];
61
+ const __nv_bfloat16* h0 = reinterpret_cast<const __nv_bfloat16*>(&raw0);
62
+ const __nv_bfloat16* h1 = reinterpret_cast<const __nv_bfloat16*>(&raw1);
63
+
64
+ float vals[16];
65
+ float amax = 0.f;
66
+ #pragma unroll
67
+ for (int i = 0; i < 8; ++i) {
68
+ vals[i] = __bfloat162float(h0[i]);
69
+ vals[8 + i] = __bfloat162float(h1[i]);
70
+ }
71
+ #pragma unroll
72
+ for (int i = 0; i < 16; ++i) {
73
+ const float a = fabsf(vals[i]);
74
+ if (a > amax) amax = a;
75
+ }
76
+
77
+ float desired = amax / 6.f;
78
+ if (desired < 1e-12f) desired = 1e-12f;
79
+ __nv_fp8_e4m3 bs_q = __nv_fp8_e4m3(fmaxf(desired, 0.f));
80
+ const float bs_dq = static_cast<float>(bs_q);
81
+
82
+ dst_sfa[layout(row, block_idx * 16, 0)] =
83
+ *reinterpret_cast<uint8_t*>(&bs_q);
84
+
85
+ const float inv_bs = 1.f / bs_dq;
86
+ uint2 out;
87
+ uint8_t* ob = reinterpret_cast<uint8_t*>(&out);
88
+ #pragma unroll
89
+ for (int p = 0; p < 8; ++p) {
90
+ const uint8_t lo = fp32_to_e2m1_bvec(vals[2 * p] * inv_bs);
91
+ const uint8_t hi = fp32_to_e2m1_bvec(vals[2 * p + 1] * inv_bs);
92
+ ob[p] = static_cast<uint8_t>(lo | (hi << 4));
93
+ }
94
+ dst_packed[row * n_blocks + block_idx] = out;
95
+ }
96
+
97
+ } // namespace
98
+
99
+ #endif // FV_HAVE_CUTLASS
100
+
101
+ int quantize_fp4_dynamic_sfa_bf16_vec(
102
+ const void* src_bf16, void* dst_packed, void* dst_sfa,
103
+ int N, int D, bool is_sfb, cudaStream_t stream) {
104
+ #if FV_HAVE_CUTLASS
105
+ if (D % 16 != 0) return -1;
106
+ if ((reinterpret_cast<uintptr_t>(src_bf16) & 15) ||
107
+ (reinterpret_cast<uintptr_t>(dst_packed) & 7)) return -1;
108
+ const int n_blocks = D / 16;
109
+ const int threads = 128;
110
+ dim3 grid((n_blocks + threads - 1) / threads, N);
111
+
112
+ auto shape = cute::make_shape(
113
+ is_sfb ? 1 : N,
114
+ is_sfb ? N : 1,
115
+ D, 1);
116
+
117
+ if (is_sfb) {
118
+ auto layout = CfgVecB::tile_atom_to_shape_SFB(shape);
119
+ kernel_quantize_fp4_sfa_bf16_vec<<<grid, threads, 0, stream>>>(
120
+ reinterpret_cast<const int4*>(src_bf16),
121
+ reinterpret_cast<uint2*>(dst_packed),
122
+ reinterpret_cast<uint8_t*>(dst_sfa),
123
+ layout, N, D >> 3);
124
+ } else {
125
+ auto layout = CfgVecB::tile_atom_to_shape_SFA(shape);
126
+ kernel_quantize_fp4_sfa_bf16_vec<<<grid, threads, 0, stream>>>(
127
+ reinterpret_cast<const int4*>(src_bf16),
128
+ reinterpret_cast<uint2*>(dst_packed),
129
+ reinterpret_cast<uint8_t*>(dst_sfa),
130
+ layout, N, D >> 3);
131
+ }
132
+ const cudaError_t e = cudaGetLastError();
133
+ return (e == cudaSuccess) ? 0 : -static_cast<int>(e);
134
+ #else
135
+ (void)src_bf16; (void)dst_packed; (void)dst_sfa;
136
+ (void)N; (void)D; (void)is_sfb; (void)stream;
137
+ return -2;
138
+ #endif
139
+ }
140
+
141
+ } // namespace fp4
142
+ } // namespace flash_rt
csrc/quantize/quantize_fp4_sfa_bf16.cuh ADDED
@@ -0,0 +1,24 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ // ============================================================================
2
+ // FlashRT — bf16-input fused NVFP4 quantize + CUTLASS SFA/SFB scale write.
3
+ //
4
+ // bf16 companion of quantize_fp4_dynamic_sfa_fp16 for pipelines whose
5
+ // activations are bf16 (GR00T N1.7 DiT). Additive: new symbols only.
6
+ // ============================================================================
7
+ #pragma once
8
+
9
+ #include <cuda_runtime.h>
10
+
11
+ namespace flash_rt {
12
+ namespace fp4 {
13
+
14
+ // Quantize a bf16 [N, D] row-major tensor to packed e2m1 [N, D/2] plus
15
+ // UE4M3 per-16-element scales written directly at the CUTLASS
16
+ // tile-interleaved SFA/SFB offsets. Vectorized (16-byte loads, 8-byte
17
+ // packed stores). Returns 0 on success, -1 on unsupported shape or
18
+ // misaligned buffers, -2 when built without CUTLASS.
19
+ int quantize_fp4_dynamic_sfa_bf16_vec(
20
+ const void* src_bf16, void* dst_packed, void* dst_sfa,
21
+ int N, int D, bool is_sfb, cudaStream_t stream);
22
+
23
+ } // namespace fp4
24
+ } // namespace flash_rt
examples/README.md ADDED
@@ -0,0 +1,10 @@
 
 
 
 
 
 
 
 
 
 
 
1
+ # fp4-gemm Examples
2
+
3
+ These examples show direct Hub-style usage of `flashrt/fp4-gemm`.
4
+
5
+ ```bash
6
+ python fp4-gemm/examples/fp4_gemm_linear.py
7
+ ```
8
+
9
+ The quantization helper is included for validation and small examples. In a
10
+ runtime, weights should normally be prepacked and loaded as FP4/SFA/SFB buffers.
examples/fp4_gemm_linear.py ADDED
@@ -0,0 +1,29 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ #!/usr/bin/env python3
2
+ """Minimal Hub-style call for flashrt/fp4-gemm."""
3
+
4
+ from __future__ import annotations
5
+
6
+ import torch
7
+ from kernels import get_kernel
8
+
9
+
10
+ def main() -> None:
11
+ if not torch.cuda.is_available():
12
+ raise SystemExit("CUDA is required")
13
+
14
+ ops = get_kernel("flashrt/fp4-gemm", version=1, trust_remote_code=True)
15
+
16
+ x = torch.randn((32, 256), device="cuda", dtype=torch.float16)
17
+ w = torch.randn((512, 256), device="cuda", dtype=torch.float16)
18
+
19
+ a_packed, sfa = ops.quantize_fp4_sfa_fp16(x, is_sfb=False)
20
+ b_packed, sfb = ops.quantize_fp4_sfa_fp16(w, is_sfb=True)
21
+ y = ops.nvfp4_gemm_bf16(a_packed, b_packed, sfa, sfb, alpha=1.0)
22
+
23
+ print("a_packed", tuple(a_packed.shape), a_packed.dtype)
24
+ print("b_packed", tuple(b_packed.shape), b_packed.dtype)
25
+ print("output", tuple(y.shape), y.dtype)
26
+
27
+
28
+ if __name__ == "__main__":
29
+ main()
flake.nix ADDED
@@ -0,0 +1,18 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ description = "Flake for FlashRT FP4 GEMM kernels";
3
+
4
+ inputs = {
5
+ # huggingface/kernels#741: CUTLASS 4.4.2 plus the corrected 4.5.2 hash.
6
+ kernel-builder.url = "github:huggingface/kernels/870e825d881664e39f9287a27a74ef63ff3c545e";
7
+ };
8
+
9
+ outputs =
10
+ {
11
+ self,
12
+ kernel-builder,
13
+ }:
14
+ kernel-builder.lib.genKernelFlakeOutputs {
15
+ inherit self;
16
+ path = ./.;
17
+ };
18
+ }
tests/test_fp4_gemm.py ADDED
@@ -0,0 +1,668 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ #!/usr/bin/env python3
2
+ """Correctness tests for fp4-gemm."""
3
+
4
+ from __future__ import annotations
5
+
6
+ import argparse
7
+ import importlib
8
+ import json
9
+ import os
10
+ import sys
11
+ from dataclasses import asdict, dataclass
12
+ from pathlib import Path
13
+
14
+ import torch
15
+
16
+
17
+ ROOT = Path(__file__).resolve().parents[2]
18
+ PACKAGE = ROOT / "fp4-gemm"
19
+ REGISTRATION_INCLUDE = (
20
+ ROOT.parent
21
+ / "kernels"
22
+ / "kernel-builder"
23
+ / "src"
24
+ / "pyproject"
25
+ / "templates"
26
+ / "torch"
27
+ )
28
+ DEFAULT_CUTLASS_INCLUDE = (
29
+ ROOT.parent
30
+ / "flashrt_pr31_review"
31
+ / "third_party"
32
+ / "cutlass"
33
+ / "include"
34
+ )
35
+
36
+ SHAPES = {
37
+ "small_m16_n128_k128": (16, 128, 128),
38
+ "small_m32_n256_k256": (32, 256, 256),
39
+ "mlp_tile_m64_n512_k512": (64, 512, 512),
40
+ }
41
+
42
+ SM110_SHAPES = {
43
+ "pi05_action_gate_up": (51, 16384, 2048),
44
+ "pi05_action_down": (51, 2048, 8192),
45
+ "groot_n17_dit_qkv": (41, 4608, 1536),
46
+ "groot_n17_dit_ffn_up": (41, 6144, 1536),
47
+ "groot_n17_dit_ffn_down": (41, 1536, 6144),
48
+ "groot_legacy_dit_qkv": (51, 4608, 1536),
49
+ "groot_backbone_gate_up": (277, 16384, 2048),
50
+ "cosmos_edge_action": (64, 9216, 2048),
51
+ "lingbot_action_gate_up": (105, 16384, 2048),
52
+ }
53
+
54
+ EPILOGUE_SHAPES = {
55
+ "epilogue_tile": (64, 512, 512),
56
+ "motus_up": (360, 14336, 3072),
57
+ "motus_down": (360, 3072, 14336),
58
+ }
59
+
60
+ MODES = {
61
+ "smoke": ["small_m16_n128_k128"],
62
+ "full": list(SHAPES),
63
+ "thor-models": list(SM110_SHAPES),
64
+ }
65
+
66
+
67
+ @dataclass
68
+ class Metrics:
69
+ shape: str
70
+ M: int
71
+ N: int
72
+ K: int
73
+ workload: str
74
+ variant: int | None
75
+ max_abs: float
76
+ mean_abs: float
77
+ p99_abs: float
78
+ cosine: float
79
+ passed: bool
80
+
81
+
82
+ class SourceOps:
83
+ def __init__(self, namespace: str) -> None:
84
+ self._ops = getattr(torch.ops, namespace)
85
+
86
+ @staticmethod
87
+ def sfa_size_bytes(rows: int, dim: int) -> int:
88
+ n_blocks = dim // 16
89
+ n_row_super = (rows + 127) // 128
90
+ n_col_super = (n_blocks + 3) // 4
91
+ return n_row_super * n_col_super * 512
92
+
93
+ def alloc_fp4(self, rows: int, dim: int):
94
+ return (
95
+ torch.empty((rows, dim // 2), device="cuda", dtype=torch.uint8),
96
+ torch.empty((self.sfa_size_bytes(rows, dim),), device="cuda", dtype=torch.uint8),
97
+ )
98
+
99
+ def quantize_fp4_sfa_fp16(self, x, packed, sfa, is_sfb=False):
100
+ self._ops.quantize_fp4_sfa_fp16(x, packed, sfa, bool(is_sfb))
101
+
102
+ def quantize_fp4_sfa_bf16(self, x, packed, sfa, is_sfb=False):
103
+ self._ops.quantize_fp4_sfa_bf16(x, packed, sfa, bool(is_sfb))
104
+
105
+ def dequantize_fp4_sfa_fp16(self, packed, sfa, out, is_sfb=False):
106
+ self._ops.dequantize_fp4_sfa_fp16(packed, sfa, out, bool(is_sfb))
107
+
108
+ def nvfp4_gemm_bf16(self, a, b, sfa, sfb, out, alpha=1.0, variant=0):
109
+ self._ops.nvfp4_gemm_bf16(a, b, sfa, sfb, out, float(alpha), int(variant))
110
+
111
+ def nvfp4_gemm_bias_bf16(self, a, b, sfa, sfb, bias, out):
112
+ self._ops.nvfp4_gemm_bias_bf16(a, b, sfa, sfb, bias, out)
113
+
114
+ def nvfp4_gemm_bias_residual_bf16(
115
+ self, a, b, sfa, sfb, bias, residual, out
116
+ ):
117
+ self._ops.nvfp4_gemm_bias_residual_bf16(
118
+ a, b, sfa, sfb, bias, residual, out
119
+ )
120
+
121
+ def nvfp4_gemm_residual_bf16(self, a, b, sfa, sfb, residual, out, alpha=1.0):
122
+ self._ops.nvfp4_gemm_residual_bf16(
123
+ a, b, sfa, sfb, residual, out, float(alpha)
124
+ )
125
+
126
+ def nvfp4_gemm_bias_gelu_bf16(self, a, b, sfa, sfb, bias, out, alpha=1.0):
127
+ self._ops.nvfp4_gemm_bias_gelu_bf16(
128
+ a, b, sfa, sfb, bias, out, float(alpha)
129
+ )
130
+
131
+ def nvfp4_gemm_bias_gelu_nvfp4(
132
+ self, a, b, sfa, sfb, bias, out_packed, out_sfa, alpha=1.0
133
+ ):
134
+ self._ops.nvfp4_gemm_bias_gelu_nvfp4(
135
+ a, b, sfa, sfb, bias, out_packed, out_sfa, float(alpha)
136
+ )
137
+
138
+ def nvfp4_gemm_streamk_bf16(self, a, b, sfa, sfb, out, alpha=1.0):
139
+ self._ops.nvfp4_gemm_streamk_bf16(a, b, sfa, sfb, out, float(alpha))
140
+
141
+ def nvfp4_gemm_streamk_bias_bf16(
142
+ self, a, b, sfa, sfb, bias, out, alpha=1.0
143
+ ):
144
+ self._ops.nvfp4_gemm_streamk_bias_bf16(
145
+ a, b, sfa, sfb, bias, out, float(alpha)
146
+ )
147
+
148
+
149
+ class InstalledOps:
150
+ """Adapt the public return-value API to the in-place test interface."""
151
+
152
+ def __init__(self, module) -> None:
153
+ self._module = module
154
+
155
+ def sfa_size_bytes(self, rows: int, dim: int) -> int:
156
+ return int(self._module.sfa_size_bytes(rows, dim))
157
+
158
+ def alloc_fp4(self, rows: int, dim: int):
159
+ return (
160
+ torch.empty((rows, dim // 2), device="cuda", dtype=torch.uint8),
161
+ torch.empty(
162
+ (self._module.sfa_size_bytes(rows, dim),),
163
+ device="cuda",
164
+ dtype=torch.uint8,
165
+ ),
166
+ )
167
+
168
+ def quantize_fp4_sfa_fp16(self, x, packed, sfa, is_sfb=False):
169
+ self._module.quantize_fp4_sfa_fp16(
170
+ x, packed=packed, sfa=sfa, is_sfb=bool(is_sfb)
171
+ )
172
+
173
+ def quantize_fp4_sfa_bf16(self, x, packed, sfa, is_sfb=False):
174
+ self._module.quantize_fp4_sfa_bf16(
175
+ x, packed=packed, sfa=sfa, is_sfb=bool(is_sfb)
176
+ )
177
+
178
+ def dequantize_fp4_sfa_fp16(self, packed, sfa, out, is_sfb=False):
179
+ self._module.dequantize_fp4_sfa_fp16(
180
+ packed, sfa, out=out, is_sfb=bool(is_sfb)
181
+ )
182
+
183
+ def nvfp4_gemm_bf16(self, a, b, sfa, sfb, out, alpha=1.0, variant=0):
184
+ self._module.nvfp4_gemm_bf16(
185
+ a,
186
+ b,
187
+ sfa,
188
+ sfb,
189
+ alpha=float(alpha),
190
+ out=out,
191
+ variant=int(variant),
192
+ )
193
+
194
+ def nvfp4_gemm_bias_bf16(self, a, b, sfa, sfb, bias, out):
195
+ self._module.nvfp4_gemm_bias_bf16(
196
+ a, b, sfa, sfb, bias, out=out
197
+ )
198
+
199
+ def nvfp4_gemm_bias_residual_bf16(
200
+ self, a, b, sfa, sfb, bias, residual, out
201
+ ):
202
+ self._module.nvfp4_gemm_bias_residual_bf16(
203
+ a, b, sfa, sfb, bias, residual, out=out
204
+ )
205
+
206
+ def nvfp4_gemm_residual_bf16(self, a, b, sfa, sfb, residual, out, alpha=1.0):
207
+ self._module.nvfp4_gemm_residual_bf16(
208
+ a, b, sfa, sfb, residual, alpha=float(alpha), out=out
209
+ )
210
+
211
+ def nvfp4_gemm_bias_gelu_bf16(self, a, b, sfa, sfb, bias, out, alpha=1.0):
212
+ self._module.nvfp4_gemm_bias_gelu_bf16(
213
+ a, b, sfa, sfb, bias, alpha=float(alpha), out=out
214
+ )
215
+
216
+ def nvfp4_gemm_bias_gelu_nvfp4(
217
+ self, a, b, sfa, sfb, bias, out_packed, out_sfa, alpha=1.0
218
+ ):
219
+ self._module.nvfp4_gemm_bias_gelu_nvfp4(
220
+ a, b, sfa, sfb, bias, alpha=float(alpha),
221
+ out_packed=out_packed, out_sfa=out_sfa,
222
+ )
223
+
224
+ def nvfp4_gemm_streamk_bf16(self, a, b, sfa, sfb, out, alpha=1.0):
225
+ self._module.nvfp4_gemm_streamk_bf16(
226
+ a, b, sfa, sfb, alpha=float(alpha), out=out
227
+ )
228
+
229
+ def nvfp4_gemm_streamk_bias_bf16(
230
+ self, a, b, sfa, sfb, bias, out, alpha=1.0
231
+ ):
232
+ self._module.nvfp4_gemm_streamk_bias_bf16(
233
+ a, b, sfa, sfb, bias, alpha=float(alpha), out=out
234
+ )
235
+
236
+
237
+ def _current_arch_list() -> str:
238
+ major, minor = torch.cuda.get_device_capability(0)
239
+ if (major, minor) == (11, 0):
240
+ return "11.0a"
241
+ if major >= 12:
242
+ return "12.0a"
243
+ return f"{major}.{minor}"
244
+
245
+
246
+ def load_source_ops() -> SourceOps:
247
+ from torch.utils.cpp_extension import load
248
+
249
+ cutlass_include = Path(os.environ.get("FLASHRT_CUTLASS_INCLUDE", str(DEFAULT_CUTLASS_INCLUDE)))
250
+ if not REGISTRATION_INCLUDE.is_dir():
251
+ raise RuntimeError(f"missing kernel-builder registration include: {REGISTRATION_INCLUDE}")
252
+ if not cutlass_include.is_dir():
253
+ raise RuntimeError(f"missing CUTLASS include path: {cutlass_include}")
254
+ os.environ.setdefault("TORCH_CUDA_ARCH_LIST", _current_arch_list())
255
+ namespace = "fp4_gemm_source_test"
256
+ capability = torch.cuda.get_device_capability(0)
257
+ if capability == (11, 0):
258
+ gemm_sources = [
259
+ str(PACKAGE / "csrc" / "gemm" / "fp4" / "cutlass_nvfp4_w4a16_gemm_sm100.cu"),
260
+ str(PACKAGE / "csrc" / "gemm" / "fp4" / "cutlass_fp4_gemm_bias_bf16_sm100.cu"),
261
+ str(PACKAGE / "csrc" / "quantize" / "quantize_fp4_sfa_bf16.cu"),
262
+ str(PACKAGE / "csrc" / "gemm" / "fp4" / "sm110_dispatch.cu"),
263
+ ]
264
+ source_define = "-DFLASHRT_FP4_GEMM_SOURCE_SM110_ONLY"
265
+ else:
266
+ gemm_sources = [
267
+ str(PACKAGE / "csrc" / "gemm" / "fp4" / "cutlass_nvfp4_w4a16_gemm_sm120.cu"),
268
+ str(PACKAGE / "csrc" / "gemm" / "fp4" / "fp4_w4a4_mma_warpsplit_sm120.cu"),
269
+ str(PACKAGE / "csrc" / "gemm" / "fp4" / "cutlass_nvfp4_gemm_bias_gelu_bf16out_sm120.cu"),
270
+ str(PACKAGE / "csrc" / "gemm" / "fp4" / "cutlass_nvfp4_gemm_bias_gelu_fp4out_sm120.cu"),
271
+ str(PACKAGE / "csrc" / "gemm" / "fp4" / "cutlass_nvfp4_gemm_dn_streamk_bias_sm120.cu"),
272
+ ]
273
+ source_define = None
274
+ load(
275
+ name=namespace,
276
+ sources=[
277
+ str(PACKAGE / "torch-ext" / "torch_binding.cpp"),
278
+ *gemm_sources,
279
+ str(PACKAGE / "csrc" / "quantize" / "quantize_fp4_sfa.cu"),
280
+ str(PACKAGE / "csrc" / "dequantize_fp4_sfa.cu"),
281
+ ],
282
+ extra_include_paths=[
283
+ str(PACKAGE / "csrc"),
284
+ str(cutlass_include),
285
+ str(REGISTRATION_INCLUDE),
286
+ ],
287
+ extra_cflags=[flag for flag in ["-O3", "-DCUDA_KERNEL", source_define] if flag],
288
+ extra_cuda_cflags=[
289
+ "-O3",
290
+ "--expt-relaxed-constexpr",
291
+ "--expt-extended-lambda",
292
+ "-DCUDA_KERNEL",
293
+ "-DCUTLASS_ARCH_MMA_SM100_SUPPORTED=1",
294
+ *([source_define] if source_define else []),
295
+ ],
296
+ verbose=False,
297
+ )
298
+ return SourceOps(namespace)
299
+
300
+
301
+ def load_installed_ops(artifact: str | None):
302
+ if artifact:
303
+ sys.path.insert(0, artifact)
304
+ try:
305
+ return InstalledOps(importlib.import_module("fp4_gemm"))
306
+ finally:
307
+ if artifact:
308
+ sys.path.remove(artifact)
309
+
310
+
311
+ def make_inputs(m: int, n: int, k: int, seed: int):
312
+ gen = torch.Generator(device="cuda")
313
+ gen.manual_seed(seed)
314
+ a = (torch.randn((m, k), device="cuda", generator=gen) * 0.25).to(torch.float16).contiguous()
315
+ b = (torch.randn((n, k), device="cuda", generator=gen) * 0.25).to(torch.float16).contiguous()
316
+ return a, b
317
+
318
+
319
+ def metrics(got: torch.Tensor, expected: torch.Tensor) -> tuple[float, float, float, float]:
320
+ diff = (got.float() - expected.float()).abs().flatten()
321
+ return (
322
+ float(diff.max().item()),
323
+ float(diff.mean().item()),
324
+ float(torch.quantile(diff, 0.99).item()),
325
+ float(torch.nn.functional.cosine_similarity(got.float().flatten(), expected.float().flatten(), dim=0).item()),
326
+ )
327
+
328
+
329
+ def check_bf16_threshold(max_abs: float, mean_abs: float, p99_abs: float, cosine: float) -> bool:
330
+ return max_abs <= 0.125 and mean_abs <= 0.005 and p99_abs <= 0.03125 and cosine >= 0.999
331
+
332
+
333
+ def select_sm110_variant(shape: tuple[int, int, int]) -> int:
334
+ _m, n, k = shape
335
+ if n >= 4 * k:
336
+ return 1
337
+ if n == 3 * k:
338
+ return 2
339
+ return 0
340
+
341
+
342
+ def prepare_quantized(ops: SourceOps, m: int, n: int, k: int):
343
+ a_fp16, b_fp16 = make_inputs(m, n, k, seed=7000 + m + n + k)
344
+ a_packed, sfa = ops.alloc_fp4(m, k)
345
+ b_packed, sfb = ops.alloc_fp4(n, k)
346
+ ops.quantize_fp4_sfa_fp16(a_fp16, a_packed, sfa, False)
347
+ ops.quantize_fp4_sfa_fp16(b_fp16, b_packed, sfb, True)
348
+ a_deq = torch.empty_like(a_fp16)
349
+ b_deq = torch.empty_like(b_fp16)
350
+ ops.dequantize_fp4_sfa_fp16(a_packed, sfa, a_deq, False)
351
+ ops.dequantize_fp4_sfa_fp16(b_packed, sfb, b_deq, True)
352
+ torch.cuda.synchronize()
353
+ expected = (a_deq.float() @ b_deq.float().T).to(torch.bfloat16)
354
+ return a_packed, b_packed, sfa, sfb, expected
355
+
356
+
357
+ def prepare_quantized_full(ops: SourceOps, m: int, n: int, k: int):
358
+ a_fp16, b_fp16 = make_inputs(m, n, k, seed=9000 + m + n + k)
359
+ a_packed, sfa = ops.alloc_fp4(m, k)
360
+ b_packed, sfb = ops.alloc_fp4(n, k)
361
+ ops.quantize_fp4_sfa_fp16(a_fp16, a_packed, sfa, False)
362
+ ops.quantize_fp4_sfa_fp16(b_fp16, b_packed, sfb, True)
363
+ a_deq = torch.empty_like(a_fp16)
364
+ b_deq = torch.empty_like(b_fp16)
365
+ ops.dequantize_fp4_sfa_fp16(a_packed, sfa, a_deq, False)
366
+ ops.dequantize_fp4_sfa_fp16(b_packed, sfb, b_deq, True)
367
+ return a_packed, b_packed, sfa, sfb, a_deq, b_deq
368
+
369
+
370
+ def run_case(ops: SourceOps, name: str, shape: tuple[int, int, int]) -> list[Metrics]:
371
+ m, n, k = shape
372
+ a_packed, b_packed, sfa, sfb, expected = prepare_quantized(ops, m, n, k)
373
+ results: list[Metrics] = []
374
+ variants = (-1, 0, 1, 2) if torch.cuda.get_device_capability(0) == (11, 0) else (0, 1, 2)
375
+ for variant in variants:
376
+ out = torch.empty((m, n), device="cuda", dtype=torch.bfloat16)
377
+ ops.nvfp4_gemm_bf16(a_packed, b_packed, sfa, sfb, out, 1.0, variant)
378
+ torch.cuda.synchronize()
379
+ max_abs, mean_abs, p99_abs, cosine = metrics(out, expected)
380
+ results.append(
381
+ Metrics(
382
+ shape=name,
383
+ M=m,
384
+ N=n,
385
+ K=k,
386
+ workload="nvfp4_gemm_bf16",
387
+ variant=variant,
388
+ max_abs=max_abs,
389
+ mean_abs=mean_abs,
390
+ p99_abs=p99_abs,
391
+ cosine=cosine,
392
+ passed=check_bf16_threshold(max_abs, mean_abs, p99_abs, cosine),
393
+ )
394
+ )
395
+
396
+ return results
397
+
398
+
399
+ def result_row(name, shape, workload, got, expected, *, fp4_output=False):
400
+ max_abs, mean_abs, p99_abs, cosine = metrics(got, expected)
401
+ if fp4_output:
402
+ mean_magnitude = float(expected.float().abs().mean().item())
403
+ rms = float(expected.float().square().mean().sqrt().item())
404
+ passed = (
405
+ cosine >= 0.9993
406
+ and mean_abs / max(mean_magnitude, 1e-12) <= 0.01
407
+ and p99_abs / max(rms, 1e-12) <= 0.15
408
+ )
409
+ else:
410
+ passed = (
411
+ cosine >= 0.999
412
+ and mean_abs <= 0.008
413
+ and p99_abs <= 0.0625
414
+ )
415
+ return Metrics(
416
+ shape=name,
417
+ M=shape[0],
418
+ N=shape[1],
419
+ K=shape[2],
420
+ workload=workload,
421
+ variant=None,
422
+ max_abs=max_abs,
423
+ mean_abs=mean_abs,
424
+ p99_abs=p99_abs,
425
+ cosine=cosine,
426
+ passed=passed,
427
+ )
428
+
429
+
430
+ def run_epilogue_case(ops, name: str, shape: tuple[int, int, int]):
431
+ m, n, k = shape
432
+ a, b, sfa, sfb, a_deq, b_deq = prepare_quantized_full(ops, m, n, k)
433
+ matmul = a_deq.float() @ b_deq.float().T
434
+ bias = (torch.randn(n, device="cuda") * 0.02).to(torch.bfloat16)
435
+ residual = torch.randn((m, n), device="cuda", dtype=torch.bfloat16)
436
+ rows = []
437
+
438
+ out = torch.empty((m, n), device="cuda", dtype=torch.bfloat16)
439
+ ops.nvfp4_gemm_residual_bf16(a, b, sfa, sfb, residual, out)
440
+ expected = (matmul + residual.float()).to(torch.bfloat16)
441
+ rows.append(result_row(name, shape, "nvfp4_gemm_residual_bf16", out, expected))
442
+
443
+ ops.nvfp4_gemm_bias_gelu_bf16(a, b, sfa, sfb, bias, out)
444
+ expected_gelu = torch.nn.functional.gelu(
445
+ matmul + bias.float().view(1, -1), approximate="tanh"
446
+ ).to(torch.bfloat16)
447
+ rows.append(result_row(name, shape, "nvfp4_gemm_bias_gelu_bf16", out, expected_gelu))
448
+
449
+ out_packed, out_sfa = ops.alloc_fp4(m, n)
450
+ ops.nvfp4_gemm_bias_gelu_nvfp4(
451
+ a, b, sfa, sfb, bias, out_packed, out_sfa
452
+ )
453
+ out_deq = torch.empty((m, n), device="cuda", dtype=torch.float16)
454
+ ops.dequantize_fp4_sfa_fp16(out_packed, out_sfa, out_deq, False)
455
+ staged_packed, staged_sfa = ops.alloc_fp4(m, n)
456
+ ops.quantize_fp4_sfa_fp16(
457
+ expected_gelu.to(torch.float16), staged_packed, staged_sfa, False
458
+ )
459
+ staged_deq = torch.empty_like(out_deq)
460
+ ops.dequantize_fp4_sfa_fp16(
461
+ staged_packed, staged_sfa, staged_deq, False
462
+ )
463
+ rows.append(
464
+ result_row(
465
+ name, shape, "nvfp4_gemm_bias_gelu_nvfp4",
466
+ out_deq, staged_deq, fp4_output=True,
467
+ )
468
+ )
469
+
470
+ ops.nvfp4_gemm_streamk_bf16(a, b, sfa, sfb, out)
471
+ expected_linear = matmul.to(torch.bfloat16)
472
+ rows.append(result_row(name, shape, "nvfp4_gemm_streamk_bf16", out, expected_linear))
473
+
474
+ ops.nvfp4_gemm_streamk_bias_bf16(a, b, sfa, sfb, bias, out)
475
+ expected_bias = (matmul + bias.float().view(1, -1)).to(torch.bfloat16)
476
+ rows.append(result_row(name, shape, "nvfp4_gemm_streamk_bias_bf16", out, expected_bias))
477
+ return rows
478
+
479
+
480
+ def run_sm110_epilogue_case(ops, name: str, shape: tuple[int, int, int]):
481
+ m, n, k = shape
482
+ a, b, sfa, sfb, a_deq, b_deq = prepare_quantized_full(ops, m, n, k)
483
+ matmul = a_deq.float() @ b_deq.float().T
484
+ bias = (torch.randn(n, device="cuda") * 0.02).to(torch.bfloat16)
485
+ residual = torch.randn((m, n), device="cuda", dtype=torch.bfloat16)
486
+ rows = []
487
+
488
+ out = torch.empty_like(residual)
489
+ ops.nvfp4_gemm_bias_bf16(a, b, sfa, sfb, bias, out)
490
+ expected_bias = (matmul + bias.float().view(1, -1)).to(torch.bfloat16)
491
+ rows.append(result_row(
492
+ name, shape, "nvfp4_gemm_bias_bf16", out, expected_bias
493
+ ))
494
+
495
+ before = residual.clone()
496
+ ops.nvfp4_gemm_bias_residual_bf16(
497
+ a, b, sfa, sfb, bias, residual, residual
498
+ )
499
+ expected_residual = (
500
+ matmul + bias.float().view(1, -1) + before.float()
501
+ ).to(torch.bfloat16)
502
+ rows.append(result_row(
503
+ name, shape, "nvfp4_gemm_bias_residual_bf16",
504
+ residual, expected_residual,
505
+ ))
506
+
507
+ out_packed, out_sfa = ops.alloc_fp4(m, n)
508
+ ops.nvfp4_gemm_bias_gelu_nvfp4(
509
+ a, b, sfa, sfb, bias, out_packed, out_sfa
510
+ )
511
+ out_deq = torch.empty((m, n), device="cuda", dtype=torch.float16)
512
+ ops.dequantize_fp4_sfa_fp16(out_packed, out_sfa, out_deq, False)
513
+ expected_gelu = torch.nn.functional.gelu(
514
+ matmul + bias.float().view(1, -1), approximate="tanh"
515
+ ).to(torch.bfloat16)
516
+ staged_packed, staged_sfa = ops.alloc_fp4(m, n)
517
+ ops.quantize_fp4_sfa_bf16(
518
+ expected_gelu, staged_packed, staged_sfa, False
519
+ )
520
+ staged_deq = torch.empty_like(out_deq)
521
+ ops.dequantize_fp4_sfa_fp16(
522
+ staged_packed, staged_sfa, staged_deq, False
523
+ )
524
+ rows.append(result_row(
525
+ name, shape, "nvfp4_gemm_bias_gelu_nvfp4",
526
+ out_deq, staged_deq, fp4_output=True,
527
+ ))
528
+ return rows
529
+
530
+
531
+ def check_installed_compile(ops: InstalledOps) -> dict[str, object]:
532
+ a_packed, b_packed, sfa, sfb, _ = prepare_quantized(ops, 128, 128, 128)
533
+
534
+ def call(a, b, scale_a, scale_b):
535
+ return ops._module.nvfp4_gemm_bf16(a, b, scale_a, scale_b)
536
+
537
+ eager = call(a_packed, b_packed, sfa, sfb)
538
+ compiled = torch.compile(call, fullgraph=True)
539
+ got = compiled(a_packed, b_packed, sfa, sfb)
540
+ torch.cuda.synchronize()
541
+ max_abs = float((got.float() - eager.float()).abs().max().item())
542
+ passed = bool(
543
+ got.dtype == torch.bfloat16
544
+ and got.shape == eager.shape
545
+ and torch.equal(got, eager)
546
+ )
547
+ return {
548
+ "fullgraph": True,
549
+ "dtype": str(got.dtype),
550
+ "shape": list(got.shape),
551
+ "max_abs": max_abs,
552
+ "exact": bool(torch.equal(got, eager)),
553
+ "passed": passed,
554
+ }
555
+
556
+
557
+ def check_bf16_quantizer(ops) -> dict[str, object]:
558
+ """Require the direct BF16 producer to preserve the established layout."""
559
+ cases = [
560
+ (1, 5120, False),
561
+ (1, 6144, False),
562
+ (1, 17408, False),
563
+ (16, 2048, False),
564
+ (128, 512, False),
565
+ (64, 1024, True),
566
+ ]
567
+ rows = []
568
+ for case_index, (m, k, is_sfb) in enumerate(cases):
569
+ torch.manual_seed(8100 + case_index)
570
+ x = (torch.randn((m, k), device="cuda") * 1.5).to(torch.bfloat16)
571
+ direct_packed, direct_sfa = ops.alloc_fp4(m, k)
572
+ compat_packed, compat_sfa = ops.alloc_fp4(m, k)
573
+ # CUTLASS SFA/SFB buffers contain alignment padding that producers do
574
+ # not write or consume. Zero it so a full-buffer equality check still
575
+ # proves every mapped scale byte lands at the same address.
576
+ direct_sfa.zero_()
577
+ compat_sfa.zero_()
578
+ ops.quantize_fp4_sfa_bf16(x, direct_packed, direct_sfa, is_sfb)
579
+ ops.quantize_fp4_sfa_fp16(
580
+ x.to(torch.float16), compat_packed, compat_sfa, is_sfb
581
+ )
582
+ torch.cuda.synchronize()
583
+ packed_exact = bool(torch.equal(direct_packed, compat_packed))
584
+ sfa_exact = bool(torch.equal(direct_sfa, compat_sfa))
585
+ direct_deq = torch.empty((m, k), device="cuda", dtype=torch.float16)
586
+ compat_deq = torch.empty_like(direct_deq)
587
+ ops.dequantize_fp4_sfa_fp16(
588
+ direct_packed, direct_sfa, direct_deq, is_sfb
589
+ )
590
+ ops.dequantize_fp4_sfa_fp16(
591
+ compat_packed, compat_sfa, compat_deq, is_sfb
592
+ )
593
+ torch.cuda.synchronize()
594
+ dequant_exact = bool(torch.equal(direct_deq, compat_deq))
595
+ rows.append(
596
+ {
597
+ "shape": [m, k],
598
+ "is_sfb": is_sfb,
599
+ "packed_exact": packed_exact,
600
+ "sfa_exact": sfa_exact,
601
+ "dequant_exact": dequant_exact,
602
+ "passed": packed_exact and sfa_exact and dequant_exact,
603
+ }
604
+ )
605
+ return {"rows": rows, "passed": all(row["passed"] for row in rows)}
606
+
607
+
608
+ def main() -> int:
609
+ parser = argparse.ArgumentParser()
610
+ parser.add_argument("--backend", choices=["source", "installed"], default="source")
611
+ parser.add_argument("--artifact", default=None)
612
+ parser.add_argument("--mode", choices=sorted(MODES), default="smoke")
613
+ parser.add_argument("--json-out", default=None)
614
+ args = parser.parse_args()
615
+
616
+ if not torch.cuda.is_available():
617
+ raise RuntimeError("CUDA is required")
618
+ ops = load_source_ops() if args.backend == "source" else load_installed_ops(args.artifact)
619
+
620
+ results: list[Metrics] = []
621
+ selected_shapes = SM110_SHAPES if args.mode == "thor-models" else SHAPES
622
+ for name in MODES[args.mode]:
623
+ results.extend(run_case(ops, name, selected_shapes[name]))
624
+ capability = torch.cuda.get_device_capability(0)
625
+ if args.mode == "full" and capability != (11, 0):
626
+ for name, shape in EPILOGUE_SHAPES.items():
627
+ results.extend(run_epilogue_case(ops, name, shape))
628
+ if capability == (11, 0) and args.mode in {"full", "thor-models"}:
629
+ for name in (
630
+ "groot_n17_dit_qkv",
631
+ "groot_n17_dit_ffn_up",
632
+ "groot_n17_dit_ffn_down",
633
+ ):
634
+ results.extend(run_sm110_epilogue_case(
635
+ ops, name, SM110_SHAPES[name]
636
+ ))
637
+ compile_check = None
638
+ bf16_quantizer_check = check_bf16_quantizer(ops)
639
+ if args.backend == "installed" and args.mode == "full":
640
+ compile_check = check_installed_compile(ops)
641
+ passed = sum(1 for item in results if item.passed)
642
+ total = len(results)
643
+ if compile_check is not None:
644
+ total += 1
645
+ passed += int(bool(compile_check["passed"]))
646
+ total += 1
647
+ passed += int(bool(bf16_quantizer_check["passed"]))
648
+ payload = {
649
+ "backend": args.backend,
650
+ "mode": args.mode,
651
+ "device": torch.cuda.get_device_name(),
652
+ "torch": torch.__version__,
653
+ "passed": passed,
654
+ "total": total,
655
+ "results": [asdict(item) for item in results],
656
+ "compile_check": compile_check,
657
+ "bf16_quantizer_check": bf16_quantizer_check,
658
+ }
659
+ print(json.dumps(payload, indent=2))
660
+ if args.json_out:
661
+ out = Path(args.json_out)
662
+ out.parent.mkdir(parents=True, exist_ok=True)
663
+ out.write_text(json.dumps(payload, indent=2) + "\n")
664
+ return 0 if passed == total else 1
665
+
666
+
667
+ if __name__ == "__main__":
668
+ raise SystemExit(main())
torch-ext/fp4_gemm/__init__.py ADDED
@@ -0,0 +1,358 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """FlashRT FP4 GEMM kernels."""
2
+
3
+ from __future__ import annotations
4
+
5
+ import torch
6
+
7
+ from ._ops import add_op_namespace_prefix, ops
8
+
9
+
10
+ def sfa_size_bytes(rows: int, dim: int) -> int:
11
+ if rows <= 0 or dim <= 0 or dim % 16 != 0:
12
+ raise ValueError("rows must be positive and dim must be positive/divisible by 16")
13
+ n_blocks = dim // 16
14
+ n_row_super = (rows + 127) // 128
15
+ n_col_super = (n_blocks + 3) // 4
16
+ return n_row_super * n_col_super * 512
17
+
18
+
19
+ def _alloc_fp4(rows: int, dim: int, device: torch.device | str):
20
+ return (
21
+ torch.empty((rows, dim // 2), device=device, dtype=torch.uint8),
22
+ torch.empty((sfa_size_bytes(rows, dim),), device=device, dtype=torch.uint8),
23
+ )
24
+
25
+
26
+ @torch.library.register_fake(add_op_namespace_prefix("nvfp4_gemm_bf16"))
27
+ def _linear_fake(
28
+ a_packed: torch.Tensor,
29
+ b_packed: torch.Tensor,
30
+ sfa: torch.Tensor,
31
+ sfb: torch.Tensor,
32
+ out: torch.Tensor,
33
+ alpha: float = 1.0,
34
+ variant: int = -1,
35
+ ) -> None:
36
+ return None
37
+
38
+
39
+ @torch.library.register_fake(add_op_namespace_prefix("fp4_w4a4_gemv_warpsplit_bf16"))
40
+ def _gemv_warpsplit_fake(a_packed, b_packed, sfa, sfb, out, alpha: float = 1.0, warps: int = 4, stages: int = 4) -> None:
41
+ if a_packed.shape[0] != 1:
42
+ raise RuntimeError("warp-split GEMV serves M=1 only")
43
+ if out.shape != (1, b_packed.shape[0]):
44
+ raise RuntimeError("out must have shape (1, N)")
45
+ return None
46
+
47
+
48
+ @torch.library.register_fake(add_op_namespace_prefix("fp4_w4a16_linear_bf16"))
49
+ def _legacy_linear_fake(
50
+ a_packed: torch.Tensor,
51
+ b_packed: torch.Tensor,
52
+ sfa: torch.Tensor,
53
+ sfb: torch.Tensor,
54
+ out: torch.Tensor,
55
+ alpha: float = 1.0,
56
+ variant: int = -1,
57
+ ) -> None:
58
+ return None
59
+
60
+
61
+ @torch.library.register_fake(add_op_namespace_prefix("nvfp4_gemm_bias_bf16"))
62
+ def _bias_fake(a, b, sfa, sfb, bias, out) -> None:
63
+ return None
64
+
65
+
66
+ @torch.library.register_fake(add_op_namespace_prefix("nvfp4_gemm_bias_residual_bf16"))
67
+ def _bias_residual_fake(a, b, sfa, sfb, bias, residual, out) -> None:
68
+ return None
69
+
70
+
71
+ @torch.library.register_fake(add_op_namespace_prefix("quantize_fp4_sfa_fp16"))
72
+ def _quant_fake(x: torch.Tensor, packed: torch.Tensor, sfa: torch.Tensor, is_sfb: bool = False) -> None:
73
+ return None
74
+
75
+
76
+ @torch.library.register_fake(add_op_namespace_prefix("quantize_fp4_sfa_bf16"))
77
+ def _quant_bf16_fake(x: torch.Tensor, packed: torch.Tensor, sfa: torch.Tensor, is_sfb: bool = False) -> None:
78
+ return None
79
+
80
+
81
+ @torch.library.register_fake(add_op_namespace_prefix("dequantize_fp4_sfa_fp16"))
82
+ def _dequant_fake(packed: torch.Tensor, sfa: torch.Tensor, out: torch.Tensor, is_sfb: bool = False) -> None:
83
+ return None
84
+
85
+
86
+ @torch.library.register_fake(add_op_namespace_prefix("nvfp4_gemm_residual_bf16"))
87
+ def _residual_fake(a, b, sfa, sfb, residual, out, alpha: float = 1.0) -> None:
88
+ return None
89
+
90
+
91
+ @torch.library.register_fake(add_op_namespace_prefix("nvfp4_gemm_bias_gelu_bf16"))
92
+ def _bias_gelu_fake(a, b, sfa, sfb, bias, out, alpha: float = 1.0) -> None:
93
+ return None
94
+
95
+
96
+ @torch.library.register_fake(add_op_namespace_prefix("nvfp4_gemm_bias_gelu_nvfp4"))
97
+ def _bias_gelu_nvfp4_fake(
98
+ a, b, sfa, sfb, bias, out_packed, out_sfa, alpha: float = 1.0
99
+ ) -> None:
100
+ return None
101
+
102
+
103
+ @torch.library.register_fake(add_op_namespace_prefix("nvfp4_gemm_streamk_bf16"))
104
+ def _streamk_fake(a, b, sfa, sfb, out, alpha: float = 1.0) -> None:
105
+ return None
106
+
107
+
108
+ @torch.library.register_fake(add_op_namespace_prefix("nvfp4_gemm_streamk_bias_bf16"))
109
+ def _streamk_bias_fake(a, b, sfa, sfb, bias, out, alpha: float = 1.0) -> None:
110
+ return None
111
+
112
+
113
+ def quantize_fp4_sfa_fp16(
114
+ x: torch.Tensor,
115
+ packed: torch.Tensor | None = None,
116
+ sfa: torch.Tensor | None = None,
117
+ is_sfb: bool = False,
118
+ ):
119
+ if packed is None or sfa is None:
120
+ packed, sfa = _alloc_fp4(x.shape[0], x.shape[1], x.device)
121
+ ops.quantize_fp4_sfa_fp16(x, packed, sfa, bool(is_sfb))
122
+ return packed, sfa
123
+
124
+
125
+ def quantize_fp4_sfa_bf16(
126
+ x: torch.Tensor,
127
+ packed: torch.Tensor | None = None,
128
+ sfa: torch.Tensor | None = None,
129
+ is_sfb: bool = False,
130
+ ):
131
+ """Quantize BF16 directly to packed E2M1 and CUTLASS SFA/SFB."""
132
+ if packed is None or sfa is None:
133
+ packed, sfa = _alloc_fp4(x.shape[0], x.shape[1], x.device)
134
+ ops.quantize_fp4_sfa_bf16(x, packed, sfa, bool(is_sfb))
135
+ return packed, sfa
136
+
137
+
138
+ def dequantize_fp4_sfa_fp16(
139
+ packed: torch.Tensor,
140
+ sfa: torch.Tensor,
141
+ out: torch.Tensor | None = None,
142
+ is_sfb: bool = False,
143
+ ) -> torch.Tensor:
144
+ if out is None:
145
+ out = torch.empty((packed.shape[0], packed.shape[1] * 2), device=packed.device, dtype=torch.float16)
146
+ ops.dequantize_fp4_sfa_fp16(packed, sfa, out, bool(is_sfb))
147
+ return out
148
+
149
+
150
+ def nvfp4_gemm_bf16(
151
+ a_packed: torch.Tensor,
152
+ b_packed: torch.Tensor,
153
+ sfa: torch.Tensor,
154
+ sfb: torch.Tensor,
155
+ alpha: float = 1.0,
156
+ out: torch.Tensor | None = None,
157
+ variant: int = -1,
158
+ ) -> torch.Tensor:
159
+ if out is None:
160
+ out = torch.empty((a_packed.shape[0], b_packed.shape[0]), device=a_packed.device, dtype=torch.bfloat16)
161
+ ops.nvfp4_gemm_bf16(a_packed, b_packed, sfa, sfb, out, float(alpha), int(variant))
162
+ return out
163
+
164
+
165
+ def fp4_w4a4_gemv_warpsplit_bf16(
166
+ a_packed: torch.Tensor,
167
+ b_packed: torch.Tensor,
168
+ sfa: torch.Tensor,
169
+ sfb: torch.Tensor,
170
+ *,
171
+ alpha: float = 1.0,
172
+ warps: int = 4,
173
+ stages: int = 4,
174
+ out: Optional[torch.Tensor] = None,
175
+ ) -> torch.Tensor:
176
+ """Warp-split-K NVFP4 W4A4 GEMV for the M=1 decode row (SM120).
177
+
178
+ Splits K across warps inside one block with a shared-memory reduce -
179
+ no cross-block intermediate, so it stays safe under CUDA-graph
180
+ replay - and fills the SMs the tiled GEMM underfills at long-K
181
+ small-M decode shapes. Same packed/scale layouts as the linear
182
+ entry points."""
183
+ if out is None:
184
+ out = torch.empty((1, b_packed.shape[0]), device=a_packed.device, dtype=torch.bfloat16)
185
+ ops.fp4_w4a4_gemv_warpsplit_bf16(a_packed, b_packed, sfa, sfb, out, float(alpha), int(warps), int(stages))
186
+ return out
187
+
188
+
189
+ def fp4_w4a16_linear_bf16(
190
+ a_packed: torch.Tensor,
191
+ b_packed: torch.Tensor,
192
+ sfa: torch.Tensor,
193
+ sfb: torch.Tensor,
194
+ alpha: float = 1.0,
195
+ out: torch.Tensor | None = None,
196
+ variant: int = -1,
197
+ ) -> torch.Tensor:
198
+ """Compatibility alias for :func:`nvfp4_gemm_bf16`."""
199
+ return nvfp4_gemm_bf16(
200
+ a_packed, b_packed, sfa, sfb, alpha=alpha, out=out, variant=variant
201
+ )
202
+
203
+
204
+ def nvfp4_gemm_bias_bf16(
205
+ a_packed: torch.Tensor,
206
+ b_packed: torch.Tensor,
207
+ sfa: torch.Tensor,
208
+ sfb: torch.Tensor,
209
+ bias: torch.Tensor,
210
+ *,
211
+ out: torch.Tensor | None = None,
212
+ ) -> torch.Tensor:
213
+ """SM110 NVFP4 GEMM with a fused per-column BF16 bias."""
214
+ if out is None:
215
+ out = torch.empty(
216
+ (a_packed.shape[0], b_packed.shape[0]),
217
+ device=a_packed.device,
218
+ dtype=torch.bfloat16,
219
+ )
220
+ ops.nvfp4_gemm_bias_bf16(a_packed, b_packed, sfa, sfb, bias, out)
221
+ return out
222
+
223
+
224
+ def nvfp4_gemm_bias_residual_bf16(
225
+ a_packed: torch.Tensor,
226
+ b_packed: torch.Tensor,
227
+ sfa: torch.Tensor,
228
+ sfb: torch.Tensor,
229
+ bias: torch.Tensor,
230
+ residual: torch.Tensor,
231
+ *,
232
+ out: torch.Tensor | None = None,
233
+ ) -> torch.Tensor:
234
+ """SM110 NVFP4 GEMM with fused BF16 bias and residual add."""
235
+ if out is None:
236
+ out = torch.empty_like(residual)
237
+ ops.nvfp4_gemm_bias_residual_bf16(
238
+ a_packed, b_packed, sfa, sfb, bias, residual, out
239
+ )
240
+ return out
241
+
242
+
243
+ def nvfp4_gemm_residual_bf16(
244
+ a_packed: torch.Tensor,
245
+ b_packed: torch.Tensor,
246
+ sfa: torch.Tensor,
247
+ sfb: torch.Tensor,
248
+ residual: torch.Tensor,
249
+ alpha: float = 1.0,
250
+ out: torch.Tensor | None = None,
251
+ ) -> torch.Tensor:
252
+ if out is None:
253
+ out = torch.empty_like(residual)
254
+ ops.nvfp4_gemm_residual_bf16(
255
+ a_packed, b_packed, sfa, sfb, residual, out, float(alpha)
256
+ )
257
+ return out
258
+
259
+
260
+ def nvfp4_gemm_bias_gelu_bf16(
261
+ a_packed: torch.Tensor,
262
+ b_packed: torch.Tensor,
263
+ sfa: torch.Tensor,
264
+ sfb: torch.Tensor,
265
+ bias: torch.Tensor,
266
+ alpha: float = 1.0,
267
+ out: torch.Tensor | None = None,
268
+ ) -> torch.Tensor:
269
+ if out is None:
270
+ out = torch.empty(
271
+ (a_packed.shape[0], b_packed.shape[0]),
272
+ device=a_packed.device,
273
+ dtype=torch.bfloat16,
274
+ )
275
+ ops.nvfp4_gemm_bias_gelu_bf16(
276
+ a_packed, b_packed, sfa, sfb, bias, out, float(alpha)
277
+ )
278
+ return out
279
+
280
+
281
+ def nvfp4_gemm_bias_gelu_nvfp4(
282
+ a_packed: torch.Tensor,
283
+ b_packed: torch.Tensor,
284
+ sfa: torch.Tensor,
285
+ sfb: torch.Tensor,
286
+ bias: torch.Tensor,
287
+ alpha: float = 1.0,
288
+ out_packed: torch.Tensor | None = None,
289
+ out_sfa: torch.Tensor | None = None,
290
+ ) -> tuple[torch.Tensor, torch.Tensor]:
291
+ m, n = a_packed.shape[0], b_packed.shape[0]
292
+ if out_packed is None:
293
+ out_packed = torch.empty((m, n // 2), device=a_packed.device, dtype=torch.uint8)
294
+ if out_sfa is None:
295
+ out_sfa = torch.empty((sfa_size_bytes(m, n),), device=a_packed.device, dtype=torch.uint8)
296
+ ops.nvfp4_gemm_bias_gelu_nvfp4(
297
+ a_packed, b_packed, sfa, sfb, bias, out_packed, out_sfa, float(alpha)
298
+ )
299
+ return out_packed, out_sfa
300
+
301
+
302
+ def nvfp4_gemm_streamk_bf16(
303
+ a_packed: torch.Tensor,
304
+ b_packed: torch.Tensor,
305
+ sfa: torch.Tensor,
306
+ sfb: torch.Tensor,
307
+ alpha: float = 1.0,
308
+ out: torch.Tensor | None = None,
309
+ ) -> torch.Tensor:
310
+ if out is None:
311
+ out = torch.empty(
312
+ (a_packed.shape[0], b_packed.shape[0]),
313
+ device=a_packed.device,
314
+ dtype=torch.bfloat16,
315
+ )
316
+ ops.nvfp4_gemm_streamk_bf16(
317
+ a_packed, b_packed, sfa, sfb, out, float(alpha)
318
+ )
319
+ return out
320
+
321
+
322
+ def nvfp4_gemm_streamk_bias_bf16(
323
+ a_packed: torch.Tensor,
324
+ b_packed: torch.Tensor,
325
+ sfa: torch.Tensor,
326
+ sfb: torch.Tensor,
327
+ bias: torch.Tensor,
328
+ alpha: float = 1.0,
329
+ out: torch.Tensor | None = None,
330
+ ) -> torch.Tensor:
331
+ if out is None:
332
+ out = torch.empty(
333
+ (a_packed.shape[0], b_packed.shape[0]),
334
+ device=a_packed.device,
335
+ dtype=torch.bfloat16,
336
+ )
337
+ ops.nvfp4_gemm_streamk_bias_bf16(
338
+ a_packed, b_packed, sfa, sfb, bias, out, float(alpha)
339
+ )
340
+ return out
341
+
342
+
343
+ __all__ = [
344
+ "dequantize_fp4_sfa_fp16",
345
+ "fp4_w4a16_linear_bf16",
346
+ "fp4_w4a4_gemv_warpsplit_bf16",
347
+ "nvfp4_gemm_bf16",
348
+ "nvfp4_gemm_bias_bf16",
349
+ "nvfp4_gemm_bias_gelu_bf16",
350
+ "nvfp4_gemm_bias_gelu_nvfp4",
351
+ "nvfp4_gemm_bias_residual_bf16",
352
+ "nvfp4_gemm_residual_bf16",
353
+ "nvfp4_gemm_streamk_bf16",
354
+ "nvfp4_gemm_streamk_bias_bf16",
355
+ "quantize_fp4_sfa_fp16",
356
+ "quantize_fp4_sfa_bf16",
357
+ "sfa_size_bytes",
358
+ ]
torch-ext/torch_binding.cpp ADDED
@@ -0,0 +1,606 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ // SPDX-License-Identifier: Apache-2.0
2
+
3
+ #include <torch/all.h>
4
+ #include <torch/library.h>
5
+
6
+ #include <limits>
7
+
8
+ #if defined(CUDA_KERNEL)
9
+ #include <ATen/cuda/CUDAContext.h>
10
+ #include <c10/cuda/CUDAGuard.h>
11
+ #endif
12
+
13
+ #include "dequantize_fp4_sfa.cuh"
14
+ #if !defined(FLASHRT_FP4_GEMM_SOURCE_SM110_ONLY)
15
+ #include "gemm/fp4/cutlass_nvfp4_gemm_bias_gelu_bf16out_sm120.cuh"
16
+ #include "gemm/fp4/cutlass_nvfp4_gemm_bias_gelu_fp4out_sm120.cuh"
17
+ #include "gemm/fp4/cutlass_nvfp4_gemm_dn_streamk_bias_sm120.cuh"
18
+ #include "gemm/fp4/cutlass_nvfp4_w4a16_gemm_sm120.cuh"
19
+ #include "gemm/fp4/fp4_w4a4_mma_warpsplit_sm120.cuh"
20
+ #endif
21
+ #include "gemm/fp4/sm110_dispatch.cuh"
22
+ #include "quantize/quantize_fp4_sfa.cuh"
23
+ #include "registration.h"
24
+ #include "torch_binding.h"
25
+
26
+ flash_rt::hub::Sm110GemmDispatch flash_rt::hub::sm110_gemm_dispatch = nullptr;
27
+ flash_rt::hub::Sm110GemmBiasDispatch
28
+ flash_rt::hub::sm110_gemm_bias_dispatch = nullptr;
29
+ flash_rt::hub::Sm110GemmBiasResidualDispatch
30
+ flash_rt::hub::sm110_gemm_bias_residual_dispatch = nullptr;
31
+ flash_rt::hub::Sm110GemmBiasGeluFp4Dispatch
32
+ flash_rt::hub::sm110_gemm_bias_gelu_fp4_dispatch = nullptr;
33
+ flash_rt::hub::Sm110QuantizeBf16Dispatch
34
+ flash_rt::hub::sm110_quantize_bf16_dispatch = nullptr;
35
+
36
+ namespace {
37
+
38
+ void check_cuda_contiguous(torch::Tensor const& tensor, const char* name) {
39
+ TORCH_CHECK(tensor.is_cuda(), name, " must be a CUDA tensor");
40
+ TORCH_CHECK(tensor.is_contiguous(), name, " must be contiguous");
41
+ }
42
+
43
+ void check_uint8_cuda(torch::Tensor const& tensor, const char* name) {
44
+ check_cuda_contiguous(tensor, name);
45
+ TORCH_CHECK(tensor.scalar_type() == torch::kUInt8,
46
+ name, " must have dtype torch.uint8");
47
+ }
48
+
49
+ void check_fp16_cuda(torch::Tensor const& tensor, const char* name) {
50
+ check_cuda_contiguous(tensor, name);
51
+ TORCH_CHECK(tensor.scalar_type() == torch::kFloat16,
52
+ name, " must have dtype torch.float16");
53
+ }
54
+
55
+ void check_bf16_cuda(torch::Tensor const& tensor, const char* name) {
56
+ check_cuda_contiguous(tensor, name);
57
+ TORCH_CHECK(tensor.scalar_type() == torch::kBFloat16,
58
+ name, " must have dtype torch.bfloat16");
59
+ }
60
+
61
+ int checked_int(int64_t value, const char* name) {
62
+ TORCH_CHECK(value > 0 && value <= std::numeric_limits<int>::max(),
63
+ name, " must fit in positive int");
64
+ return static_cast<int>(value);
65
+ }
66
+
67
+ int64_t swizzled_bytes(int64_t rows, int64_t dim) {
68
+ TORCH_CHECK(rows > 0 && dim > 0 && dim % 16 == 0,
69
+ "rows must be positive and dim must be positive/divisible by 16");
70
+ const int64_t n_blocks = dim / 16;
71
+ const int64_t n_row_super = (rows + 127) / 128;
72
+ const int64_t n_col_super = (n_blocks + 3) / 4;
73
+ return n_row_super * n_col_super * 512;
74
+ }
75
+
76
+ void check_same_device(torch::Tensor const& a, torch::Tensor const& b,
77
+ const char* a_name, const char* b_name) {
78
+ TORCH_CHECK(a.get_device() == b.get_device(),
79
+ a_name, " and ", b_name, " must be on the same CUDA device");
80
+ }
81
+
82
+ struct GemmShape {
83
+ int64_t m;
84
+ int64_t n;
85
+ int64_t k;
86
+ };
87
+
88
+ #if defined(CUDA_KERNEL)
89
+ cudaDeviceProp const* current_device_properties(torch::Tensor const& anchor) {
90
+ return at::cuda::getDeviceProperties(anchor.get_device());
91
+ }
92
+
93
+ void require_sm120(torch::Tensor const& anchor, const char* operation) {
94
+ auto const* props = current_device_properties(anchor);
95
+ TORCH_CHECK(props->major == 12 && props->minor == 0,
96
+ operation, " is an SM120 fused epilogue; got SM",
97
+ props->major, props->minor,
98
+ ". On SM110 use nvfp4_gemm_bf16 with fp4-fused-ops producers.");
99
+ }
100
+ #endif
101
+
102
+ GemmShape check_fp4_gemm_inputs(
103
+ torch::Tensor const& a_packed,
104
+ torch::Tensor const& b_packed,
105
+ torch::Tensor const& sfa,
106
+ torch::Tensor const& sfb) {
107
+ check_uint8_cuda(a_packed, "a_packed");
108
+ check_uint8_cuda(b_packed, "b_packed");
109
+ check_uint8_cuda(sfa, "sfa");
110
+ check_uint8_cuda(sfb, "sfb");
111
+ TORCH_CHECK(a_packed.dim() == 2, "a_packed must have shape (M, K / 2)");
112
+ TORCH_CHECK(b_packed.dim() == 2, "b_packed must have shape (N, K / 2)");
113
+ const int64_t m = a_packed.size(0);
114
+ const int64_t n = b_packed.size(0);
115
+ const int64_t k_half = a_packed.size(1);
116
+ TORCH_CHECK(m > 0 && n > 0 && k_half > 0, "M, N, and K must be positive");
117
+ TORCH_CHECK(b_packed.size(1) == k_half,
118
+ "a_packed and b_packed must have the same K / 2 dimension");
119
+ const int64_t k = k_half * 2;
120
+ TORCH_CHECK(k % 16 == 0, "K must be divisible by 16");
121
+ TORCH_CHECK(sfa.numel() >= swizzled_bytes(m, k),
122
+ "sfa is too small for CUTLASS SFA layout");
123
+ TORCH_CHECK(sfb.numel() >= swizzled_bytes(n, k),
124
+ "sfb is too small for CUTLASS SFB layout");
125
+ check_same_device(a_packed, b_packed, "a_packed", "b_packed");
126
+ check_same_device(a_packed, sfa, "a_packed", "sfa");
127
+ check_same_device(a_packed, sfb, "a_packed", "sfb");
128
+ return {m, n, k};
129
+ }
130
+
131
+ int select_sm110_variant(GemmShape const& shape, int64_t requested) {
132
+ if (requested >= 0) return static_cast<int>(requested);
133
+ if (shape.n >= 4 * shape.k) return 1;
134
+ if (shape.n == 3 * shape.k) return 2;
135
+ return 0;
136
+ }
137
+
138
+ } // namespace
139
+
140
+ void fp4_w4a4_gemv_warpsplit_bf16(
141
+ torch::Tensor const& a_packed,
142
+ torch::Tensor const& b_packed,
143
+ torch::Tensor const& sfa,
144
+ torch::Tensor const& sfb,
145
+ torch::Tensor& out,
146
+ double alpha,
147
+ int64_t warps,
148
+ int64_t stages) {
149
+ auto shape = check_fp4_gemm_inputs(a_packed, b_packed, sfa, sfb);
150
+ check_bf16_cuda(out, "out");
151
+ TORCH_CHECK(shape.m == 1,
152
+ "warp-split GEMV serves the M=1 decode row only");
153
+ TORCH_CHECK(out.sizes() == torch::IntArrayRef({shape.m, shape.n}),
154
+ "out must have shape (1, N)");
155
+ TORCH_CHECK(warps == 2 || warps == 4 || warps == 8,
156
+ "warps must be 2, 4 or 8");
157
+ TORCH_CHECK(stages == 3 || stages == 4 || stages == 6,
158
+ "stages must be 3, 4 or 6");
159
+ TORCH_CHECK(shape.n % 8 == 0, "N must be a multiple of 8");
160
+ TORCH_CHECK(shape.k % 64 == 0 && (shape.k / 64) % warps == 0,
161
+ "K must be a multiple of 64*warps");
162
+ check_same_device(a_packed, out, "a_packed", "out");
163
+ #if defined(CUDA_KERNEL)
164
+ at::cuda::CUDAGuard device_guard(a_packed.device());
165
+ auto const* props = current_device_properties(a_packed);
166
+ TORCH_CHECK(props->major == 12 && props->minor == 0,
167
+ "the warp-split GEMV is an SM120 kernel; got SM",
168
+ props->major, props->minor);
169
+ #if defined(FLASHRT_FP4_GEMM_SOURCE_SM110_ONLY)
170
+ TORCH_CHECK(false, "SM120 FP4 GEMM source is not present in this build");
171
+ #else
172
+ auto stream = at::cuda::getCurrentCUDAStream(a_packed.get_device()).stream();
173
+ const int rc = flash_rt::gemm::fp4_w4a4_mma_sm120_warpsplit_bf16out(
174
+ a_packed.data_ptr(), b_packed.data_ptr(), out.data_ptr(),
175
+ checked_int(shape.n, "N"), checked_int(shape.k, "K"),
176
+ sfa.data_ptr(), sfb.data_ptr(), static_cast<float>(alpha),
177
+ static_cast<int>(warps), static_cast<int>(stages), stream);
178
+ TORCH_CHECK(rc == 0, "fp4_w4a4_gemv_warpsplit_bf16 failed with rc=", rc);
179
+ #endif
180
+ #else
181
+ TORCH_CHECK(false, "fp4-gemm was not built with CUDA support");
182
+ #endif
183
+ }
184
+
185
+ void fp4_w4a16_linear_bf16(
186
+ torch::Tensor const& a_packed,
187
+ torch::Tensor const& b_packed,
188
+ torch::Tensor const& sfa,
189
+ torch::Tensor const& sfb,
190
+ torch::Tensor& out,
191
+ double alpha,
192
+ int64_t variant) {
193
+ auto shape = check_fp4_gemm_inputs(a_packed, b_packed, sfa, sfb);
194
+ check_bf16_cuda(out, "out");
195
+ TORCH_CHECK(out.sizes() == torch::IntArrayRef({shape.m, shape.n}),
196
+ "out must have shape (M, N)");
197
+ TORCH_CHECK(variant >= -1 && variant <= 2,
198
+ "variant must be -1(auto), 0(default), 1(widen), or 2(pingpong)");
199
+ check_same_device(a_packed, out, "a_packed", "out");
200
+ #if defined(CUDA_KERNEL)
201
+ at::cuda::CUDAGuard device_guard(a_packed.device());
202
+ auto const* props = current_device_properties(a_packed);
203
+ TORCH_CHECK((props->major == 11 && props->minor == 0) ||
204
+ (props->major == 12 && props->minor == 0),
205
+ "nvfp4_gemm_bf16 requires SM110 or SM120; got SM",
206
+ props->major, props->minor);
207
+ auto stream = at::cuda::getCurrentCUDAStream(a_packed.get_device()).stream();
208
+ if (props->major == 11) {
209
+ variant = select_sm110_variant(shape, variant);
210
+ TORCH_CHECK(flash_rt::hub::sm110_gemm_dispatch != nullptr,
211
+ "SM110 FP4 GEMM source is not present in this build");
212
+ flash_rt::hub::sm110_gemm_dispatch(
213
+ a_packed.data_ptr(), b_packed.data_ptr(), out.data_ptr(),
214
+ checked_int(shape.m, "M"), checked_int(shape.n, "N"),
215
+ checked_int(shape.k, "K"), sfa.data_ptr(), sfb.data_ptr(),
216
+ static_cast<float>(alpha), variant, stream);
217
+ } else {
218
+ #if defined(FLASHRT_FP4_GEMM_SOURCE_SM110_ONLY)
219
+ TORCH_CHECK(false, "SM120 FP4 GEMM source is not present in this build");
220
+ #else
221
+ if (variant == 1) {
222
+ flash_rt::gemm::fp4_w4a16_gemm_sm120_bf16out_widen(
223
+ a_packed.data_ptr(), b_packed.data_ptr(), out.data_ptr(),
224
+ checked_int(shape.m, "M"), checked_int(shape.n, "N"), checked_int(shape.k, "K"),
225
+ sfa.data_ptr(), sfb.data_ptr(), static_cast<float>(alpha), stream);
226
+ } else if (variant == 2) {
227
+ flash_rt::gemm::fp4_w4a16_gemm_sm120_bf16out_pingpong(
228
+ a_packed.data_ptr(), b_packed.data_ptr(), out.data_ptr(),
229
+ checked_int(shape.m, "M"), checked_int(shape.n, "N"), checked_int(shape.k, "K"),
230
+ sfa.data_ptr(), sfb.data_ptr(), static_cast<float>(alpha), stream);
231
+ } else {
232
+ flash_rt::gemm::fp4_w4a16_gemm_sm120_bf16out(
233
+ a_packed.data_ptr(), b_packed.data_ptr(), out.data_ptr(),
234
+ checked_int(shape.m, "M"), checked_int(shape.n, "N"), checked_int(shape.k, "K"),
235
+ sfa.data_ptr(), sfb.data_ptr(), static_cast<float>(alpha), stream);
236
+ }
237
+ #endif
238
+ }
239
+ #endif
240
+ }
241
+
242
+ void nvfp4_gemm_bias_bf16(
243
+ torch::Tensor const& a_packed,
244
+ torch::Tensor const& b_packed,
245
+ torch::Tensor const& sfa,
246
+ torch::Tensor const& sfb,
247
+ torch::Tensor const& bias,
248
+ torch::Tensor& out) {
249
+ auto shape = check_fp4_gemm_inputs(a_packed, b_packed, sfa, sfb);
250
+ check_bf16_cuda(bias, "bias");
251
+ check_bf16_cuda(out, "out");
252
+ TORCH_CHECK(bias.dim() == 1 && bias.numel() == shape.n,
253
+ "bias must have shape (N,)");
254
+ TORCH_CHECK(out.sizes() == torch::IntArrayRef({shape.m, shape.n}),
255
+ "out must have shape (M, N)");
256
+ check_same_device(a_packed, bias, "a_packed", "bias");
257
+ check_same_device(a_packed, out, "a_packed", "out");
258
+ #if defined(CUDA_KERNEL)
259
+ at::cuda::CUDAGuard device_guard(a_packed.device());
260
+ auto const* props = current_device_properties(a_packed);
261
+ TORCH_CHECK(props->major == 11 && props->minor == 0,
262
+ "nvfp4_gemm_bias_bf16 currently requires SM110; got SM",
263
+ props->major, props->minor);
264
+ TORCH_CHECK(flash_rt::hub::sm110_gemm_bias_dispatch != nullptr,
265
+ "SM110 fused-bias FP4 GEMM source is not present in this build");
266
+ auto stream = at::cuda::getCurrentCUDAStream(a_packed.get_device()).stream();
267
+ const int rc = flash_rt::hub::sm110_gemm_bias_dispatch(
268
+ a_packed.data_ptr(), sfa.data_ptr(), b_packed.data_ptr(), sfb.data_ptr(),
269
+ bias.data_ptr(), out.data_ptr(), checked_int(shape.m, "M"),
270
+ checked_int(shape.n, "N"), checked_int(shape.k, "K"), stream);
271
+ TORCH_CHECK(rc == 0, "nvfp4_gemm_bias_bf16 failed with rc=", rc);
272
+ #endif
273
+ }
274
+
275
+ void nvfp4_gemm_bias_residual_bf16(
276
+ torch::Tensor const& a_packed,
277
+ torch::Tensor const& b_packed,
278
+ torch::Tensor const& sfa,
279
+ torch::Tensor const& sfb,
280
+ torch::Tensor const& bias,
281
+ torch::Tensor const& residual,
282
+ torch::Tensor& out) {
283
+ auto shape = check_fp4_gemm_inputs(a_packed, b_packed, sfa, sfb);
284
+ check_bf16_cuda(bias, "bias");
285
+ check_bf16_cuda(residual, "residual");
286
+ check_bf16_cuda(out, "out");
287
+ TORCH_CHECK(bias.dim() == 1 && bias.numel() == shape.n,
288
+ "bias must have shape (N,)");
289
+ TORCH_CHECK(residual.sizes() == torch::IntArrayRef({shape.m, shape.n}),
290
+ "residual must have shape (M, N)");
291
+ TORCH_CHECK(out.sizes() == residual.sizes(), "out must match residual");
292
+ check_same_device(a_packed, bias, "a_packed", "bias");
293
+ check_same_device(a_packed, residual, "a_packed", "residual");
294
+ check_same_device(a_packed, out, "a_packed", "out");
295
+ #if defined(CUDA_KERNEL)
296
+ at::cuda::CUDAGuard device_guard(a_packed.device());
297
+ auto const* props = current_device_properties(a_packed);
298
+ TORCH_CHECK(props->major == 11 && props->minor == 0,
299
+ "nvfp4_gemm_bias_residual_bf16 currently requires SM110; got SM",
300
+ props->major, props->minor);
301
+ TORCH_CHECK(flash_rt::hub::sm110_gemm_bias_residual_dispatch != nullptr,
302
+ "SM110 bias-residual FP4 GEMM source is not present in this build");
303
+ auto stream = at::cuda::getCurrentCUDAStream(a_packed.get_device()).stream();
304
+ const int rc = flash_rt::hub::sm110_gemm_bias_residual_dispatch(
305
+ a_packed.data_ptr(), sfa.data_ptr(), b_packed.data_ptr(), sfb.data_ptr(),
306
+ bias.data_ptr(), residual.data_ptr(), out.data_ptr(),
307
+ checked_int(shape.m, "M"), checked_int(shape.n, "N"),
308
+ checked_int(shape.k, "K"), stream);
309
+ TORCH_CHECK(rc == 0, "nvfp4_gemm_bias_residual_bf16 failed with rc=", rc);
310
+ #endif
311
+ }
312
+
313
+ void nvfp4_gemm_residual_bf16(
314
+ torch::Tensor const& a_packed,
315
+ torch::Tensor const& b_packed,
316
+ torch::Tensor const& sfa,
317
+ torch::Tensor const& sfb,
318
+ torch::Tensor const& residual,
319
+ torch::Tensor& out,
320
+ double alpha) {
321
+ auto shape = check_fp4_gemm_inputs(a_packed, b_packed, sfa, sfb);
322
+ check_bf16_cuda(residual, "residual");
323
+ check_bf16_cuda(out, "out");
324
+ TORCH_CHECK(residual.sizes() == torch::IntArrayRef({shape.m, shape.n}),
325
+ "residual must have shape (M, N)");
326
+ TORCH_CHECK(out.sizes() == residual.sizes(), "out must match residual");
327
+ check_same_device(a_packed, residual, "a_packed", "residual");
328
+ check_same_device(a_packed, out, "a_packed", "out");
329
+ #if defined(CUDA_KERNEL)
330
+ at::cuda::CUDAGuard device_guard(a_packed.device());
331
+ require_sm120(a_packed, "nvfp4_gemm_residual_bf16");
332
+ auto stream = at::cuda::getCurrentCUDAStream(a_packed.get_device()).stream();
333
+ #if !defined(FLASHRT_FP4_GEMM_SOURCE_SM110_ONLY)
334
+ flash_rt::gemm::fp4_w4a16_gemm_residual_sm120_bf16out(
335
+ a_packed.data_ptr(), b_packed.data_ptr(), residual.data_ptr(),
336
+ out.data_ptr(), checked_int(shape.m, "M"), checked_int(shape.n, "N"),
337
+ checked_int(shape.k, "K"), sfa.data_ptr(), sfb.data_ptr(),
338
+ static_cast<float>(alpha), stream);
339
+ #endif
340
+ #endif
341
+ }
342
+
343
+ void nvfp4_gemm_bias_gelu_bf16(
344
+ torch::Tensor const& a_packed,
345
+ torch::Tensor const& b_packed,
346
+ torch::Tensor const& sfa,
347
+ torch::Tensor const& sfb,
348
+ torch::Tensor const& bias,
349
+ torch::Tensor& out,
350
+ double alpha) {
351
+ auto shape = check_fp4_gemm_inputs(a_packed, b_packed, sfa, sfb);
352
+ check_bf16_cuda(bias, "bias");
353
+ check_bf16_cuda(out, "out");
354
+ TORCH_CHECK(bias.dim() == 1 && bias.numel() == shape.n,
355
+ "bias must have shape (N,)");
356
+ TORCH_CHECK(out.sizes() == torch::IntArrayRef({shape.m, shape.n}),
357
+ "out must have shape (M, N)");
358
+ check_same_device(a_packed, bias, "a_packed", "bias");
359
+ check_same_device(a_packed, out, "a_packed", "out");
360
+ #if defined(CUDA_KERNEL)
361
+ at::cuda::CUDAGuard device_guard(a_packed.device());
362
+ require_sm120(a_packed, "nvfp4_gemm_bias_gelu_bf16");
363
+ auto stream = at::cuda::getCurrentCUDAStream(a_packed.get_device()).stream();
364
+ #if !defined(FLASHRT_FP4_GEMM_SOURCE_SM110_ONLY)
365
+ flash_rt::gemm::fp4_w4a16_gemm_bias_gelu_bf16out_sm120(
366
+ a_packed.data_ptr(), b_packed.data_ptr(), sfa.data_ptr(), sfb.data_ptr(),
367
+ bias.data_ptr(), out.data_ptr(), checked_int(shape.m, "M"),
368
+ checked_int(shape.n, "N"), checked_int(shape.k, "K"),
369
+ static_cast<float>(alpha), stream);
370
+ #endif
371
+ #endif
372
+ }
373
+
374
+ void nvfp4_gemm_bias_gelu_nvfp4(
375
+ torch::Tensor const& a_packed,
376
+ torch::Tensor const& b_packed,
377
+ torch::Tensor const& sfa,
378
+ torch::Tensor const& sfb,
379
+ torch::Tensor const& bias,
380
+ torch::Tensor& out_packed,
381
+ torch::Tensor& out_sfa,
382
+ double alpha) {
383
+ auto shape = check_fp4_gemm_inputs(a_packed, b_packed, sfa, sfb);
384
+ check_bf16_cuda(bias, "bias");
385
+ check_uint8_cuda(out_packed, "out_packed");
386
+ check_uint8_cuda(out_sfa, "out_sfa");
387
+ TORCH_CHECK(bias.dim() == 1 && bias.numel() == shape.n,
388
+ "bias must have shape (N,)");
389
+ TORCH_CHECK(shape.n % 2 == 0 &&
390
+ out_packed.sizes() ==
391
+ torch::IntArrayRef({shape.m, shape.n / 2}),
392
+ "out_packed must have shape (M, N / 2)");
393
+ TORCH_CHECK(out_sfa.numel() >= swizzled_bytes(shape.m, shape.n),
394
+ "out_sfa is too small for output scale layout");
395
+ check_same_device(a_packed, bias, "a_packed", "bias");
396
+ check_same_device(a_packed, out_packed, "a_packed", "out_packed");
397
+ check_same_device(a_packed, out_sfa, "a_packed", "out_sfa");
398
+ #if defined(CUDA_KERNEL)
399
+ at::cuda::CUDAGuard device_guard(a_packed.device());
400
+ auto const* props = current_device_properties(a_packed);
401
+ auto stream = at::cuda::getCurrentCUDAStream(a_packed.get_device()).stream();
402
+ if (props->major == 11 && props->minor == 0) {
403
+ TORCH_CHECK(flash_rt::hub::sm110_gemm_bias_gelu_fp4_dispatch != nullptr,
404
+ "SM110 bias-GELU-FP4 GEMM source is not present in this build");
405
+ const int rc = flash_rt::hub::sm110_gemm_bias_gelu_fp4_dispatch(
406
+ a_packed.data_ptr(), sfa.data_ptr(), b_packed.data_ptr(), sfb.data_ptr(),
407
+ bias.data_ptr(), out_packed.data_ptr(), out_sfa.data_ptr(),
408
+ checked_int(shape.m, "M"), checked_int(shape.n, "N"),
409
+ checked_int(shape.k, "K"), stream);
410
+ TORCH_CHECK(rc == 0, "nvfp4_gemm_bias_gelu_nvfp4 failed with rc=", rc);
411
+ return;
412
+ }
413
+ require_sm120(a_packed, "nvfp4_gemm_bias_gelu_nvfp4");
414
+ #if !defined(FLASHRT_FP4_GEMM_SOURCE_SM110_ONLY)
415
+ flash_rt::gemm::fp4_w4a16_gemm_bias_gelu_fp4out_sm120(
416
+ a_packed.data_ptr(), b_packed.data_ptr(), sfa.data_ptr(), sfb.data_ptr(),
417
+ bias.data_ptr(), out_packed.data_ptr(), out_sfa.data_ptr(),
418
+ checked_int(shape.m, "M"), checked_int(shape.n, "N"),
419
+ checked_int(shape.k, "K"), static_cast<float>(alpha), stream);
420
+ #endif
421
+ #endif
422
+ }
423
+
424
+ void nvfp4_gemm_streamk_bf16(
425
+ torch::Tensor const& a_packed,
426
+ torch::Tensor const& b_packed,
427
+ torch::Tensor const& sfa,
428
+ torch::Tensor const& sfb,
429
+ torch::Tensor& out,
430
+ double alpha) {
431
+ auto shape = check_fp4_gemm_inputs(a_packed, b_packed, sfa, sfb);
432
+ check_bf16_cuda(out, "out");
433
+ TORCH_CHECK(out.sizes() == torch::IntArrayRef({shape.m, shape.n}),
434
+ "out must have shape (M, N)");
435
+ check_same_device(a_packed, out, "a_packed", "out");
436
+ #if defined(CUDA_KERNEL)
437
+ at::cuda::CUDAGuard device_guard(a_packed.device());
438
+ require_sm120(a_packed, "nvfp4_gemm_streamk_bf16");
439
+ auto stream = at::cuda::getCurrentCUDAStream(a_packed.get_device()).stream();
440
+ #if !defined(FLASHRT_FP4_GEMM_SOURCE_SM110_ONLY)
441
+ flash_rt::gemm::fp4_w4a16_gemm_dn_streamk_bf16out_sm120(
442
+ a_packed.data_ptr(), b_packed.data_ptr(), sfa.data_ptr(), sfb.data_ptr(),
443
+ out.data_ptr(), checked_int(shape.m, "M"), checked_int(shape.n, "N"),
444
+ checked_int(shape.k, "K"), static_cast<float>(alpha), stream);
445
+ #endif
446
+ #endif
447
+ }
448
+
449
+ void nvfp4_gemm_streamk_bias_bf16(
450
+ torch::Tensor const& a_packed,
451
+ torch::Tensor const& b_packed,
452
+ torch::Tensor const& sfa,
453
+ torch::Tensor const& sfb,
454
+ torch::Tensor const& bias,
455
+ torch::Tensor& out,
456
+ double alpha) {
457
+ auto shape = check_fp4_gemm_inputs(a_packed, b_packed, sfa, sfb);
458
+ check_bf16_cuda(bias, "bias");
459
+ check_bf16_cuda(out, "out");
460
+ TORCH_CHECK(bias.dim() == 1 && bias.numel() == shape.n,
461
+ "bias must have shape (N,)");
462
+ TORCH_CHECK(out.sizes() == torch::IntArrayRef({shape.m, shape.n}),
463
+ "out must have shape (M, N)");
464
+ check_same_device(a_packed, bias, "a_packed", "bias");
465
+ check_same_device(a_packed, out, "a_packed", "out");
466
+ #if defined(CUDA_KERNEL)
467
+ at::cuda::CUDAGuard device_guard(a_packed.device());
468
+ require_sm120(a_packed, "nvfp4_gemm_streamk_bias_bf16");
469
+ auto stream = at::cuda::getCurrentCUDAStream(a_packed.get_device()).stream();
470
+ #if !defined(FLASHRT_FP4_GEMM_SOURCE_SM110_ONLY)
471
+ flash_rt::gemm::fp4_w4a16_gemm_dn_streamk_bias_bf16out_sm120(
472
+ a_packed.data_ptr(), b_packed.data_ptr(), sfa.data_ptr(), sfb.data_ptr(),
473
+ bias.data_ptr(), out.data_ptr(), checked_int(shape.m, "M"),
474
+ checked_int(shape.n, "N"), checked_int(shape.k, "K"),
475
+ static_cast<float>(alpha), stream);
476
+ #endif
477
+ #endif
478
+ }
479
+
480
+ void quantize_fp4_sfa_fp16(
481
+ torch::Tensor const& x,
482
+ torch::Tensor& packed,
483
+ torch::Tensor& sfa,
484
+ bool is_sfb) {
485
+ check_fp16_cuda(x, "x");
486
+ check_uint8_cuda(packed, "packed");
487
+ check_uint8_cuda(sfa, "sfa");
488
+ TORCH_CHECK(x.dim() == 2, "x must have shape (rows, dim)");
489
+ const int64_t rows = x.size(0);
490
+ const int64_t dim = x.size(1);
491
+ TORCH_CHECK(dim % 16 == 0, "x.shape[1] must be divisible by 16");
492
+ TORCH_CHECK(packed.sizes() == torch::IntArrayRef({rows, dim / 2}),
493
+ "packed must have shape (rows, dim / 2)");
494
+ TORCH_CHECK(sfa.numel() >= swizzled_bytes(rows, dim),
495
+ "sfa is too small for CUTLASS SFA/SFB layout");
496
+ check_same_device(x, packed, "x", "packed");
497
+ check_same_device(x, sfa, "x", "sfa");
498
+ #if defined(CUDA_KERNEL)
499
+ at::cuda::CUDAGuard device_guard(x.device());
500
+ auto stream = at::cuda::getCurrentCUDAStream(x.get_device()).stream();
501
+ const int rc = flash_rt::fp4::quantize_fp4_dynamic_sfa_fp16(
502
+ x.data_ptr(), packed.data_ptr(), sfa.data_ptr(),
503
+ checked_int(rows, "rows"), checked_int(dim, "dim"), is_sfb, stream);
504
+ TORCH_CHECK(rc == 0, "quantize_fp4_dynamic_sfa_fp16 failed with rc=", rc);
505
+ #endif
506
+ }
507
+
508
+ void quantize_fp4_sfa_bf16(
509
+ torch::Tensor const& x,
510
+ torch::Tensor& packed,
511
+ torch::Tensor& sfa,
512
+ bool is_sfb) {
513
+ check_bf16_cuda(x, "x");
514
+ check_uint8_cuda(packed, "packed");
515
+ check_uint8_cuda(sfa, "sfa");
516
+ TORCH_CHECK(x.dim() == 2, "x must have shape (rows, dim)");
517
+ const int64_t rows = x.size(0);
518
+ const int64_t dim = x.size(1);
519
+ TORCH_CHECK(dim % 16 == 0, "x.shape[1] must be divisible by 16");
520
+ TORCH_CHECK(packed.sizes() == torch::IntArrayRef({rows, dim / 2}),
521
+ "packed must have shape (rows, dim / 2)");
522
+ TORCH_CHECK(sfa.numel() >= swizzled_bytes(rows, dim),
523
+ "sfa is too small for CUTLASS SFA/SFB layout");
524
+ check_same_device(x, packed, "x", "packed");
525
+ check_same_device(x, sfa, "x", "sfa");
526
+ #if defined(CUDA_KERNEL)
527
+ at::cuda::CUDAGuard device_guard(x.device());
528
+ auto stream = at::cuda::getCurrentCUDAStream(x.get_device()).stream();
529
+ auto const* props = current_device_properties(x);
530
+ int rc = 0;
531
+ if (props->major == 11 && props->minor == 0) {
532
+ TORCH_CHECK(flash_rt::hub::sm110_quantize_bf16_dispatch != nullptr,
533
+ "SM110 vectorized BF16 FP4 quantizer is not present in this build");
534
+ rc = flash_rt::hub::sm110_quantize_bf16_dispatch(
535
+ x.data_ptr(), packed.data_ptr(), sfa.data_ptr(),
536
+ checked_int(rows, "rows"), checked_int(dim, "dim"), is_sfb, stream);
537
+ } else {
538
+ rc = flash_rt::fp4::quantize_fp4_dynamic_sfa_bf16(
539
+ x.data_ptr(), packed.data_ptr(), sfa.data_ptr(),
540
+ checked_int(rows, "rows"), checked_int(dim, "dim"), is_sfb, stream);
541
+ }
542
+ TORCH_CHECK(rc == 0, "quantize_fp4_dynamic_sfa_bf16 failed with rc=", rc);
543
+ #endif
544
+ }
545
+
546
+ void dequantize_fp4_sfa_fp16(
547
+ torch::Tensor const& packed,
548
+ torch::Tensor const& sfa,
549
+ torch::Tensor& out,
550
+ bool is_sfb) {
551
+ check_uint8_cuda(packed, "packed");
552
+ check_uint8_cuda(sfa, "sfa");
553
+ check_fp16_cuda(out, "out");
554
+ TORCH_CHECK(out.dim() == 2, "out must have shape (rows, dim)");
555
+ const int64_t rows = out.size(0);
556
+ const int64_t dim = out.size(1);
557
+ TORCH_CHECK(dim % 16 == 0, "out.shape[1] must be divisible by 16");
558
+ TORCH_CHECK(packed.sizes() == torch::IntArrayRef({rows, dim / 2}),
559
+ "packed must have shape (rows, dim / 2)");
560
+ TORCH_CHECK(sfa.numel() >= swizzled_bytes(rows, dim),
561
+ "sfa is too small for CUTLASS SFA layout");
562
+ check_same_device(packed, sfa, "packed", "sfa");
563
+ check_same_device(packed, out, "packed", "out");
564
+ #if defined(CUDA_KERNEL)
565
+ at::cuda::CUDAGuard device_guard(packed.device());
566
+ auto stream = at::cuda::getCurrentCUDAStream(packed.get_device()).stream();
567
+ flash_rt::fused_fp4::dequantize_fp4_sfa_fp16(
568
+ reinterpret_cast<const uint8_t*>(packed.data_ptr()),
569
+ reinterpret_cast<const uint8_t*>(sfa.data_ptr()),
570
+ reinterpret_cast<__half*>(out.data_ptr()),
571
+ checked_int(rows, "rows"), checked_int(dim, "dim"), is_sfb, stream);
572
+ #endif
573
+ }
574
+
575
+ TORCH_LIBRARY_EXPAND(TORCH_EXTENSION_NAME, ops) {
576
+ ops.def("nvfp4_gemm_bf16(Tensor a_packed, Tensor b_packed, Tensor sfa, Tensor sfb, Tensor! out, float alpha=1.0, int variant=-1) -> ()");
577
+ ops.def("fp4_w4a16_linear_bf16(Tensor a_packed, Tensor b_packed, Tensor sfa, Tensor sfb, Tensor! out, float alpha=1.0, int variant=-1) -> ()");
578
+ ops.def("fp4_w4a4_gemv_warpsplit_bf16(Tensor a_packed, Tensor b_packed, Tensor sfa, Tensor sfb, Tensor! out, float alpha=1.0, int warps=4, int stages=4) -> ()");
579
+ ops.def("nvfp4_gemm_bias_bf16(Tensor a_packed, Tensor b_packed, Tensor sfa, Tensor sfb, Tensor bias, Tensor! out) -> ()");
580
+ ops.def("nvfp4_gemm_bias_residual_bf16(Tensor a_packed, Tensor b_packed, Tensor sfa, Tensor sfb, Tensor bias, Tensor residual, Tensor! out) -> ()");
581
+ ops.def("nvfp4_gemm_residual_bf16(Tensor a_packed, Tensor b_packed, Tensor sfa, Tensor sfb, Tensor residual, Tensor! out, float alpha=1.0) -> ()");
582
+ ops.def("nvfp4_gemm_bias_gelu_bf16(Tensor a_packed, Tensor b_packed, Tensor sfa, Tensor sfb, Tensor bias, Tensor! out, float alpha=1.0) -> ()");
583
+ ops.def("nvfp4_gemm_bias_gelu_nvfp4(Tensor a_packed, Tensor b_packed, Tensor sfa, Tensor sfb, Tensor bias, Tensor! out_packed, Tensor! out_sfa, float alpha=1.0) -> ()");
584
+ ops.def("nvfp4_gemm_streamk_bf16(Tensor a_packed, Tensor b_packed, Tensor sfa, Tensor sfb, Tensor! out, float alpha=1.0) -> ()");
585
+ ops.def("nvfp4_gemm_streamk_bias_bf16(Tensor a_packed, Tensor b_packed, Tensor sfa, Tensor sfb, Tensor bias, Tensor! out, float alpha=1.0) -> ()");
586
+ ops.def("quantize_fp4_sfa_fp16(Tensor x, Tensor! packed, Tensor! sfa, bool is_sfb=False) -> ()");
587
+ ops.def("quantize_fp4_sfa_bf16(Tensor x, Tensor! packed, Tensor! sfa, bool is_sfb=False) -> ()");
588
+ ops.def("dequantize_fp4_sfa_fp16(Tensor packed, Tensor sfa, Tensor! out, bool is_sfb=False) -> ()");
589
+ #if defined(CUDA_KERNEL)
590
+ ops.impl("nvfp4_gemm_bf16", torch::kCUDA, &fp4_w4a16_linear_bf16);
591
+ ops.impl("fp4_w4a16_linear_bf16", torch::kCUDA, &fp4_w4a16_linear_bf16);
592
+ ops.impl("fp4_w4a4_gemv_warpsplit_bf16", torch::kCUDA, &fp4_w4a4_gemv_warpsplit_bf16);
593
+ ops.impl("nvfp4_gemm_bias_bf16", torch::kCUDA, &nvfp4_gemm_bias_bf16);
594
+ ops.impl("nvfp4_gemm_bias_residual_bf16", torch::kCUDA, &nvfp4_gemm_bias_residual_bf16);
595
+ ops.impl("nvfp4_gemm_residual_bf16", torch::kCUDA, &nvfp4_gemm_residual_bf16);
596
+ ops.impl("nvfp4_gemm_bias_gelu_bf16", torch::kCUDA, &nvfp4_gemm_bias_gelu_bf16);
597
+ ops.impl("nvfp4_gemm_bias_gelu_nvfp4", torch::kCUDA, &nvfp4_gemm_bias_gelu_nvfp4);
598
+ ops.impl("nvfp4_gemm_streamk_bf16", torch::kCUDA, &nvfp4_gemm_streamk_bf16);
599
+ ops.impl("nvfp4_gemm_streamk_bias_bf16", torch::kCUDA, &nvfp4_gemm_streamk_bias_bf16);
600
+ ops.impl("quantize_fp4_sfa_fp16", torch::kCUDA, &quantize_fp4_sfa_fp16);
601
+ ops.impl("quantize_fp4_sfa_bf16", torch::kCUDA, &quantize_fp4_sfa_bf16);
602
+ ops.impl("dequantize_fp4_sfa_fp16", torch::kCUDA, &dequantize_fp4_sfa_fp16);
603
+ #endif
604
+ }
605
+
606
+ REGISTER_EXTENSION(TORCH_EXTENSION_NAME)
torch-ext/torch_binding.h ADDED
@@ -0,0 +1,31 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ // SPDX-License-Identifier: Apache-2.0
2
+ #pragma once
3
+
4
+ #include <torch/all.h>
5
+
6
+ void fp4_w4a16_linear_bf16(
7
+ torch::Tensor const& a_packed,
8
+ torch::Tensor const& b_packed,
9
+ torch::Tensor const& sfa,
10
+ torch::Tensor const& sfb,
11
+ torch::Tensor& out,
12
+ double alpha,
13
+ int64_t variant);
14
+
15
+ void quantize_fp4_sfa_fp16(
16
+ torch::Tensor const& x,
17
+ torch::Tensor& packed,
18
+ torch::Tensor& sfa,
19
+ bool is_sfb);
20
+
21
+ void quantize_fp4_sfa_bf16(
22
+ torch::Tensor const& x,
23
+ torch::Tensor& packed,
24
+ torch::Tensor& sfa,
25
+ bool is_sfb);
26
+
27
+ void dequantize_fp4_sfa_fp16(
28
+ torch::Tensor const& packed,
29
+ torch::Tensor const& sfa,
30
+ torch::Tensor& out,
31
+ bool is_sfb);