Promote latest kernel artifacts to main
Browse files- CARD.md +55 -0
- README.md +113 -6
- SYNC.md +49 -0
- VALIDATION.md +83 -0
- benchmarks/RESULTS.md +57 -0
- build.toml +74 -0
- build/torch211-cxx11-cu130-aarch64-linux/__init__.py +51 -0
- build/torch211-cxx11-cu130-aarch64-linux/{_fp4_gemm_cuda_b46a817.abi3.so → _fp4_gemm_cuda_8a66d8b.abi3.so} +2 -2
- build/torch211-cxx11-cu130-aarch64-linux/_ops.py +3 -3
- build/torch211-cxx11-cu130-aarch64-linux/metadata.json +5 -9
- csrc/cutlass/util/packed_stride.hpp +570 -0
- csrc/dequantize_fp4_sfa.cu +89 -0
- csrc/dequantize_fp4_sfa.cuh +21 -0
- csrc/gemm/fp4/cutlass_fp4_gemm_bias_bf16_sm100.cu +208 -0
- csrc/gemm/fp4/cutlass_fp4_gemm_bias_bf16_sm100.cuh +47 -0
- csrc/gemm/fp4/cutlass_nvfp4_gemm_bias_gelu_bf16out_sm120.cu +212 -0
- csrc/gemm/fp4/cutlass_nvfp4_gemm_bias_gelu_bf16out_sm120.cuh +37 -0
- csrc/gemm/fp4/cutlass_nvfp4_gemm_bias_gelu_fp4out_sm120.cu +234 -0
- csrc/gemm/fp4/cutlass_nvfp4_gemm_bias_gelu_fp4out_sm120.cuh +40 -0
- csrc/gemm/fp4/cutlass_nvfp4_gemm_dn_streamk_bias_sm120.cu +307 -0
- csrc/gemm/fp4/cutlass_nvfp4_gemm_dn_streamk_bias_sm120.cuh +53 -0
- csrc/gemm/fp4/cutlass_nvfp4_w4a16_gemm_sm100.cu +411 -0
- csrc/gemm/fp4/cutlass_nvfp4_w4a16_gemm_sm100.cuh +65 -0
- csrc/gemm/fp4/cutlass_nvfp4_w4a16_gemm_sm120.cu +690 -0
- csrc/gemm/fp4/cutlass_nvfp4_w4a16_gemm_sm120.cuh +112 -0
- csrc/gemm/fp4/fp4_w4a4_mma_warpsplit_sm120.cu +147 -0
- csrc/gemm/fp4/fp4_w4a4_mma_warpsplit_sm120.cuh +20 -0
- csrc/gemm/fp4/sm110_dispatch.cu +50 -0
- csrc/gemm/fp4/sm110_dispatch.cuh +45 -0
- csrc/quantize/quantize_fp4_sfa.cu +194 -0
- csrc/quantize/quantize_fp4_sfa.cuh +41 -0
- csrc/quantize/quantize_fp4_sfa_bf16.cu +142 -0
- csrc/quantize/quantize_fp4_sfa_bf16.cuh +24 -0
- examples/README.md +10 -0
- examples/fp4_gemm_linear.py +29 -0
- flake.nix +18 -0
- tests/test_fp4_gemm.py +668 -0
- torch-ext/fp4_gemm/__init__.py +358 -0
- torch-ext/torch_binding.cpp +606 -0
- 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 |
-
#
|
| 2 |
|
| 3 |
-
|
| 4 |
-
that resolve repositories through the default Hugging Face model repo API.
|
| 5 |
|
| 6 |
-
|
|
|
|
|
|
|
|
|
|
| 7 |
|
| 8 |
-
|
| 9 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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:
|
| 3 |
-
size
|
|
|
|
| 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
|
| 3 |
-
ops = torch.ops.
|
| 4 |
|
| 5 |
def add_op_namespace_prefix(op_name: str):
|
| 6 |
-
return f"
|
|
|
|
| 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": "
|
| 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": "
|
| 17 |
-
"
|
| 18 |
-
"_ops.py": "
|
| 19 |
"fp4_gemm/__init__.py": "v6p5XMfQzddhi1fLSAw4HX9CyS0rQsidvu9VsT01xi4="
|
| 20 |
}
|
| 21 |
},
|
| 22 |
"provenance": {
|
| 23 |
"kernel": {
|
| 24 |
-
"sha": "
|
| 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);
|