File size: 1,327 Bytes
9022070
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
# flashrt/fp8-gemm

FlashRT native CUDA FP8 GEMV/GEMM kernels for low-latency transformer and
diffuser linear layers.

The block-128 scaled API supports Ada `sm_89` and Blackwell `sm_120a`.
The per-tensor API supports Blackwell `sm_110a` (Jetson AGX Thor) and
`sm_120a`. SM110 uses the production FlashRT Sq/T1/Wide CUTLASS family and has
been swept across PI0.5, GROOT, Cosmos Edge, and LingBot VLA projection shapes.

## Functions

- `fp8_linear_bf16(input, weight, alpha=1.0, out=None, variant=0)`
- `fp8_linear_residual_bf16(input, weight, residual, alpha=1.0, variant=0)`
- `fp8_linear_bias_bf16(input, weight, bias, alpha=1.0, out=None)`
- `fp8_linear_bias_residual_bf16(input, weight, bias, residual, alpha=1.0)`
- `fp8_linear_bias_gelu_bf16(input, weight, bias, alpha=1.0, out=None)`
- `fp8_blockwise_linear_bf16(input, weight, input_scale, weight_scale, out=None)`
- `select_fp8_linear_tile(m, n, k, variant=0)`

On SM110, keep `variant=0` for the tuned public dispatcher. Variants `1`, `2`,
and `3` force Sq, T1, and Wide respectively for diagnostics.

SM110 also provides BF16 bias, in-place bias+residual, and tanh-GELU+bias
epilogues for SigLIP-style projection and MLP shapes. The validated large-M
plain GEMM band is `M=65..1024`.

See the repository README for shape contracts, validation status, and examples.