MCPcopy Create free account
hub / github.com/bitsandbytes-foundation/bitsandbytes / test_linear_no_igemmlt

Function test_linear_no_igemmlt

tests/test_linear8bitlt.py:28–66  ·  view source on GitHub ↗
(device)

Source from the content-addressed store, hash-verified

26# https://github.com/bigscience-workshop/petals/blob/main/tests/test_linear8bitlt.py
27@pytest.mark.parametrize("device", get_available_devices())
28def test_linear_no_igemmlt(device):
29 linear = torch.nn.Linear(1024, 3072)
30 x = torch.randn(3, 1024, dtype=torch.half)
31 linear_custom = Linear8bitLt(
32 linear.in_features,
33 linear.out_features,
34 linear.bias is not None,
35 has_fp16_weights=False,
36 threshold=6.0,
37 )
38
39 # TODO: Remove, this is no longer implemented
40 linear_custom.state.force_no_igemmlt = True
41
42 linear_custom.weight = bnb.nn.Int8Params(
43 linear.weight.data.clone(),
44 requires_grad=False,
45 has_fp16_weights=False,
46 ).to(linear.weight.dtype)
47 linear_custom.bias = linear.bias
48 linear_custom = linear_custom.to(device)
49 linear = linear.half().to(device)
50
51 x_ref = x.clone().to(device).requires_grad_(True)
52 x_ours = x.clone().to(device).requires_grad_(True)
53 fx_ref = linear(x_ref).float()
54 grad_proj = torch.randn_like(fx_ref)
55 (fx_ref * grad_proj).mean().backward()
56
57 fx_ours = linear_custom(x_ours).float()
58 (fx_ours * grad_proj).mean().backward()
59
60 assert linear_custom.state.CB is not None
61 assert not linear_custom.state.has_fp16_weights
62
63 idx = torch.isclose(fx_ref, fx_ours, atol=0.02, rtol=1e-5)
64 assert (idx == 0).sum().item() < fx_ref.numel() * 2.5e-4
65 torch.testing.assert_close(fx_ref, fx_ours, atol=0.03, rtol=1e-5)
66 torch.testing.assert_close(x_ref.grad, x_ours.grad, atol=0.01, rtol=1e-5)
67
68
69@pytest.mark.parametrize("device", get_available_devices())

Callers

nothing calls this directly

Calls 4

toMethod · 0.95
Linear8bitLtClass · 0.90
toMethod · 0.45
backwardMethod · 0.45

Tested by

no test coverage detected