import torch
import torch.nn as nn
class Model(nn.Module):
_M_CHUNK = 16 * 1024
def __init__(self, group_n: int, group_k: int) -> None:
super().__init__()
self.group_n = group_n
self.group_k = group_k
def forward(self, a: torch.Tensor, b: torch.Tensor, a_scales: torch.Tensor, b_scales: torch.Tensor) -> torch.Tensor:
(m, k) = a.shape
n = b.shape[0]
num_k_blocks = k // self.group_k
b_blocks = b.float().view(n, num_k_blocks, self.group_k).permute(1, 0, 2).contiguous()
b_scale_per_n = b_scales.repeat_interleave(self.group_n, dim=0).t().contiguous()
result = torch.empty(m, n, dtype=torch.float32, device=a.device)
chunk = min(m, self._M_CHUNK) if m > 0 else 1
for mc_start in range(0, m, chunk):
mc_end = min(mc_start + chunk, m)
cm = mc_end - mc_start
a_blocks = a[mc_start:mc_end].float().view(cm, num_k_blocks, self.group_k).permute(1, 0, 2).contiguous()
dots = torch.bmm(a_blocks, b_blocks.transpose(-2, -1))
a_scale_chunk = a_scales[mc_start:mc_end].t().contiguous()
scaled = dots * a_scale_chunk.unsqueeze(-1) * b_scale_per_n.unsqueeze(1)
result[mc_start:mc_end] = scaled.sum(dim=0)
return result| # | m | n | k | group_n | group_k | TSOL(XPU-A) | TProd | S |
|---|---|---|---|---|---|---|---|---|
| 0 | 349 | 2048 | 2048 | 128 | 128 | 6.36 us | 61.80 us | 10.3% |
| 1 | 1024 | 2048 | 2048 | 128 | 128 | 18.67 us | 73.90 us | 25.3% |
| 2 | 4533 | 2048 | 2048 | 128 | 128 | 82.66 us | 300.00 us | 27.6% |
| 3 | 6918 | 2048 | 2048 | 128 | 128 | 126.15 us | 440.40 us | 28.6% |
| 4 | 20480 | 2048 | 2048 | 128 | 128 | 373.46 us | 1.24 ms | 30.2% |
| 5 | 51200 | 2048 | 2048 | 128 | 128 | 933.65 us | 3.10 ms | 30.1% |
| 6 | 65536 | 2048 | 2048 | 128 | 128 | 1.20 ms | 3.97 ms | 30.1% |
| 7 | 98304 | 2048 | 2048 | 128 | 128 | 1.79 ms | 5.93 ms | 30.2% |
| 8 | 118784 | 2048 | 2048 | 128 | 128 | 2.17 ms | 7.18 ms | 30.2% |
| 9 | 250880 | 2048 | 2048 | 128 | 128 | 4.57 ms | 15.13 ms | 30.2% |
| 10 | 516096 | 2048 | 2048 | 128 | 128 | 9.41 ms | 31.14 ms | 30.2% |
| 11 | 704512 | 2048 | 2048 | 128 | 128 | 12.85 ms | 42.52 ms | 30.2% |
| 12 | 237 | 4096 | 4096 | 128 | 128 | 17.29 us | 118.30 us | 14.6% |
| 13 | 3072 | 4096 | 4096 | 128 | 128 | 224.10 us | 801.90 us | 27.9% |
| 14 | 4096 | 4096 | 4096 | 128 | 128 | 298.80 us | 1.06 ms | 28.2% |
| 15 | 8192 | 4096 | 4096 | 128 | 128 | 597.60 us | 2.09 ms | 28.6% |
| 16 | 16384 | 4096 | 4096 | 128 | 128 | 1.20 ms | 4.19 ms | 28.5% |
| 17 | 24576 | 4096 | 4096 | 128 | 128 | 1.79 ms | 6.24 ms | 28.8% |
| 18 | 49152 | 4096 | 4096 | 128 | 128 | 3.59 ms | 12.43 ms | 28.8% |
| 19 | 65536 | 4096 | 4096 | 128 | 128 | 4.78 ms | 16.55 ms | 28.9% |
| 20 | 98304 | 4096 | 4096 | 128 | 128 | 7.17 ms | 24.81 ms | 28.9% |
| 21 | 131072 | 4096 | 4096 | 128 | 128 | 9.56 ms | 33.05 ms | 28.9% |
| 22 | 258048 | 4096 | 4096 | 128 | 128 | 18.82 ms | 65.10 ms | 28.9% |
| 23 | 299520 | 4096 | 4096 | 128 | 128 | 21.85 ms | 75.57 ms | 28.9% |
import torch
_M_CHUNK = 64 * 1024
def _make_inputs(m: int, n: int, k: int, group_n: int, group_k: int) -> dict[str, torch.Tensor]:
fp8_max = torch.finfo(torch.float8_e4m3fn).max
b_raw = torch.randn(n, k, dtype=torch.float32, device='cuda')
b_max = b_raw.view(n // group_n, group_n, k // group_k, group_k).amax(dim=1).amax(dim=-1)
b_scales = torch.clamp(b_max / fp8_max, min=1e-10)
b_quant = (b_raw / b_scales.repeat_interleave(group_n, dim=0).repeat_interleave(group_k, dim=1)).clamp(-fp8_max, fp8_max).to(torch.float8_e4m3fn)
del b_raw
a_quant = torch.empty(m, k, dtype=torch.float8_e4m3fn, device='cuda')
a_scales = torch.empty(m, k // group_k, dtype=torch.float32, device='cuda')
chunk = min(m, _M_CHUNK) if m > 0 else 1
for mc_start in range(0, m, chunk):
mc_end = min(mc_start + chunk, m)
a_raw_c = torch.randn(mc_end - mc_start, k, dtype=torch.float32, device='cuda')
a_max_c = a_raw_c.abs().view(mc_end - mc_start, k // group_k, group_k).amax(dim=-1)
a_scales_c = torch.clamp(a_max_c / fp8_max, min=1e-10)
a_quant_c = (a_raw_c / a_scales_c.repeat_interleave(group_k, dim=1)).clamp(-fp8_max, fp8_max).to(torch.float8_e4m3fn)
a_quant[mc_start:mc_end] = a_quant_c
a_scales[mc_start:mc_end] = a_scales_c
return {'a': a_quant, 'b': b_quant, 'a_scales': a_scales, 'b_scales': b_scales}