[NPU]Releasing redundant memory of w13_weight and nz when the ascend_fuseep feature is enabled (#19813)
This commit is contained in:
@@ -529,13 +529,11 @@ class NpuFuseEPMoE(DeepEPMoE):
|
|||||||
return weight.view(*original_shape[:dim], -1, *original_shape[dim + 1 :])
|
return weight.view(*original_shape[:dim], -1, *original_shape[dim + 1 :])
|
||||||
|
|
||||||
def _process_weights_after_loading(self, layer: torch.nn.Module) -> None:
|
def _process_weights_after_loading(self, layer: torch.nn.Module) -> None:
|
||||||
cpu_w13 = layer.w13_weight.transpose(1, 2).cpu()
|
cpu_w13 = layer.w13_weight.data.transpose(1, 2).cpu()
|
||||||
w13 = self.reshape_w13_weight(cpu_w13, -1).npu()
|
layer.w13_weight.data = self.reshape_w13_weight(cpu_w13, -1).npu()
|
||||||
w13 = npu_format_cast(w13)
|
layer.w13_weight.data = npu_format_cast(layer.w13_weight.data)
|
||||||
layer.w13_weight = torch.nn.Parameter(w13, requires_grad=False)
|
|
||||||
|
|
||||||
w2 = npu_format_cast(layer.w2_weight)
|
layer.w2_weight.data = npu_format_cast(layer.w2_weight.data)
|
||||||
layer.w2_weight = torch.nn.Parameter(w2, requires_grad=False)
|
|
||||||
|
|
||||||
w13_scale = layer.w13_weight_scale.data.squeeze(-1).contiguous()
|
w13_scale = layer.w13_weight_scale.data.squeeze(-1).contiguous()
|
||||||
w13_scale = self.permute_w13_weight_scale(w13_scale, 128)
|
w13_scale = self.permute_w13_weight_scale(w13_scale, 128)
|
||||||
|
|||||||
Reference in New Issue
Block a user