[NPU] [Quantization] w4a4 MoE layer support (#18924)
This commit is contained in:
@@ -14,6 +14,92 @@ if TYPE_CHECKING:
|
||||
from sglang.srt.layers.quantization.base_config import QuantizationConfig
|
||||
|
||||
|
||||
def npu_fused_experts_w4a4(
|
||||
hidden_states: torch.Tensor,
|
||||
w13: torch.Tensor,
|
||||
w13_scale: torch.Tensor,
|
||||
w2: torch.Tensor,
|
||||
w2_scale: torch.Tensor,
|
||||
topk_weights: torch.Tensor,
|
||||
topk_ids: torch.Tensor,
|
||||
top_k: int,
|
||||
):
|
||||
original_shape = hidden_states.shape
|
||||
original_dtype = hidden_states.dtype
|
||||
scale_dtype = original_dtype if original_dtype == torch.bfloat16 else torch.float32
|
||||
if len(original_shape) == 3:
|
||||
hidden_states = hidden_states.view(-1, hidden_states.shape[-1])
|
||||
num_tokens = hidden_states.shape[0]
|
||||
num_experts = w13.shape[0]
|
||||
|
||||
hidden_states, expanded_row_idx, expert_tokens, _ = (
|
||||
torch.ops.npu.npu_moe_init_routing_v2(
|
||||
hidden_states,
|
||||
topk_ids,
|
||||
active_num=num_tokens * top_k,
|
||||
expert_num=num_experts,
|
||||
expert_tokens_num_type=1,
|
||||
expert_tokens_num_flag=True,
|
||||
active_expert_range=[0, num_experts],
|
||||
quant_mode=-1,
|
||||
)
|
||||
)
|
||||
expert_tokens = expert_tokens.to(torch.int64)
|
||||
|
||||
# gmm1: gate_up_proj
|
||||
hidden_states, pertoken_scale = torch.ops.npu.npu_dynamic_quant(
|
||||
hidden_states, dst_type=torch.quint4x2
|
||||
)
|
||||
scale_args13 = {
|
||||
"scale": [w13_scale],
|
||||
"per_token_scale": [pertoken_scale],
|
||||
}
|
||||
|
||||
hidden_states = torch.ops.npu.npu_grouped_matmul(
|
||||
x=[hidden_states],
|
||||
weight=[w13],
|
||||
**scale_args13,
|
||||
split_item=2,
|
||||
group_list_type=1,
|
||||
group_type=0,
|
||||
group_list=expert_tokens,
|
||||
output_dtype=original_dtype,
|
||||
)[0]
|
||||
# act_fn: swiglu
|
||||
hidden_states = torch.ops.npu.npu_swiglu(hidden_states)
|
||||
hidden_states, pertoken_scale = torch.ops.npu.npu_dynamic_quant(hidden_states)
|
||||
|
||||
scale_args2 = {
|
||||
"scale": [w2_scale.to(scale_dtype)],
|
||||
"per_token_scale": [pertoken_scale],
|
||||
}
|
||||
# gmm2: down_proj
|
||||
hidden_states = torch.ops.npu.npu_grouped_matmul(
|
||||
x=[hidden_states],
|
||||
weight=[w2],
|
||||
**scale_args2,
|
||||
split_item=2,
|
||||
group_list_type=1,
|
||||
group_type=0,
|
||||
group_list=expert_tokens,
|
||||
output_dtype=original_dtype,
|
||||
)[0]
|
||||
|
||||
final_hidden_states = torch.ops.npu.npu_moe_finalize_routing(
|
||||
hidden_states,
|
||||
skip1=None,
|
||||
skip2=None,
|
||||
bias=None,
|
||||
scales=topk_weights,
|
||||
expanded_src_to_dst_row=expanded_row_idx,
|
||||
export_for_source_row=topk_ids,
|
||||
drop_pad_mode=2,
|
||||
)
|
||||
if len(original_shape) == 3:
|
||||
final_hidden_states = final_hidden_states.view(original_shape)
|
||||
return final_hidden_states
|
||||
|
||||
|
||||
def npu_fused_experts(
|
||||
hidden_states: torch.Tensor,
|
||||
w13: torch.Tensor,
|
||||
@@ -231,6 +317,74 @@ class _NPUFusedMoEMethodBase(FusedMoEMethodBase):
|
||||
self.quant_config = quant_config
|
||||
|
||||
|
||||
class NPUW4A4Int4DynamicMoEMethod(_NPUFusedMoEMethodBase):
|
||||
|
||||
def process_weights_after_loading(self, layer: torch.nn.Module) -> None:
|
||||
layer.w13_weight.data = npu_format_cast(layer.w13_weight.data.transpose(1, 2))
|
||||
layer.w13_weight.data = self._pack_to_int32(
|
||||
layer.w13_weight.data.to(torch.int32)
|
||||
)
|
||||
|
||||
layer.w2_weight.data = npu_format_cast(layer.w2_weight.data.transpose(1, 2))
|
||||
|
||||
scale_np = layer.w13_weight_scale.data.cpu().numpy()
|
||||
scale_np.dtype = np.uint32
|
||||
scale_uint64_tensor = torch.from_numpy(scale_np.astype(np.int64)).npu()
|
||||
|
||||
layer.w13_weight_scale = torch.nn.Parameter(
|
||||
scale_uint64_tensor.squeeze(-1), requires_grad=False
|
||||
)
|
||||
layer.w2_weight_scale = torch.nn.Parameter(
|
||||
layer.w2_weight_scale.data.squeeze(-1), requires_grad=False
|
||||
)
|
||||
|
||||
# Compressed-tensors format doesn't have this field
|
||||
if hasattr(layer, "w13_weight_offset"):
|
||||
layer.w13_weight_offset = torch.nn.Parameter(
|
||||
layer.w13_weight_offset.data.squeeze(-1),
|
||||
requires_grad=False,
|
||||
)
|
||||
if hasattr(layer, "w2_weight_offset"):
|
||||
layer.w2_weight_offset = torch.nn.Parameter(
|
||||
layer.w2_weight_offset.data.squeeze(-1),
|
||||
requires_grad=False,
|
||||
)
|
||||
|
||||
def _pack_to_int32(self, weight: torch.Tensor):
|
||||
# pack 8 int4 to int32, we use a int32 to represent a int4
|
||||
assert (
|
||||
weight.shape[-1] % 8 == 0
|
||||
), "the last dim of weight needs to be divided by 8"
|
||||
new_weight = torch.ops.npu.npu_convert_weight_to_int4pack(weight.flatten(0, 1))
|
||||
new_weight = new_weight.view(weight.shape[0], weight.shape[1], -1)
|
||||
return new_weight
|
||||
|
||||
def apply(
|
||||
self,
|
||||
layer,
|
||||
dispatch_output: "StandardDispatchOutput",
|
||||
) -> "CombineInput":
|
||||
from sglang.srt.layers.moe.token_dispatcher import StandardCombineInput
|
||||
|
||||
x = dispatch_output.hidden_states
|
||||
topk_output = dispatch_output.topk_output
|
||||
|
||||
topk_weights, topk_ids, _ = topk_output
|
||||
topk_ids = topk_ids.to(torch.int32)
|
||||
topk_weights = topk_weights.to(x.dtype)
|
||||
output = npu_fused_experts_w4a4(
|
||||
hidden_states=x,
|
||||
w13=layer.w13_weight,
|
||||
w13_scale=layer.w13_weight_scale,
|
||||
w2=layer.w2_weight,
|
||||
w2_scale=layer.w2_weight_scale,
|
||||
topk_weights=topk_weights,
|
||||
topk_ids=topk_ids,
|
||||
top_k=topk_ids.shape[1],
|
||||
)
|
||||
return StandardCombineInput(hidden_states=output)
|
||||
|
||||
|
||||
class NPUW8A8Int8DynamicMoEMethod(_NPUFusedMoEMethodBase):
|
||||
|
||||
def process_weights_after_loading(self, layer: torch.nn.Module) -> None:
|
||||
|
||||
Reference in New Issue
Block a user