File size: 5,155 Bytes
88b8ef2
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
# fp4-gemm

FlashRT native Blackwell NVFP4 A4W4 GEMM kernels.

This package consumes packed FP4 E2M1 tensors plus CUTLASS Sm1xx SFA/SFB scale
buffers and produces BF16 output. It is designed to pair with
`flashrt/fp4-fused-ops` and other static low-bit transformer/diffuser runtime
paths.

## Available Functions

- `sfa_size_bytes(rows, dim)`
- `quantize_fp4_sfa_fp16(x, packed=None, sfa=None, is_sfb=False)`
- `quantize_fp4_sfa_bf16(x, packed=None, sfa=None, is_sfb=False)`
- `dequantize_fp4_sfa_fp16(packed, sfa, out=None, is_sfb=False)`
- `nvfp4_gemm_bf16(a_packed, b_packed, sfa, sfb, alpha=1.0, out=None, variant=-1)`
- `nvfp4_gemm_bias_bf16(a_packed, b_packed, sfa, sfb, bias, out=None)`
- `nvfp4_gemm_bias_residual_bf16(a_packed, b_packed, sfa, sfb, bias, residual, out=None)`
- `nvfp4_gemm_residual_bf16(a_packed, b_packed, sfa, sfb, residual, alpha=1.0, out=None)`
- `nvfp4_gemm_bias_gelu_bf16(a_packed, b_packed, sfa, sfb, bias, alpha=1.0, out=None)`
- `nvfp4_gemm_bias_gelu_nvfp4(a_packed, b_packed, sfa, sfb, bias, alpha=1.0, out_packed=None, out_sfa=None)`
- `nvfp4_gemm_streamk_bf16(a_packed, b_packed, sfa, sfb, alpha=1.0, out=None)`
- `nvfp4_gemm_streamk_bias_bf16(a_packed, b_packed, sfa, sfb, bias, alpha=1.0, out=None)`
- `fp4_w4a16_linear_bf16(...)` is retained as a compatibility alias

## Tensor Contract

- `a_packed`: `torch.uint8`, shape `(M, K / 2)`.
- `b_packed`: `torch.uint8`, shape `(N, K / 2)`.
- `sfa`: `torch.uint8`, CUTLASS SFA layout for `(M, K)`.
- `sfb`: `torch.uint8`, CUTLASS SFB layout for `(N, K)`.
- output: `torch.bfloat16`, shape `(M, N)`.
- `K` must be divisible by 16.
- Targets: Blackwell `sm_110a` (Jetson AGX Thor, CUDA 13+) and `sm_120a`
  (RTX Blackwell, CUDA 12.8+).

`variant` selects the CUTLASS schedule:

- `-1`: architecture-aware auto-dispatch (public default).
- `0`: default `<128,128,256>` cooperative schedule.
- `1`: widen `<128,256,128>` schedule, intended for very large `N`.
- `2`: pingpong schedule for A/B testing shape-specific wins.

The canonical linear API and FP4/SFA quantize/dequantize helpers are available
on both SM110 and SM120. SM110 additionally provides the GROOT N1.7 production
epilogues `nvfp4_gemm_bias_bf16`, `nvfp4_gemm_bias_residual_bf16`, and
`nvfp4_gemm_bias_gelu_nvfp4`. The latter emits packed FP4 plus CUTLASS SFA so
the following projection can consume it without a BF16 materialization and a
standalone quantization launch. Stream-K and the older BF16 GELU epilogue keep
their existing SM120 dispatch and reject unsupported architectures explicitly.

The SM110 release gate includes the production `(M,N,K)` shapes
`(41,4608,1536)`, `(41,6144,1536)`, and `(41,1536,6144)`, plus the legacy
`M=51` compatibility row. The kernels are the native sources used by FlashRT's
GROOT N1.7 Thor NVFP4 pipeline.

## Minimal Usage

```python
from kernels import get_kernel
import torch

ops = get_kernel("flashrt/fp4-gemm", version=1, trust_remote_code=True)

x = torch.randn((32, 256), device="cuda", dtype=torch.float16)
w = torch.randn((512, 256), device="cuda", dtype=torch.float16)

a_packed, sfa = ops.quantize_fp4_sfa_fp16(x, is_sfb=False)
b_packed, sfb = ops.quantize_fp4_sfa_fp16(w, is_sfb=True)

y = ops.nvfp4_gemm_bf16(a_packed, b_packed, sfa, sfb, alpha=1.0)
```

For BF16 model activations, use the direct producer so the hot path does not
materialize an intermediate FP16 tensor:

```python
x_bf16 = torch.randn((1, 5120), device="cuda", dtype=torch.bfloat16)
a_packed, sfa = ops.quantize_fp4_sfa_bf16(x_bf16)
```

The BF16 entry writes the same E2M1 bytes and CUTLASS SFA/SFB layout as
`quantize_fp4_sfa_fp16(x_bf16.to(torch.float16))` for finite FP16-range
inputs. It is an additive API; the existing FP16 producer remains unchanged.

The quantize/dequantize helpers are included for examples and validation. A
production runtime should keep weights prepacked and should avoid quantizing in
the hot path unless that producer kernel is part of the intended low-bit block.

Use the bias/GELU and residual variants to avoid returning to BF16
elementwise code between low-bit GEMMs. Stream-K variants are selected only
for the validated large down-projection shapes; unsupported shapes reject
rather than silently selecting a losing schedule.

## Validation

```bash
python fp4-gemm/tests/test_fp4_gemm.py --backend source --mode full
python fp4-gemm/tests/test_fp4_gemm.py --backend installed --mode full \
  --artifact fp4-gemm/build/torch211-cxx11-cu128-x86_64-linux
python fp4-gemm/benchmarks/benchmark.py --backend installed --mode headline \
  --artifact fp4-gemm/build/torch211-cxx11-cu128-x86_64-linux

# Thor model-shape gate
python fp4-gemm/tests/test_fp4_gemm.py --backend installed \
  --mode thor-models \
  --artifact fp4-gemm/build/torch211-cxx11-cu130-aarch64-linux
```

The correctness reference dequantizes the same FP4/SFA and FP4/SFB inputs used
by the kernel, then computes the PyTorch GEMM reference from those dequantized
low-bit values.

The producer gate also checks the BF16 direct entry byte-for-byte against the
established FP16 compatibility chain at decode widths 5120, 6144 and 17408,
plus multi-row activation and SFB layouts.