(self, device, dtype, has_bias)
| 79 | @pytest.mark.parametrize("dtype", [torch.float16, torch.bfloat16, torch.float32], ids=id_formatter("dtype")) |
| 80 | @pytest.mark.parametrize("has_bias", TRUE_FALSE) |
| 81 | def test_int8_scaled_mm(self, device, dtype, has_bias): |
| 82 | A = torch.randint(-128, 127, (10, 20), dtype=torch.int8, device=device) |
| 83 | B = torch.randint(-128, 127, (30, 20), dtype=torch.int8, device=device) |
| 84 | row_stats = torch.randn(10, dtype=torch.float32, device=device) |
| 85 | col_stats = torch.randn(30, dtype=torch.float32, device=device) |
| 86 | bias = torch.randn(30, dtype=dtype, device=device) if has_bias else None |
| 87 | out = torch.ops.bitsandbytes.int8_scaled_mm(A, B, row_stats, col_stats, bias=bias, dtype=dtype) |
| 88 | |
| 89 | assert out.shape == (10, 30) |
| 90 | assert out.dtype == dtype |
| 91 | assert out.device == A.device |
| 92 | |
| 93 | opcheck(torch.ops.bitsandbytes.int8_scaled_mm, (A, B, row_stats, col_stats, bias, dtype)) |
| 94 | |
| 95 | |
| 96 | class TestInt8BlockwiseQuantOps: |
nothing calls this directly
no outgoing calls
no test coverage detected