Super tiny expose transform_scale_ue8m0 API for RL frameworks (#13323)

This commit is contained in:
fzyzcjy
2025-11-15 17:31:04 +08:00
committed by GitHub
parent 1d3d42bda0
commit d971f22898

View File

@@ -460,7 +460,7 @@ def _requant_weight_ue8m0(
weight_block_size=weight_block_size,
)
out_s = _transform_scale_ue8m0(out_s, mn=out_w.shape[-2])
out_s = transform_scale_ue8m0(out_s, mn=out_w.shape[-2])
return out_w, out_s
@@ -492,11 +492,11 @@ def quant_weight_ue8m0(
def transform_scale_ue8m0_inplace(param, mn):
param.data = _transform_scale_ue8m0(param.data, mn=mn)
param.data = transform_scale_ue8m0(param.data, mn=mn)
# NOTE copy and modified from DeepGEMM
def _transform_scale_ue8m0(sf, mn):
def transform_scale_ue8m0(sf, mn):
import deep_gemm.utils.layout
sf = sf.index_select(-2, torch.arange(mn, device=sf.device) // 128)
@@ -507,7 +507,7 @@ def _transform_scale_ue8m0(sf, mn):
def inverse_transform_scale_ue8m0(sf_packed, mn):
sf_fp32 = _inverse_transform_scale_ue8m0_impl(sf_packed)
# Can call consistency check every time since this is only called on startup
sf_packed_recreated = _transform_scale_ue8m0(sf_fp32, mn=mn)
sf_packed_recreated = transform_scale_ue8m0(sf_fp32, mn=mn)
assert torch.all(
sf_packed == sf_packed_recreated
), f"{sf_packed=} {sf_packed_recreated}"