[CPU] [BF16] Call fused_experts_cpu, weight_packed_linear and bmm_cpu kernel in DeepSeek model (#6641)
Co-authored-by: Thien Tran <gau.nernst@yahoo.com.sg>
This commit is contained in:
@@ -2457,6 +2457,77 @@ def cpu_has_amx_support():
|
||||
return torch._C._cpu._is_amx_tile_supported() and is_intel_amx_backend_available
|
||||
|
||||
|
||||
def prepack_weight_if_needed(weight):
|
||||
if weight.device != torch.device("cpu"):
|
||||
return weight
|
||||
if not cpu_has_amx_support():
|
||||
return weight
|
||||
|
||||
return torch.ops.sgl_kernel.convert_weight_packed(weight)
|
||||
|
||||
|
||||
# TODO: currently gemm kernel has the below requirements:
|
||||
# OC % TILE_N == 0, where TILE_N = 16
|
||||
# IC % TILE_K == 0, where TILE_K = 32
|
||||
def dim_is_supported(weight):
|
||||
return weight.size(0) % 16 == 0 and weight.size(1) % 32 == 0
|
||||
|
||||
|
||||
def _process_weight_after_loading(module, weight_names, transpose_dims=None) -> None:
|
||||
# Pack weight for get better performance on CPU
|
||||
devices = {getattr(module, weight_name).device for weight_name in weight_names}
|
||||
assert len(devices) == 1, f"Expects all weights to be on the same device"
|
||||
device = devices.pop()
|
||||
|
||||
if transpose_dims:
|
||||
assert len(weight_names) == len(
|
||||
transpose_dims
|
||||
), "len(weight_names) should be equal to len(transpose_dims)"
|
||||
|
||||
for i, weight_name in enumerate(weight_names):
|
||||
weight_tensor = getattr(module, weight_name)
|
||||
|
||||
# We don't pack weight or use intel amx backend if any weight of this module has unsupported dim.
|
||||
if not dim_is_supported(weight_tensor):
|
||||
logger.warning(
|
||||
f"Expects weight.size(0) % 16 == 0 and weight.size(1) % 32 == 0 "
|
||||
f"but {weight_tensor.size(0)=} and {weight_tensor.size(1)=} in {module}. "
|
||||
f"{module} won't use intel amx backend."
|
||||
)
|
||||
module.use_intel_amx_backend = False
|
||||
return
|
||||
|
||||
if transpose_dims and transpose_dims[i]:
|
||||
weight_tensor = weight_tensor.transpose(*transpose_dims[i])
|
||||
|
||||
packed_weight = torch.nn.Parameter(
|
||||
prepack_weight_if_needed(weight_tensor),
|
||||
requires_grad=False,
|
||||
)
|
||||
packed_weight.__dict__ = weight_tensor.__dict__
|
||||
setattr(module, weight_name, packed_weight)
|
||||
|
||||
module.use_intel_amx_backend = (
|
||||
device == torch.device("cpu") and cpu_has_amx_support()
|
||||
)
|
||||
|
||||
if (
|
||||
module.use_intel_amx_backend
|
||||
and hasattr(module, "bias")
|
||||
and module.bias is not None
|
||||
):
|
||||
module.bias = torch.nn.Parameter(module.bias.data.float(), requires_grad=False)
|
||||
|
||||
|
||||
class PackWeightMethod:
|
||||
def __init__(self, weight_names, transpose_dims=None):
|
||||
self.weight_names = weight_names
|
||||
self.transpose_dims = transpose_dims
|
||||
|
||||
def process_weights_after_loading(self, module) -> None:
|
||||
_process_weight_after_loading(module, self.weight_names, self.transpose_dims)
|
||||
|
||||
|
||||
class LazyValue:
|
||||
def __init__(self, creator: Callable):
|
||||
self._creator = creator
|
||||
|
||||
Reference in New Issue
Block a user