(
A: torch.Tensor,
row_stats: torch.Tensor,
col_stats: torch.Tensor,
dtype: Optional[torch.dtype] = None,
bias: Optional[torch.Tensor] = None,
)
| 37 | |
| 38 | @register_kernel("bitsandbytes::int8_mm_dequant", "default") |
| 39 | def _( |
| 40 | A: torch.Tensor, |
| 41 | row_stats: torch.Tensor, |
| 42 | col_stats: torch.Tensor, |
| 43 | dtype: Optional[torch.dtype] = None, |
| 44 | bias: Optional[torch.Tensor] = None, |
| 45 | ) -> torch.Tensor: |
| 46 | if A.dtype != torch.int32: |
| 47 | raise ValueError(f"A must be int32, got {A.dtype}") |
| 48 | if row_stats.dtype != torch.float32: |
| 49 | raise ValueError(f"row_stats must be float32, got {row_stats.dtype}") |
| 50 | if col_stats.dtype != torch.float32: |
| 51 | raise ValueError(f"col_stats must be float32, got {col_stats.dtype}") |
| 52 | |
| 53 | A_calc = A.view(-1, A.shape[-1]) |
| 54 | row_stats = row_stats.reshape(-1).unsqueeze(-1) |
| 55 | col_stats = col_stats.reshape(-1).unsqueeze(0) |
| 56 | |
| 57 | out = A_calc * (row_stats * col_stats) * 6.200124e-05 |
| 58 | if bias is not None: |
| 59 | out += bias |
| 60 | |
| 61 | return out.to(dtype or torch.float16) |
| 62 | |
| 63 | |
| 64 | @register_kernel("bitsandbytes::int8_mixed_scaled_mm", "default") |
nothing calls this directly
no test coverage detected