Promote latest kernel artifacts to main
Browse files- README.md +132 -6
- build/torch212-cxx11-cu132-x86_64-linux/__init__.py +4 -2
- build/torch212-cxx11-cu132-x86_64-linux/{_fp8_gemm_cuda_8fba1d7.abi3.so → _fp8_gemm_cuda_9ac1ace.abi3.so} +1 -1
- build/torch212-cxx11-cu132-x86_64-linux/_ops.py +3 -3
- build/torch212-cxx11-cu132-x86_64-linux/metadata.json +5 -5
- build/torch213-cxx11-cu130-x86_64-linux/__init__.py +4 -2
- build/torch213-cxx11-cu130-x86_64-linux/{_fp8_gemm_cuda_8fba1d7.abi3.so → _fp8_gemm_cuda_9ac1ace.abi3.so} +1 -1
- build/torch213-cxx11-cu130-x86_64-linux/_ops.py +3 -3
- build/torch213-cxx11-cu130-x86_64-linux/metadata.json +5 -5
- build/torch213-cxx11-cu132-x86_64-linux/__init__.py +4 -2
- build/torch213-cxx11-cu132-x86_64-linux/{_fp8_gemm_cuda_8fba1d7.abi3.so → _fp8_gemm_cuda_9ac1ace.abi3.so} +1 -1
- build/torch213-cxx11-cu132-x86_64-linux/_ops.py +3 -3
- build/torch213-cxx11-cu132-x86_64-linux/metadata.json +5 -5
README.md
CHANGED
|
@@ -1,9 +1,135 @@
|
|
| 1 |
-
#
|
| 2 |
|
| 3 |
-
|
| 4 |
-
|
| 5 |
|
| 6 |
-
|
|
|
|
|
|
|
|
|
|
| 7 |
|
| 8 |
-
|
| 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.
|
build/torch212-cxx11-cu132-x86_64-linux/__init__.py
CHANGED
|
@@ -169,6 +169,8 @@ def select_fp8_linear_tile(m: int, n: int, k: int, variant: int = 0) -> str:
|
|
| 169 |
raise RuntimeError("small-M dispatcher currently supports variant=0 only")
|
| 170 |
if k % 32:
|
| 171 |
raise RuntimeError("SM120 requires k divisible by 32")
|
|
|
|
|
|
|
| 172 |
if m <= 16:
|
| 173 |
if k % 256 == 0:
|
| 174 |
return "ld_fp8_gemm_16x128x256_w4" if n % 128 == 0 else "ld_fp8_gemm_16x64x256_w4"
|
|
@@ -193,7 +195,7 @@ def select_fp8_linear_tile(m: int, n: int, k: int, variant: int = 0) -> str:
|
|
| 193 |
if n % 128 == 0:
|
| 194 |
return "ld_fp8_gemm_64x128x128_w4"
|
| 195 |
return "ld_fp8_gemm_64x64x128_w4"
|
| 196 |
-
raise RuntimeError("
|
| 197 |
|
| 198 |
|
| 199 |
def fp8_linear_bf16(
|
|
@@ -209,7 +211,7 @@ def fp8_linear_bf16(
|
|
| 209 |
``(M, K)`` and ``(N, K)``. ``alpha`` is a host float, normally the product
|
| 210 |
of static per-tensor input and weight scales. SM110 uses the production
|
| 211 |
CUTLASS Sq/T1/Wide dispatcher over full model row counts; SM120 uses the
|
| 212 |
-
hand-tuned M<=64 path.
|
| 213 |
"""
|
| 214 |
|
| 215 |
if out is None:
|
|
|
|
| 169 |
raise RuntimeError("small-M dispatcher currently supports variant=0 only")
|
| 170 |
if k % 32:
|
| 171 |
raise RuntimeError("SM120 requires k divisible by 32")
|
| 172 |
+
if m > 64:
|
| 173 |
+
return "cublaslt_fp8_large_m"
|
| 174 |
if m <= 16:
|
| 175 |
if k % 256 == 0:
|
| 176 |
return "ld_fp8_gemm_16x128x256_w4" if n % 128 == 0 else "ld_fp8_gemm_16x64x256_w4"
|
|
|
|
| 195 |
if n % 128 == 0:
|
| 196 |
return "ld_fp8_gemm_64x128x128_w4"
|
| 197 |
return "ld_fp8_gemm_64x64x128_w4"
|
| 198 |
+
raise RuntimeError("M must be positive")
|
| 199 |
|
| 200 |
|
| 201 |
def fp8_linear_bf16(
|
|
|
|
| 211 |
``(M, K)`` and ``(N, K)``. ``alpha`` is a host float, normally the product
|
| 212 |
of static per-tensor input and weight scales. SM110 uses the production
|
| 213 |
CUTLASS Sq/T1/Wide dispatcher over full model row counts; SM120 uses the
|
| 214 |
+
hand-tuned M<=64 path and cuBLASLt for larger row counts.
|
| 215 |
"""
|
| 216 |
|
| 217 |
if out is None:
|
build/torch212-cxx11-cu132-x86_64-linux/{_fp8_gemm_cuda_8fba1d7.abi3.so → _fp8_gemm_cuda_9ac1ace.abi3.so}
RENAMED
|
@@ -1,3 +1,3 @@
|
|
| 1 |
version https://git-lfs.github.com/spec/v1
|
| 2 |
-
oid sha256:
|
| 3 |
size 5710472
|
|
|
|
| 1 |
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:a31a5af323997ec6b1eec10dc85d1aa54dc0fdd32c4e41aa3fb262c2609cbd93
|
| 3 |
size 5710472
|
build/torch212-cxx11-cu132-x86_64-linux/_ops.py
CHANGED
|
@@ -1,9 +1,9 @@
|
|
| 1 |
import torch
|
| 2 |
-
from . import
|
| 3 |
-
ops = torch.ops.
|
| 4 |
|
| 5 |
def add_op_namespace_prefix(op_name: str):
|
| 6 |
"""
|
| 7 |
Prefix op by namespace.
|
| 8 |
"""
|
| 9 |
-
return f"
|
|
|
|
| 1 |
import torch
|
| 2 |
+
from . import _fp8_gemm_cuda_9ac1ace
|
| 3 |
+
ops = torch.ops._fp8_gemm_cuda_9ac1ace
|
| 4 |
|
| 5 |
def add_op_namespace_prefix(op_name: str):
|
| 6 |
"""
|
| 7 |
Prefix op by namespace.
|
| 8 |
"""
|
| 9 |
+
return f"_fp8_gemm_cuda_9ac1ace::{op_name}"
|
build/torch212-cxx11-cu132-x86_64-linux/metadata.json
CHANGED
|
@@ -1,6 +1,6 @@
|
|
| 1 |
{
|
| 2 |
"name": "fp8-gemm",
|
| 3 |
-
"id": "
|
| 4 |
"version": 1,
|
| 5 |
"license": "Apache-2.0",
|
| 6 |
"python-depends": [],
|
|
@@ -15,9 +15,9 @@
|
|
| 15 |
"digest": {
|
| 16 |
"algorithm": "sha256",
|
| 17 |
"files": {
|
| 18 |
-
"__init__.py": "
|
| 19 |
-
"
|
| 20 |
-
"_ops.py": "
|
| 21 |
}
|
| 22 |
},
|
| 23 |
"provenance": {
|
|
@@ -27,7 +27,7 @@
|
|
| 27 |
"dirty": false
|
| 28 |
},
|
| 29 |
"kernel": {
|
| 30 |
-
"sha": "
|
| 31 |
"dirty": false
|
| 32 |
}
|
| 33 |
}
|
|
|
|
| 1 |
{
|
| 2 |
"name": "fp8-gemm",
|
| 3 |
+
"id": "_fp8_gemm_cuda_9ac1ace",
|
| 4 |
"version": 1,
|
| 5 |
"license": "Apache-2.0",
|
| 6 |
"python-depends": [],
|
|
|
|
| 15 |
"digest": {
|
| 16 |
"algorithm": "sha256",
|
| 17 |
"files": {
|
| 18 |
+
"__init__.py": "7UODiuG0STVFP2F4ZTuu+2z8Ih/u8572gTNX+uLBP/s=",
|
| 19 |
+
"_fp8_gemm_cuda_9ac1ace.abi3.so": "oxpa8yOZfsax7sENyF0apU3A/dMsTkGqP7JiwmCcvZM=",
|
| 20 |
+
"_ops.py": "8623JQCjzAtJO6ZqPnLQaFzDwRCUC0pkQWHUlwQbtAc="
|
| 21 |
}
|
| 22 |
},
|
| 23 |
"provenance": {
|
|
|
|
| 27 |
"dirty": false
|
| 28 |
},
|
| 29 |
"kernel": {
|
| 30 |
+
"sha": "9ac1ace56531e1dee1125f5dd2f1cafc77b890fa",
|
| 31 |
"dirty": false
|
| 32 |
}
|
| 33 |
}
|
build/torch213-cxx11-cu130-x86_64-linux/__init__.py
CHANGED
|
@@ -169,6 +169,8 @@ def select_fp8_linear_tile(m: int, n: int, k: int, variant: int = 0) -> str:
|
|
| 169 |
raise RuntimeError("small-M dispatcher currently supports variant=0 only")
|
| 170 |
if k % 32:
|
| 171 |
raise RuntimeError("SM120 requires k divisible by 32")
|
|
|
|
|
|
|
| 172 |
if m <= 16:
|
| 173 |
if k % 256 == 0:
|
| 174 |
return "ld_fp8_gemm_16x128x256_w4" if n % 128 == 0 else "ld_fp8_gemm_16x64x256_w4"
|
|
@@ -193,7 +195,7 @@ def select_fp8_linear_tile(m: int, n: int, k: int, variant: int = 0) -> str:
|
|
| 193 |
if n % 128 == 0:
|
| 194 |
return "ld_fp8_gemm_64x128x128_w4"
|
| 195 |
return "ld_fp8_gemm_64x64x128_w4"
|
| 196 |
-
raise RuntimeError("
|
| 197 |
|
| 198 |
|
| 199 |
def fp8_linear_bf16(
|
|
@@ -209,7 +211,7 @@ def fp8_linear_bf16(
|
|
| 209 |
``(M, K)`` and ``(N, K)``. ``alpha`` is a host float, normally the product
|
| 210 |
of static per-tensor input and weight scales. SM110 uses the production
|
| 211 |
CUTLASS Sq/T1/Wide dispatcher over full model row counts; SM120 uses the
|
| 212 |
-
hand-tuned M<=64 path.
|
| 213 |
"""
|
| 214 |
|
| 215 |
if out is None:
|
|
|
|
| 169 |
raise RuntimeError("small-M dispatcher currently supports variant=0 only")
|
| 170 |
if k % 32:
|
| 171 |
raise RuntimeError("SM120 requires k divisible by 32")
|
| 172 |
+
if m > 64:
|
| 173 |
+
return "cublaslt_fp8_large_m"
|
| 174 |
if m <= 16:
|
| 175 |
if k % 256 == 0:
|
| 176 |
return "ld_fp8_gemm_16x128x256_w4" if n % 128 == 0 else "ld_fp8_gemm_16x64x256_w4"
|
|
|
|
| 195 |
if n % 128 == 0:
|
| 196 |
return "ld_fp8_gemm_64x128x128_w4"
|
| 197 |
return "ld_fp8_gemm_64x64x128_w4"
|
| 198 |
+
raise RuntimeError("M must be positive")
|
| 199 |
|
| 200 |
|
| 201 |
def fp8_linear_bf16(
|
|
|
|
| 211 |
``(M, K)`` and ``(N, K)``. ``alpha`` is a host float, normally the product
|
| 212 |
of static per-tensor input and weight scales. SM110 uses the production
|
| 213 |
CUTLASS Sq/T1/Wide dispatcher over full model row counts; SM120 uses the
|
| 214 |
+
hand-tuned M<=64 path and cuBLASLt for larger row counts.
|
| 215 |
"""
|
| 216 |
|
| 217 |
if out is None:
|
build/torch213-cxx11-cu130-x86_64-linux/{_fp8_gemm_cuda_8fba1d7.abi3.so → _fp8_gemm_cuda_9ac1ace.abi3.so}
RENAMED
|
@@ -1,3 +1,3 @@
|
|
| 1 |
version https://git-lfs.github.com/spec/v1
|
| 2 |
-
oid sha256:
|
| 3 |
size 5665200
|
|
|
|
| 1 |
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:311c4b58ff077ff68730c8690d1a44cfabf4815c9f277da3859d8226eed8f47b
|
| 3 |
size 5665200
|
build/torch213-cxx11-cu130-x86_64-linux/_ops.py
CHANGED
|
@@ -1,9 +1,9 @@
|
|
| 1 |
import torch
|
| 2 |
-
from . import
|
| 3 |
-
ops = torch.ops.
|
| 4 |
|
| 5 |
def add_op_namespace_prefix(op_name: str):
|
| 6 |
"""
|
| 7 |
Prefix op by namespace.
|
| 8 |
"""
|
| 9 |
-
return f"
|
|
|
|
| 1 |
import torch
|
| 2 |
+
from . import _fp8_gemm_cuda_9ac1ace
|
| 3 |
+
ops = torch.ops._fp8_gemm_cuda_9ac1ace
|
| 4 |
|
| 5 |
def add_op_namespace_prefix(op_name: str):
|
| 6 |
"""
|
| 7 |
Prefix op by namespace.
|
| 8 |
"""
|
| 9 |
+
return f"_fp8_gemm_cuda_9ac1ace::{op_name}"
|
build/torch213-cxx11-cu130-x86_64-linux/metadata.json
CHANGED
|
@@ -1,6 +1,6 @@
|
|
| 1 |
{
|
| 2 |
"name": "fp8-gemm",
|
| 3 |
-
"id": "
|
| 4 |
"version": 1,
|
| 5 |
"license": "Apache-2.0",
|
| 6 |
"python-depends": [],
|
|
@@ -15,9 +15,9 @@
|
|
| 15 |
"digest": {
|
| 16 |
"algorithm": "sha256",
|
| 17 |
"files": {
|
| 18 |
-
"__init__.py": "
|
| 19 |
-
"
|
| 20 |
-
"_ops.py": "
|
| 21 |
}
|
| 22 |
},
|
| 23 |
"provenance": {
|
|
@@ -27,7 +27,7 @@
|
|
| 27 |
"dirty": false
|
| 28 |
},
|
| 29 |
"kernel": {
|
| 30 |
-
"sha": "
|
| 31 |
"dirty": false
|
| 32 |
}
|
| 33 |
}
|
|
|
|
| 1 |
{
|
| 2 |
"name": "fp8-gemm",
|
| 3 |
+
"id": "_fp8_gemm_cuda_9ac1ace",
|
| 4 |
"version": 1,
|
| 5 |
"license": "Apache-2.0",
|
| 6 |
"python-depends": [],
|
|
|
|
| 15 |
"digest": {
|
| 16 |
"algorithm": "sha256",
|
| 17 |
"files": {
|
| 18 |
+
"__init__.py": "7UODiuG0STVFP2F4ZTuu+2z8Ih/u8572gTNX+uLBP/s=",
|
| 19 |
+
"_fp8_gemm_cuda_9ac1ace.abi3.so": "MRxLWP8Hf/aHMMhpDRpEz6v0gVyfJ32jhZ2CJu7Y9Hs=",
|
| 20 |
+
"_ops.py": "8623JQCjzAtJO6ZqPnLQaFzDwRCUC0pkQWHUlwQbtAc="
|
| 21 |
}
|
| 22 |
},
|
| 23 |
"provenance": {
|
|
|
|
| 27 |
"dirty": false
|
| 28 |
},
|
| 29 |
"kernel": {
|
| 30 |
+
"sha": "9ac1ace56531e1dee1125f5dd2f1cafc77b890fa",
|
| 31 |
"dirty": false
|
| 32 |
}
|
| 33 |
}
|
build/torch213-cxx11-cu132-x86_64-linux/__init__.py
CHANGED
|
@@ -169,6 +169,8 @@ def select_fp8_linear_tile(m: int, n: int, k: int, variant: int = 0) -> str:
|
|
| 169 |
raise RuntimeError("small-M dispatcher currently supports variant=0 only")
|
| 170 |
if k % 32:
|
| 171 |
raise RuntimeError("SM120 requires k divisible by 32")
|
|
|
|
|
|
|
| 172 |
if m <= 16:
|
| 173 |
if k % 256 == 0:
|
| 174 |
return "ld_fp8_gemm_16x128x256_w4" if n % 128 == 0 else "ld_fp8_gemm_16x64x256_w4"
|
|
@@ -193,7 +195,7 @@ def select_fp8_linear_tile(m: int, n: int, k: int, variant: int = 0) -> str:
|
|
| 193 |
if n % 128 == 0:
|
| 194 |
return "ld_fp8_gemm_64x128x128_w4"
|
| 195 |
return "ld_fp8_gemm_64x64x128_w4"
|
| 196 |
-
raise RuntimeError("
|
| 197 |
|
| 198 |
|
| 199 |
def fp8_linear_bf16(
|
|
@@ -209,7 +211,7 @@ def fp8_linear_bf16(
|
|
| 209 |
``(M, K)`` and ``(N, K)``. ``alpha`` is a host float, normally the product
|
| 210 |
of static per-tensor input and weight scales. SM110 uses the production
|
| 211 |
CUTLASS Sq/T1/Wide dispatcher over full model row counts; SM120 uses the
|
| 212 |
-
hand-tuned M<=64 path.
|
| 213 |
"""
|
| 214 |
|
| 215 |
if out is None:
|
|
|
|
| 169 |
raise RuntimeError("small-M dispatcher currently supports variant=0 only")
|
| 170 |
if k % 32:
|
| 171 |
raise RuntimeError("SM120 requires k divisible by 32")
|
| 172 |
+
if m > 64:
|
| 173 |
+
return "cublaslt_fp8_large_m"
|
| 174 |
if m <= 16:
|
| 175 |
if k % 256 == 0:
|
| 176 |
return "ld_fp8_gemm_16x128x256_w4" if n % 128 == 0 else "ld_fp8_gemm_16x64x256_w4"
|
|
|
|
| 195 |
if n % 128 == 0:
|
| 196 |
return "ld_fp8_gemm_64x128x128_w4"
|
| 197 |
return "ld_fp8_gemm_64x64x128_w4"
|
| 198 |
+
raise RuntimeError("M must be positive")
|
| 199 |
|
| 200 |
|
| 201 |
def fp8_linear_bf16(
|
|
|
|
| 211 |
``(M, K)`` and ``(N, K)``. ``alpha`` is a host float, normally the product
|
| 212 |
of static per-tensor input and weight scales. SM110 uses the production
|
| 213 |
CUTLASS Sq/T1/Wide dispatcher over full model row counts; SM120 uses the
|
| 214 |
+
hand-tuned M<=64 path and cuBLASLt for larger row counts.
|
| 215 |
"""
|
| 216 |
|
| 217 |
if out is None:
|
build/torch213-cxx11-cu132-x86_64-linux/{_fp8_gemm_cuda_8fba1d7.abi3.so → _fp8_gemm_cuda_9ac1ace.abi3.so}
RENAMED
|
@@ -1,3 +1,3 @@
|
|
| 1 |
version https://git-lfs.github.com/spec/v1
|
| 2 |
-
oid sha256:
|
| 3 |
size 5710312
|
|
|
|
| 1 |
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:19cb9b7534a0e305d81bb2b9c8e15400729c8c2e3e2ede481a6a3692ffbba3bb
|
| 3 |
size 5710312
|
build/torch213-cxx11-cu132-x86_64-linux/_ops.py
CHANGED
|
@@ -1,9 +1,9 @@
|
|
| 1 |
import torch
|
| 2 |
-
from . import
|
| 3 |
-
ops = torch.ops.
|
| 4 |
|
| 5 |
def add_op_namespace_prefix(op_name: str):
|
| 6 |
"""
|
| 7 |
Prefix op by namespace.
|
| 8 |
"""
|
| 9 |
-
return f"
|
|
|
|
| 1 |
import torch
|
| 2 |
+
from . import _fp8_gemm_cuda_9ac1ace
|
| 3 |
+
ops = torch.ops._fp8_gemm_cuda_9ac1ace
|
| 4 |
|
| 5 |
def add_op_namespace_prefix(op_name: str):
|
| 6 |
"""
|
| 7 |
Prefix op by namespace.
|
| 8 |
"""
|
| 9 |
+
return f"_fp8_gemm_cuda_9ac1ace::{op_name}"
|
build/torch213-cxx11-cu132-x86_64-linux/metadata.json
CHANGED
|
@@ -1,6 +1,6 @@
|
|
| 1 |
{
|
| 2 |
"name": "fp8-gemm",
|
| 3 |
-
"id": "
|
| 4 |
"version": 1,
|
| 5 |
"license": "Apache-2.0",
|
| 6 |
"python-depends": [],
|
|
@@ -15,9 +15,9 @@
|
|
| 15 |
"digest": {
|
| 16 |
"algorithm": "sha256",
|
| 17 |
"files": {
|
| 18 |
-
"__init__.py": "
|
| 19 |
-
"
|
| 20 |
-
"_ops.py": "
|
| 21 |
}
|
| 22 |
},
|
| 23 |
"provenance": {
|
|
@@ -27,7 +27,7 @@
|
|
| 27 |
"dirty": false
|
| 28 |
},
|
| 29 |
"kernel": {
|
| 30 |
-
"sha": "
|
| 31 |
"dirty": false
|
| 32 |
}
|
| 33 |
}
|
|
|
|
| 1 |
{
|
| 2 |
"name": "fp8-gemm",
|
| 3 |
+
"id": "_fp8_gemm_cuda_9ac1ace",
|
| 4 |
"version": 1,
|
| 5 |
"license": "Apache-2.0",
|
| 6 |
"python-depends": [],
|
|
|
|
| 15 |
"digest": {
|
| 16 |
"algorithm": "sha256",
|
| 17 |
"files": {
|
| 18 |
+
"__init__.py": "7UODiuG0STVFP2F4ZTuu+2z8Ih/u8572gTNX+uLBP/s=",
|
| 19 |
+
"_fp8_gemm_cuda_9ac1ace.abi3.so": "GcubdTSg4wXYG7K5yOFUAHKcjC4+Lt5IGmo2kv+7o7s=",
|
| 20 |
+
"_ops.py": "8623JQCjzAtJO6ZqPnLQaFzDwRCUC0pkQWHUlwQbtAc="
|
| 21 |
}
|
| 22 |
},
|
| 23 |
"provenance": {
|
|
|
|
| 27 |
"dirty": false
|
| 28 |
},
|
| 29 |
"kernel": {
|
| 30 |
+
"sha": "9ac1ace56531e1dee1125f5dd2f1cafc77b890fa",
|
| 31 |
"dirty": false
|
| 32 |
}
|
| 33 |
}
|