Initialize legacy kernels compatibility mirror
Browse files
README.md
CHANGED
|
@@ -1,135 +1,9 @@
|
|
| 1 |
-
# fp8-gemm
|
| 2 |
|
| 3 |
-
|
| 4 |
-
|
| 5 |
|
| 6 |
-
|
| 7 |
-
Tensor APIs for Hugging Face Kernel Hub. It is intended for model runtimes that
|
| 8 |
-
already hold activations and weights in FP8 and want a low-overhead BF16 output
|
| 9 |
-
linear path.
|
| 10 |
|
| 11 |
-
|
| 12 |
-
|
| 13 |
-
- `fp8_linear_bf16(input, weight, alpha=1.0, out=None, variant=0)`
|
| 14 |
-
- `fp8_linear_residual_bf16(input, weight, residual, alpha=1.0, variant=0)`
|
| 15 |
-
- `fp8_linear_bias_bf16(input, weight, bias, alpha=1.0, out=None)`
|
| 16 |
-
- `fp8_linear_bias_residual_bf16(input, weight, bias, residual, alpha=1.0)`
|
| 17 |
-
- `fp8_linear_bias_gelu_bf16(input, weight, bias, alpha=1.0, out=None)`
|
| 18 |
-
- `fp8_blockwise_linear_bf16(input, weight, input_scale, weight_scale, out=None)`
|
| 19 |
-
- `fp8_blockwise_swiglu_quantize_fp8(input, gate_up_weight, input_scale, gate_up_weight_scale, output=None, output_scale=None)`
|
| 20 |
-
- `select_fp8_linear_tile(m, n, k, variant=0)`
|
| 21 |
-
|
| 22 |
-
Tensor contract:
|
| 23 |
-
|
| 24 |
-
- `input`: `torch.float8_e4m3fn`, shape `(M, K)`, contiguous CUDA tensor.
|
| 25 |
-
- `weight`: `torch.float8_e4m3fn`, shape `(N, K)`, contiguous CUDA tensor.
|
| 26 |
-
- `out`: `torch.bfloat16`, shape `(M, N)`.
|
| 27 |
-
- `residual`: `torch.bfloat16`, shape `(1, N)` or `(N,)`, only supported for
|
| 28 |
-
the `M=1` decode GEMV path.
|
| 29 |
-
- `K % 16 == 0`; SM120 additionally requires `K % 32 == 0`.
|
| 30 |
-
- On SM120, `M == 1` uses dedicated GEMV and `2 <= M <= 64` uses small-M
|
| 31 |
-
GEMM tiles.
|
| 32 |
-
- On SM110 (Jetson AGX Thor), the per-tensor API uses the production FlashRT
|
| 33 |
-
CUTLASS Sq/T1/Wide family and supports the validated model-shape matrix from
|
| 34 |
-
decode through large vision/backbone rows. The large-M production band is
|
| 35 |
-
validated from `M=65` through `M=1024`, including PI0.5 prefill QKV, O,
|
| 36 |
-
gate/up, and down projections at `M=712..970`. `N` and `K` must be divisible
|
| 37 |
-
by 16.
|
| 38 |
-
- The three BF16 bias APIs are SM110-only. They accept BF16 `(N,)` bias and
|
| 39 |
-
preserve the same row-major FP8 `(M,K)` input and `(N,K)` weight contract.
|
| 40 |
-
The residual API updates a BF16 `(M,N)` tensor in place. The GELU API uses
|
| 41 |
-
the tanh approximation.
|
| 42 |
-
- SM110 `variant=0` is the production auto dispatcher. Diagnostic variants are
|
| 43 |
-
`1=Sq`, `2=T1`, and `3=Wide`; they are correctness-tested but should not be
|
| 44 |
-
pinned by model integrations without a shape-specific benchmark.
|
| 45 |
-
- The per-tensor kernels use Blackwell FP8 MMA instructions and are not valid
|
| 46 |
-
for SM89. SM89 support is provided by the blockwise API below.
|
| 47 |
-
- `alpha` is a host float. For per-tensor FP8 quantization, pass
|
| 48 |
-
`float(input_scale * weight_scale)` from your static calibration metadata.
|
| 49 |
-
|
| 50 |
-
The blockwise API uses a separate contract:
|
| 51 |
-
|
| 52 |
-
- `input`: FP8 E4M3 `(M, K)`.
|
| 53 |
-
- `weight`: FP8 E4M3 `(N, K)`.
|
| 54 |
-
- `input_scale`: FP32 `(M, K / 128)`.
|
| 55 |
-
- `weight_scale`: FP32 `(N / 128, K / 128)`.
|
| 56 |
-
- `N` and `K` must be divisible by 128; `M` is unrestricted.
|
| 57 |
-
- Output is BF16 `(M, N)`.
|
| 58 |
-
- On SM89, the blockwise API dispatches to the production FlashRT native
|
| 59 |
-
`mma.sync.aligned.m16n8k32` GEMM/GEMV implementation.
|
| 60 |
-
- On SM120, it dispatches to the production FlashRT CUTLASS block-scaled
|
| 61 |
-
implementation.
|
| 62 |
-
- SM110 is intentionally not claimed by the blockwise API; use the per-tensor
|
| 63 |
-
static-scale path there. Other architectures are rejected explicitly.
|
| 64 |
-
|
| 65 |
-
The fused SM89 producer accepts FP8 `(M,K)` input, FP8 `(2*N,K)` gate/up
|
| 66 |
-
weight, block-128 FP32 scales, and returns FP8 `(M,N)` plus FP32 `(M,N/128)`
|
| 67 |
-
output scales. Its public range is `1 <= M <= 256` with `N` and `K` divisible
|
| 68 |
-
by 128. It is rejected explicitly on non-SM89 GPUs.
|
| 69 |
-
|
| 70 |
-
## Minimal Usage
|
| 71 |
-
|
| 72 |
-
```python
|
| 73 |
-
from kernels import get_kernel
|
| 74 |
-
import torch
|
| 75 |
-
|
| 76 |
-
ops = get_kernel("flashrt/fp8-gemm", version=1, trust_remote_code=True)
|
| 77 |
-
|
| 78 |
-
x = torch.randn((16, 4096), device="cuda", dtype=torch.bfloat16).to(torch.float8_e4m3fn)
|
| 79 |
-
w = torch.randn((8192, 4096), device="cuda", dtype=torch.bfloat16).to(torch.float8_e4m3fn)
|
| 80 |
-
|
| 81 |
-
y = ops.fp8_linear_bf16(x, w, alpha=1.0)
|
| 82 |
-
```
|
| 83 |
-
|
| 84 |
-
SM110 bias epilogues:
|
| 85 |
-
|
| 86 |
-
```python
|
| 87 |
-
bias = torch.randn((8192,), device="cuda", dtype=torch.bfloat16)
|
| 88 |
-
residual = torch.randn((16, 8192), device="cuda", dtype=torch.bfloat16)
|
| 89 |
-
|
| 90 |
-
y = ops.fp8_linear_bias_bf16(x, w, bias, alpha=1.0)
|
| 91 |
-
ops.fp8_linear_bias_residual_bf16(x, w, bias, residual, alpha=1.0)
|
| 92 |
-
y_gelu = ops.fp8_linear_bias_gelu_bf16(x, w, bias, alpha=1.0)
|
| 93 |
-
```
|
| 94 |
-
|
| 95 |
-
Warm each distinct SM110 bias shape once before CUDA Graph capture. The
|
| 96 |
-
cuBLASLt fallback lazily creates and caches its descriptor, algorithm, and
|
| 97 |
-
workspace on the first call; replay itself performs no allocation.
|
| 98 |
-
|
| 99 |
-
Decode residual path:
|
| 100 |
-
|
| 101 |
-
```python
|
| 102 |
-
x = torch.randn((1, 4096), device="cuda", dtype=torch.bfloat16).to(torch.float8_e4m3fn)
|
| 103 |
-
w = torch.randn((4096, 4096), device="cuda", dtype=torch.bfloat16).to(torch.float8_e4m3fn)
|
| 104 |
-
residual = torch.zeros((1, 4096), device="cuda", dtype=torch.bfloat16)
|
| 105 |
-
|
| 106 |
-
ops.fp8_linear_residual_bf16(x, w, residual, alpha=1.0)
|
| 107 |
-
```
|
| 108 |
-
|
| 109 |
-
Block-128 scaling:
|
| 110 |
-
|
| 111 |
-
```python
|
| 112 |
-
m, k, n = 51, 1536, 1536
|
| 113 |
-
x = torch.randn((m, k), device="cuda").to(torch.float8_e4m3fn)
|
| 114 |
-
w = torch.randn((n, k), device="cuda").to(torch.float8_e4m3fn)
|
| 115 |
-
x_scale = torch.ones((m, k // 128), device="cuda", dtype=torch.float32)
|
| 116 |
-
w_scale = torch.ones((n // 128, k // 128), device="cuda", dtype=torch.float32)
|
| 117 |
-
|
| 118 |
-
y = ops.fp8_blockwise_linear_bf16(x, w, x_scale, w_scale)
|
| 119 |
-
```
|
| 120 |
-
|
| 121 |
-
## Validation
|
| 122 |
-
|
| 123 |
-
```bash
|
| 124 |
-
python fp8-gemm/tests/test_fp8_gemm.py --backend source --mode full
|
| 125 |
-
python fp8-gemm/benchmarks/benchmark.py --backend source --mode headline
|
| 126 |
-
python fp8-gemm/benchmarks/benchmark.py --backend source --mode pi05-prefill
|
| 127 |
-
python fp8-gemm/benchmarks/benchmark_bias.py --backend source
|
| 128 |
-
```
|
| 129 |
-
|
| 130 |
-
The SM110 full sweep covers PI0.5, GROOT N1.6/N1.7, Cosmos Edge, and LingBot
|
| 131 |
-
VLA projection families, plus decode, generic small-M, the `M=65` large-M
|
| 132 |
-
boundary, and the three SigLIP bias epilogues. Public
|
| 133 |
-
benchmark tables are only updated after source correctness, installed artifact
|
| 134 |
-
correctness, shape/tile sweeps, `torch.compile(fullgraph=True)`, CUDA Graph
|
| 135 |
-
replay, and parity against the original FlashRT native pointer entry pass.
|
|
|
|
| 1 |
+
# flashrt/fp8-gemm
|
| 2 |
|
| 3 |
+
This repository is a compatibility mirror for older `kernels` clients
|
| 4 |
+
that resolve repositories through the default Hugging Face model repo API.
|
| 5 |
|
| 6 |
+
Canonical Kernel Hub repo: https://huggingface.co/kernels/flashrt/fp8-gemm
|
|
|
|
|
|
|
|
|
|
| 7 |
|
| 8 |
+
Do not edit this mirror by hand. It is generated from the Kernel Hub
|
| 9 |
+
`vN` branches and contains the same `build/**` artifacts.
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|