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

Function _

bitsandbytes/backends/default/ops.py:39–61  ·  view source on GitHub ↗
(
    A: torch.Tensor,
    row_stats: torch.Tensor,
    col_stats: torch.Tensor,
    dtype: Optional[torch.dtype] = None,
    bias: Optional[torch.Tensor] = None,
)

Source from the content-addressed store, hash-verified

37
38@register_kernel("bitsandbytes::int8_mm_dequant", "default")
39def _(
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")

Callers

nothing calls this directly

Calls 8

_get_4bit_codeFunction · 0.85
_dequantize_4bit_computeFunction · 0.85
_optimizer_update_32bitFunction · 0.85
_int8_linear_matmul_implFunction · 0.70
toMethod · 0.45

Tested by

no test coverage detected