File size: 4,601 Bytes
7a139da
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
3cd3528
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
3be46d1
 
 
 
 
 
 
 
 
 
 
 
 
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
import pytest
import torch

import kernels

fp4 = kernels.get_kernel("phanerozoic/fp4-train", version=1, trust_remote_code=True)

requires_cuda = pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA required")
E2M1 = None


def _ref_qdq(x):
    grid = torch.tensor([0., .5, 1., 1.5, 2., 3., 4., 6.], device=x.device)
    M, K = x.shape
    amax = x.reshape(M, K // 16, 16).abs().amax(-1).clamp(min=1e-12)
    s = (amax / 6).to(torch.float8_e4m3fn).float().repeat_interleave(16, 1)
    xs = x / s
    mag = xs.abs().clamp(max=6)
    idx = torch.bucketize(mag, grid, right=False)
    hi = grid[idx.clamp(max=7)]; lo = grid[(idx - 1).clamp(min=0)]
    return (torch.where((mag - lo) < (hi - mag), lo, hi) * xs.sign()) * s


@requires_cuda
def test_encoding_matches_reference():
    x = torch.randn(256, 512, device="cuda")
    p, s = fp4.quantize(x)
    d = fp4.dequantize(p, s, *x.shape).float()
    assert (d - _ref_qdq(x)).abs().max().item() < 1e-3


@requires_cuda
def test_stochastic_rounding_unbiased():
    xb = torch.empty(4096, 16, device="cuda"); xb[:, 0] = 6.0; xb[:, 1:] = 2.3
    acc = torch.zeros_like(xb)
    for stp in range(48):
        p, s = fp4.quantize(xb, stochastic=True, step=stp)
        acc += fp4.dequantize(p, s, *xb.shape).float()
    assert abs((acc / 48)[:, 1:].mean().item() - 2.3) < 0.03


@requires_cuda
def test_error_feedback_cancels_drift():
    xb = torch.empty(4096, 16, device="cuda"); xb[:, 0] = 6.0; xb[:, 1:] = 2.3
    res = torch.zeros_like(xb, dtype=torch.float32); acc = torch.zeros_like(xb)
    for stp in range(48):
        p, s = fp4.quantize(xb, stochastic=False, residual=res, step=stp)
        acc += fp4.dequantize(p, s, *xb.shape).float()
    assert abs((acc / 48)[:, 1:].mean().item() - 2.3) < 0.03


@requires_cuda
def test_fp4_linear_trains():
    torch.manual_seed(0); d = 256
    W = [torch.randn(d, d, device="cuda", requires_grad=True) for _ in range(2)]
    X = torch.randn(512, d, device="cuda")
    Y = torch.tanh(X @ (torch.randn(d, d, device="cuda") / d**0.5))
    opt = torch.optim.Adam(W, lr=3e-3)
    losses = []
    for stp in range(150):
        h = X
        for i, w in enumerate(W):
            h = fp4.fp4_linear(h, w, stochastic=True, step=stp * 2 + i)
            if i < len(W) - 1:
                h = torch.relu(h)
        loss = (h - Y).pow(2).mean()
        opt.zero_grad(); loss.backward(); opt.step(); losses.append(loss.item())
    assert losses[-1] < losses[0] * 0.5


def _kqdq(x, sr, res, step):
    p, s = fp4.quantize(x.contiguous(), stochastic=sr, residual=res, seed=0, step=step)
    return fp4.dequantize(p, s, *x.shape).float()


@requires_cuda
def test_weight_storage_ef_beats_rtn():
    # weights stored in NVFP4 and updated in place; sub-ULP updates. Round-to-nearest
    # stalls; error feedback recovers near-bf16 loss.
    torch.manual_seed(2); d, n = 128, 2048
    Wt = [torch.randn(d, d, device="cuda") / d**0.5 for _ in range(2)]
    X = torch.randn(n, d, device="cuda"); Y = torch.relu(X @ Wt[0].t()) @ Wt[1].t()
    W0 = [torch.randn(d, d, device="cuda") / d**0.5 for _ in range(2)]

    def run(mode, steps=1000, lr=1e-3):
        W = [w.clone() for w in W0]
        res = [torch.zeros_like(w) for w in W] if mode == "ef" else [None, None]
        if mode != "bf16":
            W = [_kqdq(w, mode == "sr", r, 0) for w, r in zip(W, res)]
        for t in range(steps):
            h = torch.relu(X @ W[0].t()); e = h @ W[1].t() - Y
            g1 = 2.0 / n * (e.t() @ h); gh = (e @ W[1]) * (X @ W[0].t() > 0).float()
            g0 = 2.0 / n * (gh.t() @ X)
            for i, g in enumerate((g0, g1)):
                tgt = W[i] - lr * g
                W[i] = tgt if mode == "bf16" else _kqdq(tgt, mode == "sr", res[i], t + 1)
        return (torch.relu(X @ W[0].t()) @ W[1].t() - Y).pow(2).mean().item()

    bf16, rtn, ef = run("bf16"), run("rtn"), run("ef")
    assert ef < rtn * 0.5      # error feedback decisively beats round-to-nearest
    assert ef < bf16 * 1.3     # and recovers near-bf16 loss


@requires_cuda
def test_fp4_mm_hardware_matches_fp4_matmul():
    # native NVFP4 tensor-core GEMM is bit-close to the dequant path (both are the
    # FP4-operand product); requires a Blackwell FP4-capable GPU.
    if torch.cuda.get_device_capability(0)[0] < 10:
        pytest.skip("native FP4 _scaled_mm needs Blackwell (cc >= 10)")
    a = torch.randn(256, 1024, device="cuda")
    b = torch.randn(512, 1024, device="cuda")
    hw = fp4.fp4_mm(a, b).float()
    ref = fp4.fp4_matmul(a, b).float()
    assert (hw - ref).abs().max().item() / (ref.abs().max().item() + 1e-9) < 0.05