#025
per_token_group_quant_fp8
fp8_e4m3 vllm · both · vllm.model_executor.layers.quantization.utils.fp8_utils · importance 0.5%
Reference Implementation
reference.py import torch
import torch.nn as nn
class Model(nn.Module):
def __init__(self, group_size: int) -> None:
super().__init__()
self.group_size = group_size
def forward(self, x: torch.Tensor) -> dict[str, torch.Tensor]:
(m, n) = x.shape
if getattr(torch.version, 'hip', None) and hasattr(torch, 'float8_e4m3fnuz'):
fp8_dtype = torch.float8_e4m3fnuz
fp8_max = 224.0
else:
fp8_dtype = torch.float8_e4m3fn
fp8_max = torch.finfo(torch.float8_e4m3fn).max
x_fp32 = x.float()
x_groups = x_fp32.view(m, n // self.group_size, self.group_size)
x_abs_max = x_groups.abs().amax(dim=-1)
scales = torch.clamp(x_abs_max / fp8_max, min=1e-10)
quantized = (x_fp32 / scales.repeat_interleave(self.group_size, dim=1)).clamp(-fp8_max, fp8_max).to(fp8_dtype)
return {'quantized': quantized, 'scales': scales}
Shapes
TSOL hardware:
| # | token_count | hidden_size | TProd | S |
| 0 | 1 | 2048 | 0.00 us | 7.60 us | 0.0% |
| 1 | 37504 | 2048 | 43.93 us | 106.10 us | 41.4% |
| 2 | 75008 | 2048 | 87.86 us | 199.60 us | 44.0% |
| 3 | 85632 | 2048 | 100.30 us | 223.30 us | 44.9% |
| 4 | 128192 | 2048 | 150.15 us | 325.70 us | 46.1% |
| 5 | 256896 | 2048 | 300.91 us | 639.80 us | 47.0% |
| 6 | 500736 | 2048 | 586.52 us | 1.26 ms | 46.6% |
| 7 | 1029056 | 2048 | 1.21 ms | 2.62 ms | 46.0% |
| 8 | 2058112 | 2048 | 2.41 ms | 5.33 ms | 45.2% |
| 9 | 3087168 | 2048 | 3.62 ms | 7.13 ms | 50.7% |
| 10 | 1 | 7168 | 0.00 us | 7.60 us | 0.0% |
| 11 | 726 | 7168 | 2.98 us | 13.90 us | 21.4% |
| 12 | 969 | 7168 | 3.97 us | 15.60 us | 25.4% |
| 13 | 1933 | 7168 | 7.92 us | 23.10 us | 34.3% |
| 14 | 4192 | 7168 | 17.19 us | 47.60 us | 36.1% |
| 15 | 8667 | 7168 | 35.53 us | 91.50 us | 38.8% |
| 16 | 16768 | 7168 | 68.74 us | 154.10 us | 44.6% |
| 17 | 40448 | 7168 | 165.82 us | 355.80 us | 46.6% |
| 18 | 58688 | 7168 | 240.60 us | 515.10 us | 46.7% |
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}