liangsu9988 commited on
Commit
86e4260
·
verified ·
1 Parent(s): 53a1e54

Initialize legacy kernels compatibility mirror

Browse files
Files changed (1) hide show
  1. README.md +6 -132
README.md CHANGED
@@ -1,135 +1,9 @@
1
- # fp8-gemm
2
 
3
- FlashRT native CUDA FP8 GEMV/GEMM kernels for low-latency transformer and
4
- diffuser linear layers on NVIDIA Ada SM89 and Blackwell SM110/SM120 GPUs.
5
 
6
- This package exposes the hand-tuned FP8 E4M3 decode and small-M kernels as
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
- ## Available Functions
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.