File size: 5,015 Bytes
2222481
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
import torch

from kernels.benchmark import Benchmark

SHAPES = {
    "qkv": (8704, 6656),
    "o": (6656, 4096),
    "mlp": (39936, 6656),
    "down": (6656, 19968),
    "lm_head": (202048, 6656),
}


def _dequant(pw):
    grid = torch.tensor(
        [0.0, 0.5, 1.0, 1.5, 2.0, 3.0, 4.0, 6.0], device=pw.qweight.device
    )
    gs = pw.global_scale.float()
    out = torch.empty(pw.n, pw.k, dtype=torch.float16, device=pw.qweight.device)
    step = 16384
    for i in range(0, pw.n, step):
        qw = pw.qweight[i : i + step]
        nib = torch.stack([qw & 0xF, qw >> 4], dim=-1)
        nib = nib.reshape(qw.shape[0], pw.k).long()
        sign = torch.where((nib & 0x8) != 0, -1.0, 1.0)
        vals = grid[nib & 0x7] * sign
        sf = pw.sf_rowmajor[i : i + step].view(torch.float8_e4m3fn).float()
        scale = sf.repeat_interleave(16, dim=1)
        out[i : i + step] = ((vals * scale) / gs).to(torch.float16)
    return out


class NvFP4GemmBenchmark(Benchmark):
    seed: int = 42

    def _setup_shape(self, name, m):
        n, k = SHAPES[name]
        torch.manual_seed(self.seed)
        w = torch.randn(n, k, dtype=torch.bfloat16) * 0.02
        self.pw = self.kernel.pack(w.to(self.device))
        self.x = torch.randn(m, k, dtype=torch.bfloat16, device=self.device) * 0.125
        self.ref_w = _dequant(self.pw)
        self.xs = self.x.to(torch.float16)

    def setup_decode_qkv(self):
        self._setup_shape("qkv", 1)

    def benchmark_decode_qkv(self):
        self.out = self.kernel.gemm(self.pw, self.x)

    def verify_decode_qkv(self):
        return (self.xs @ self.ref_w.T).to(torch.bfloat16)

    def setup_decode_o(self):
        self._setup_shape("o", 1)

    def benchmark_decode_o(self):
        self.out = self.kernel.gemm(self.pw, self.x)

    def verify_decode_o(self):
        return (self.xs @ self.ref_w.T).to(torch.bfloat16)

    def setup_decode_mlp(self):
        self._setup_shape("mlp", 1)

    def benchmark_decode_mlp(self):
        self.out = self.kernel.gemm(self.pw, self.x)

    def verify_decode_mlp(self):
        return (self.xs @ self.ref_w.T).to(torch.bfloat16)

    def setup_decode_down(self):
        self._setup_shape("down", 1)

    def benchmark_decode_down(self):
        self.out = self.kernel.gemm(self.pw, self.x)

    def verify_decode_down(self):
        return (self.xs @ self.ref_w.T).to(torch.bfloat16)

    def setup_decode_lm_head(self):
        self._setup_shape("lm_head", 1)

    def benchmark_decode_lm_head(self):
        self.out = self.kernel.gemm(self.pw, self.x)

    def verify_decode_lm_head(self):
        return (self.xs @ self.ref_w.T).to(torch.bfloat16)

    def setup_prefill_mlp(self):
        self._setup_shape("mlp", 512)
        self.x = self.x * 0.2
        xq = self.kernel.quantize_reference(self.x.cpu()).to(self.device)
        self.xs = xq.to(torch.float16)

    def benchmark_prefill_mlp(self):
        self.out = self.kernel.gemm(self.pw, self.x)

    def verify_prefill_mlp(self):
        return (self.xs @ self.ref_w.T).to(torch.bfloat16)


class NvFP4RawGemvBenchmark(Benchmark):
    seed: int = 42

    def _setup_shape(self, name):
        n, k = SHAPES[name]
        torch.manual_seed(self.seed)
        w = torch.randn(n, k, dtype=torch.bfloat16) * 0.02
        pw = self.kernel.pack(w.to(self.device))
        self.qw = pw.qweight
        self.sf = pw.sf_rowmajor
        self.alpha = (1.0 / pw.global_scale).float().reshape(1)
        self.x = torch.randn(1, k, dtype=torch.bfloat16, device=self.device) * 0.125
        self.ref_w = _dequant(pw)
        self.xs = self.x.to(torch.float16)

    def setup_gemv_qkv(self):
        self._setup_shape("qkv")

    def benchmark_gemv_qkv(self):
        self.out = self.kernel._ops.ops.nvfp4_gemv(self.x, self.qw, self.sf, self.alpha)

    def verify_gemv_qkv(self):
        return (self.xs @ self.ref_w.T).to(torch.bfloat16)

    def setup_gemv_o(self):
        self._setup_shape("o")

    def benchmark_gemv_o(self):
        self.out = self.kernel._ops.ops.nvfp4_gemv(self.x, self.qw, self.sf, self.alpha)

    def verify_gemv_o(self):
        return (self.xs @ self.ref_w.T).to(torch.bfloat16)

    def setup_gemv_mlp(self):
        self._setup_shape("mlp")

    def benchmark_gemv_mlp(self):
        self.out = self.kernel._ops.ops.nvfp4_gemv(self.x, self.qw, self.sf, self.alpha)

    def verify_gemv_mlp(self):
        return (self.xs @ self.ref_w.T).to(torch.bfloat16)

    def setup_gemv_down(self):
        self._setup_shape("down")

    def benchmark_gemv_down(self):
        self.out = self.kernel._ops.ops.nvfp4_gemv(self.x, self.qw, self.sf, self.alpha)

    def verify_gemv_down(self):
        return (self.xs @ self.ref_w.T).to(torch.bfloat16)

    def setup_gemv_lm_head(self):
        self._setup_shape("lm_head")

    def benchmark_gemv_lm_head(self):
        self.out = self.kernel._ops.ops.nvfp4_gemv(self.x, self.qw, self.sf, self.alpha)

    def verify_gemv_lm_head(self):
        return (self.xs @ self.ref_w.T).to(torch.bfloat16)