[Quant][ue8m0 fix] group requant_weight_ue8m0 reduce reserved gpu memory (#31755)
Co-authored-by: root <root@johor-edge-gpu-a17-13.bec-host.baidu.com>
This commit is contained in:
@@ -1552,6 +1552,13 @@ def requant_weight_ue8m0(
|
|||||||
):
|
):
|
||||||
assert weight_block_size == [128, 128]
|
assert weight_block_size == [128, 128]
|
||||||
|
|
||||||
|
# 3D+ weights stack multiple experts (e.g. MoE); requant each group separately.
|
||||||
|
# 2D weights are a single matrix and fall through to the direct path below.
|
||||||
|
if weight.dim() > 2:
|
||||||
|
return _requant_weight_ue8m0_grouped(
|
||||||
|
weight, weight_scale_inv, weight_block_size
|
||||||
|
)
|
||||||
|
|
||||||
*_, n, k = weight.shape
|
*_, n, k = weight.shape
|
||||||
|
|
||||||
weight_dequant = block_quant_dequant(
|
weight_dequant = block_quant_dequant(
|
||||||
@@ -1571,6 +1578,47 @@ def requant_weight_ue8m0(
|
|||||||
return out_w, out_s
|
return out_w, out_s
|
||||||
|
|
||||||
|
|
||||||
|
def _requant_weight_ue8m0_grouped(
|
||||||
|
weight: torch.Tensor,
|
||||||
|
weight_scale_inv: torch.Tensor,
|
||||||
|
weight_block_size: List[int],
|
||||||
|
):
|
||||||
|
*group_dims, n, k = weight.shape
|
||||||
|
w_groups = weight.reshape(-1, n, k)
|
||||||
|
s_groups = weight_scale_inv.reshape(-1, *weight_scale_inv.shape[-2:])
|
||||||
|
num_groups = w_groups.shape[0]
|
||||||
|
|
||||||
|
out_w = None
|
||||||
|
out_s = None
|
||||||
|
for g in range(num_groups):
|
||||||
|
weight_dequant = block_quant_dequant(
|
||||||
|
w_groups[g],
|
||||||
|
s_groups[g],
|
||||||
|
weight_block_size,
|
||||||
|
torch.bfloat16,
|
||||||
|
)
|
||||||
|
w_g, s_g = quant_weight_ue8m0(
|
||||||
|
weight_dequant=weight_dequant,
|
||||||
|
weight_block_size=weight_block_size,
|
||||||
|
)
|
||||||
|
if out_w is None:
|
||||||
|
out_w = torch.empty(
|
||||||
|
(num_groups, *w_g.shape), dtype=w_g.dtype, device=w_g.device
|
||||||
|
)
|
||||||
|
out_s = torch.empty(
|
||||||
|
(num_groups, *s_g.shape), dtype=s_g.dtype, device=s_g.device
|
||||||
|
)
|
||||||
|
out_w[g] = w_g
|
||||||
|
out_s[g] = s_g
|
||||||
|
|
||||||
|
out_w = out_w.view(*group_dims, n, k)
|
||||||
|
out_s = out_s.view(*group_dims, *out_s.shape[-2:])
|
||||||
|
|
||||||
|
out_s = transform_scale_ue8m0(out_s, mn=n)
|
||||||
|
|
||||||
|
return out_w, out_s
|
||||||
|
|
||||||
|
|
||||||
def quant_weight_ue8m0(
|
def quant_weight_ue8m0(
|
||||||
weight_dequant: torch.Tensor,
|
weight_dequant: torch.Tensor,
|
||||||
weight_block_size: List[int],
|
weight_block_size: List[int],
|
||||||
|
|||||||
Reference in New Issue
Block a user