#007

fp8_dynamic_per_token_quant

fp8_e4m3 rtp-llm · both · aiter.ops.quant · importance 0.8%

Reference Implementation

reference.py
import torch
import torch.nn as nn

class Model(nn.Module):

    def __init__(self) -> None:
        super().__init__()
        self.finfo = torch.finfo(torch.float8_e4m3fnuz)

    def forward(self, x: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
        x_float = x.to(torch.float32)
        row_absmax = x_float.abs().amax(dim=-1, keepdim=True).clamp(min=1e-12)
        scale = row_absmax / self.finfo.max
        inv_scale = scale.reciprocal()
        q = (x_float * inv_scale).clamp(min=self.finfo.min, max=self.finfo.max)
        return (q.to(torch.float8_e4m3fnuz), scale.squeeze(-1))

Shapes

TSOL hardware:
# token_counthidden_size TSOL(XPU-A)TProdS
0 12048 0.00 us 7.60 us 0.0%
1 6082048 0.71 us 8.20 us 8.7%
2 8652048 1.00 us 9.10 us 11.0%
3 20322048 2.36 us 14.00 us 16.9%
4 41442048 4.81 us 21.40 us 22.5%
5 69202048 8.03 us 31.90 us 25.2%
6 162562048 18.86 us 73.60 us 25.6%
7 247122048 28.67 us 105.70 us 27.1%
8 455682048 52.86 us 183.50 us 28.8%
10 14096 0.00 us 7.60 us 0.0%
11 1274096 0.29 us 7.70 us 3.8%
12 2554096 0.59 us 7.90 us 7.5%
13 5134096 1.19 us 9.30 us 12.8%
14 10294096 2.39 us 11.90 us 20.1%
15 20414096 4.73 us 17.50 us 27.0%
16 41044096 9.52 us 28.70 us 33.2%
17 82084096 19.04 us 51.90 us 36.7%
18 15120 0.00 us 7.60 us 0.0%
19 16887168 6.85 us 22.00 us 31.1%
20 38047168 15.44 us 42.40 us 36.4%

Input Generation

input.py
import torch


def _make_inputs(token_count: int, hidden_size: int) -> dict[str, torch.Tensor]:
    x = torch.randn(token_count, hidden_size, dtype=torch.bfloat16, device='cuda')
    return {'x': x}