(device)
| 26 | # https://github.com/bigscience-workshop/petals/blob/main/tests/test_linear8bitlt.py |
| 27 | @pytest.mark.parametrize("device", get_available_devices()) |
| 28 | def 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()) |
nothing calls this directly
no test coverage detected