[AMD] Enable share expert fusion with router experts for Qwen3.5 BF16 & FP8 (#20736)
Co-authored-by: Chen, Todd <zhenchen@amd.com> Co-authored-by: jacky.cheng <yichiche@amd.com>
This commit is contained in:
co-authored by
Chen, Todd
jacky.cheng
parent
46c8a597ef
commit
ea05ea5abe
@@ -62,7 +62,7 @@ from sglang.srt.layers.moe import (
|
|||||||
)
|
)
|
||||||
from sglang.srt.layers.moe.ep_moe.layer import get_moe_impl_class
|
from sglang.srt.layers.moe.ep_moe.layer import get_moe_impl_class
|
||||||
from sglang.srt.layers.moe.fused_moe_triton import FusedMoE
|
from sglang.srt.layers.moe.fused_moe_triton import FusedMoE
|
||||||
from sglang.srt.layers.moe.topk import TopK
|
from sglang.srt.layers.moe.topk import StandardTopKOutput, TopK, TopKOutputChecker
|
||||||
from sglang.srt.layers.moe.utils import (
|
from sglang.srt.layers.moe.utils import (
|
||||||
RoutingMethodType,
|
RoutingMethodType,
|
||||||
filter_moe_weight_param_global_expert,
|
filter_moe_weight_param_global_expert,
|
||||||
@@ -88,8 +88,10 @@ from sglang.srt.server_args import get_global_server_args
|
|||||||
from sglang.srt.utils import (
|
from sglang.srt.utils import (
|
||||||
add_prefix,
|
add_prefix,
|
||||||
cpu_has_amx_support,
|
cpu_has_amx_support,
|
||||||
|
get_bool_env_var,
|
||||||
is_cpu,
|
is_cpu,
|
||||||
is_cuda,
|
is_cuda,
|
||||||
|
is_hip,
|
||||||
make_layers,
|
make_layers,
|
||||||
use_intel_amx_backend,
|
use_intel_amx_backend,
|
||||||
)
|
)
|
||||||
@@ -100,6 +102,25 @@ logger = logging.getLogger(__name__)
|
|||||||
_is_cuda = is_cuda()
|
_is_cuda = is_cuda()
|
||||||
_is_cpu = is_cpu()
|
_is_cpu = is_cpu()
|
||||||
_is_cpu_amx_available = cpu_has_amx_support()
|
_is_cpu_amx_available = cpu_has_amx_support()
|
||||||
|
_is_hip = is_hip()
|
||||||
|
_use_aiter = get_bool_env_var("SGLANG_USE_AITER") and _is_hip
|
||||||
|
|
||||||
|
|
||||||
|
def can_fuse_shared_expert(
|
||||||
|
config: PretrainedConfig,
|
||||||
|
) -> bool:
|
||||||
|
"""Whether the shared expert may be fused as an extra MoE expert (Qwen3.5 + Aiter).
|
||||||
|
|
||||||
|
Caller must still gate on ``support_shared_expert_fusion`` and ``_use_aiter``.
|
||||||
|
"""
|
||||||
|
if (
|
||||||
|
get_global_server_args().disable_shared_experts_fusion is True
|
||||||
|
or getattr(config, "shared_expert_intermediate_size", 0) <= 0
|
||||||
|
or config.shared_expert_intermediate_size != config.moe_intermediate_size
|
||||||
|
or get_moe_a2a_backend().is_deepep()
|
||||||
|
):
|
||||||
|
return False
|
||||||
|
return True
|
||||||
|
|
||||||
|
|
||||||
class Qwen2MoeMLP(nn.Module):
|
class Qwen2MoeMLP(nn.Module):
|
||||||
@@ -163,6 +184,7 @@ class Qwen2MoeSparseMoeBlock(nn.Module):
|
|||||||
alt_stream: Optional[torch.cuda.Stream] = None,
|
alt_stream: Optional[torch.cuda.Stream] = None,
|
||||||
prefix: str = "",
|
prefix: str = "",
|
||||||
is_nextn: bool = False,
|
is_nextn: bool = False,
|
||||||
|
support_shared_expert_fusion: bool = False,
|
||||||
):
|
):
|
||||||
super().__init__()
|
super().__init__()
|
||||||
self.tp_size = get_tensor_model_parallel_world_size()
|
self.tp_size = get_tensor_model_parallel_world_size()
|
||||||
@@ -173,6 +195,27 @@ class Qwen2MoeSparseMoeBlock(nn.Module):
|
|||||||
f"Tensor parallel size {self.tp_size} is greater than "
|
f"Tensor parallel size {self.tp_size} is greater than "
|
||||||
f"the number of experts {config.num_experts}."
|
f"the number of experts {config.num_experts}."
|
||||||
)
|
)
|
||||||
|
self.num_experts = config.num_experts
|
||||||
|
self.num_shared_experts = 0
|
||||||
|
self.num_fused_shared_experts = 0
|
||||||
|
if hasattr(config, "n_shared_experts"):
|
||||||
|
# config defines the number of shared experts
|
||||||
|
self.num_shared_experts = config.n_shared_experts
|
||||||
|
elif (
|
||||||
|
hasattr(config, "shared_expert_intermediate_size")
|
||||||
|
and config.shared_expert_intermediate_size > 0
|
||||||
|
):
|
||||||
|
# n_shared_experts is not defined, but shared_expert_intermediate_size is defined, so we use 1 as the number of shared experts
|
||||||
|
self.num_shared_experts = 1
|
||||||
|
|
||||||
|
self.enable_shared_expert_fusion = False # default to False
|
||||||
|
if _use_aiter:
|
||||||
|
# enable shared expert fusion when use aiter
|
||||||
|
self.enable_shared_expert_fusion = (
|
||||||
|
support_shared_expert_fusion and can_fuse_shared_expert(config)
|
||||||
|
)
|
||||||
|
if self.enable_shared_expert_fusion:
|
||||||
|
self.num_fused_shared_experts = self.num_shared_experts
|
||||||
|
|
||||||
self.topk = TopK(
|
self.topk = TopK(
|
||||||
top_k=config.num_experts_per_tok,
|
top_k=config.num_experts_per_tok,
|
||||||
@@ -182,14 +225,24 @@ class Qwen2MoeSparseMoeBlock(nn.Module):
|
|||||||
|
|
||||||
self.experts = get_moe_impl_class(quant_config)(
|
self.experts = get_moe_impl_class(quant_config)(
|
||||||
layer_id=self.layer_id,
|
layer_id=self.layer_id,
|
||||||
top_k=config.num_experts_per_tok,
|
top_k=(
|
||||||
num_experts=config.num_experts
|
config.num_experts_per_tok
|
||||||
+ get_global_server_args().ep_num_redundant_experts,
|
if not self.enable_shared_expert_fusion
|
||||||
|
else config.num_experts_per_tok + self.num_fused_shared_experts
|
||||||
|
),
|
||||||
|
num_experts=(
|
||||||
|
config.num_experts + get_global_server_args().ep_num_redundant_experts
|
||||||
|
if not self.enable_shared_expert_fusion
|
||||||
|
else config.num_experts
|
||||||
|
+ get_global_server_args().ep_num_redundant_experts
|
||||||
|
+ self.num_fused_shared_experts
|
||||||
|
),
|
||||||
hidden_size=config.hidden_size,
|
hidden_size=config.hidden_size,
|
||||||
intermediate_size=config.moe_intermediate_size,
|
intermediate_size=config.moe_intermediate_size,
|
||||||
quant_config=quant_config,
|
quant_config=quant_config,
|
||||||
prefix=add_prefix("experts", prefix),
|
prefix=add_prefix("experts", prefix),
|
||||||
routing_method_type=RoutingMethodType.RenormalizeNaive,
|
routing_method_type=RoutingMethodType.RenormalizeNaive,
|
||||||
|
num_fused_shared_experts=self.num_fused_shared_experts,
|
||||||
)
|
)
|
||||||
|
|
||||||
self.gate = ReplicatedLinear(
|
self.gate = ReplicatedLinear(
|
||||||
@@ -199,7 +252,13 @@ class Qwen2MoeSparseMoeBlock(nn.Module):
|
|||||||
quant_config=None,
|
quant_config=None,
|
||||||
prefix=add_prefix("gate", prefix),
|
prefix=add_prefix("gate", prefix),
|
||||||
)
|
)
|
||||||
if config.shared_expert_intermediate_size > 0:
|
# When enable_shared_expert_fusion, the shared expert runs inside the MoE kernel
|
||||||
|
# (via _append_shared_to_topk_output); a separate shared_expert MLP would
|
||||||
|
# double-count. If fusion is off (num_fused_shared_experts == 0), keep shared_expert.
|
||||||
|
if (
|
||||||
|
config.shared_expert_intermediate_size > 0
|
||||||
|
and not self.enable_shared_expert_fusion
|
||||||
|
):
|
||||||
self.shared_expert = Qwen2MoeMLP(
|
self.shared_expert = Qwen2MoeMLP(
|
||||||
hidden_size=config.hidden_size,
|
hidden_size=config.hidden_size,
|
||||||
intermediate_size=config.shared_expert_intermediate_size,
|
intermediate_size=config.shared_expert_intermediate_size,
|
||||||
@@ -245,6 +304,46 @@ class Qwen2MoeSparseMoeBlock(nn.Module):
|
|||||||
)
|
)
|
||||||
]
|
]
|
||||||
|
|
||||||
|
def _get_shared_expert_weights(self, hidden_states: torch.Tensor) -> torch.Tensor:
|
||||||
|
"""Return sigmoid(shared_expert_gate) for fused shared expert weights."""
|
||||||
|
if not self.enable_shared_expert_fusion or self.shared_expert_gate is None:
|
||||||
|
return None
|
||||||
|
shared_out = self.shared_expert_gate(hidden_states)
|
||||||
|
shared_logits = shared_out[0] if isinstance(shared_out, tuple) else shared_out
|
||||||
|
return F.sigmoid(shared_logits)
|
||||||
|
|
||||||
|
def _append_shared_to_topk_output(
|
||||||
|
self,
|
||||||
|
topk_output: StandardTopKOutput,
|
||||||
|
hidden_states: torch.Tensor,
|
||||||
|
) -> StandardTopKOutput:
|
||||||
|
"""Append shared expert ids and weights to topk output before fused MoE."""
|
||||||
|
if not self.enable_shared_expert_fusion:
|
||||||
|
return topk_output
|
||||||
|
shared_weights = self._get_shared_expert_weights(hidden_states)
|
||||||
|
if shared_weights is None:
|
||||||
|
return topk_output
|
||||||
|
M = topk_output.topk_ids.shape[0]
|
||||||
|
shared_expert_id = self.num_experts
|
||||||
|
shared_ids = torch.full(
|
||||||
|
(M, self.num_fused_shared_experts),
|
||||||
|
shared_expert_id,
|
||||||
|
dtype=topk_output.topk_ids.dtype,
|
||||||
|
device=topk_output.topk_ids.device,
|
||||||
|
)
|
||||||
|
shared_weights = shared_weights.expand(M, self.num_fused_shared_experts).to(
|
||||||
|
topk_output.topk_weights.dtype
|
||||||
|
)
|
||||||
|
fused_topk_ids = torch.cat([topk_output.topk_ids, shared_ids], dim=-1)
|
||||||
|
fused_topk_weights = torch.cat(
|
||||||
|
[topk_output.topk_weights, shared_weights], dim=-1
|
||||||
|
)
|
||||||
|
return StandardTopKOutput(
|
||||||
|
topk_weights=fused_topk_weights,
|
||||||
|
topk_ids=fused_topk_ids,
|
||||||
|
router_logits=topk_output.router_logits,
|
||||||
|
)
|
||||||
|
|
||||||
def _forward_shared_experts(self, hidden_states: torch.Tensor):
|
def _forward_shared_experts(self, hidden_states: torch.Tensor):
|
||||||
shared_output = None
|
shared_output = None
|
||||||
if self.shared_expert is not None:
|
if self.shared_expert is not None:
|
||||||
@@ -300,6 +399,10 @@ class Qwen2MoeSparseMoeBlock(nn.Module):
|
|||||||
# router_logits: (num_tokens, n_experts)
|
# router_logits: (num_tokens, n_experts)
|
||||||
router_logits, _ = self.gate(hidden_states)
|
router_logits, _ = self.gate(hidden_states)
|
||||||
topk_output = self.topk(hidden_states, router_logits)
|
topk_output = self.topk(hidden_states, router_logits)
|
||||||
|
if self.enable_shared_expert_fusion and TopKOutputChecker.format_is_standard(
|
||||||
|
topk_output
|
||||||
|
):
|
||||||
|
topk_output = self._append_shared_to_topk_output(topk_output, hidden_states)
|
||||||
return self.experts(hidden_states, topk_output)
|
return self.experts(hidden_states, topk_output)
|
||||||
|
|
||||||
def forward_normal_dual_stream(
|
def forward_normal_dual_stream(
|
||||||
|
|||||||
@@ -86,9 +86,11 @@ from sglang.srt.utils import (
|
|||||||
LazyValue,
|
LazyValue,
|
||||||
add_prefix,
|
add_prefix,
|
||||||
cpu_has_amx_support,
|
cpu_has_amx_support,
|
||||||
|
get_bool_env_var,
|
||||||
is_cpu,
|
is_cpu,
|
||||||
is_cuda,
|
is_cuda,
|
||||||
is_gfx95_supported,
|
is_gfx95_supported,
|
||||||
|
is_hip,
|
||||||
is_npu,
|
is_npu,
|
||||||
make_layers,
|
make_layers,
|
||||||
set_weight_attrs,
|
set_weight_attrs,
|
||||||
@@ -100,6 +102,8 @@ _is_cuda = is_cuda()
|
|||||||
_is_npu = is_npu()
|
_is_npu = is_npu()
|
||||||
_is_cpu = is_cpu()
|
_is_cpu = is_cpu()
|
||||||
_is_gfx95 = is_gfx95_supported()
|
_is_gfx95 = is_gfx95_supported()
|
||||||
|
_is_hip = is_hip()
|
||||||
|
_use_aiter = get_bool_env_var("SGLANG_USE_AITER") and _is_hip
|
||||||
_is_amx_available = cpu_has_amx_support()
|
_is_amx_available = cpu_has_amx_support()
|
||||||
|
|
||||||
|
|
||||||
@@ -528,6 +532,7 @@ class Qwen3_5LinearDecoderLayer(nn.Module):
|
|||||||
alt_stream=alt_stream,
|
alt_stream=alt_stream,
|
||||||
prefix=add_prefix("mlp", prefix.replace(".linear_attn", "")),
|
prefix=add_prefix("mlp", prefix.replace(".linear_attn", "")),
|
||||||
is_nextn=is_nextn,
|
is_nextn=is_nextn,
|
||||||
|
support_shared_expert_fusion=True,
|
||||||
)
|
)
|
||||||
is_layer_sparse = True
|
is_layer_sparse = True
|
||||||
is_previous_layer_sparse = True
|
is_previous_layer_sparse = True
|
||||||
@@ -739,6 +744,7 @@ class Qwen3_5AttentionDecoderLayer(nn.Module):
|
|||||||
alt_stream=alt_stream,
|
alt_stream=alt_stream,
|
||||||
prefix=add_prefix("mlp", prefix.replace(".self_attn", "")),
|
prefix=add_prefix("mlp", prefix.replace(".self_attn", "")),
|
||||||
is_nextn=is_nextn,
|
is_nextn=is_nextn,
|
||||||
|
support_shared_expert_fusion=True,
|
||||||
)
|
)
|
||||||
is_layer_sparse = True
|
is_layer_sparse = True
|
||||||
is_previous_layer_sparse = True
|
is_previous_layer_sparse = True
|
||||||
@@ -1494,6 +1500,20 @@ class Qwen3_5MoeForConditionalGeneration(Qwen3VLForConditionalGeneration):
|
|||||||
self.is_mrope_enabled = "mrope_section" in rope_config
|
self.is_mrope_enabled = "mrope_section" in rope_config
|
||||||
|
|
||||||
self.deepstack_visual_indexes = self.visual.deepstack_visual_indexes
|
self.deepstack_visual_indexes = self.visual.deepstack_visual_indexes
|
||||||
|
self.num_fused_shared_experts = 0
|
||||||
|
if _use_aiter:
|
||||||
|
self.num_fused_shared_experts = self._get_num_fused_shared_experts()
|
||||||
|
|
||||||
|
self.enable_shared_expert_fusion = self.num_fused_shared_experts > 0
|
||||||
|
|
||||||
|
def _get_num_fused_shared_experts(self):
|
||||||
|
if not (
|
||||||
|
hasattr(self.model, "layers")
|
||||||
|
and len(self.model.layers) > 0
|
||||||
|
and hasattr(self.model.layers[0].mlp, "num_fused_shared_experts")
|
||||||
|
):
|
||||||
|
return 0
|
||||||
|
return self.model.layers[0].mlp.num_fused_shared_experts
|
||||||
|
|
||||||
def get_embed_and_head(self):
|
def get_embed_and_head(self):
|
||||||
embed = self.model.embed_tokens.weight if self.pp_group.is_first_rank else None
|
embed = self.model.embed_tokens.weight if self.pp_group.is_first_rank else None
|
||||||
@@ -1525,13 +1545,19 @@ class Qwen3_5MoeForConditionalGeneration(Qwen3VLForConditionalGeneration):
|
|||||||
("in_proj_ba.", "in_proj_a.", 1),
|
("in_proj_ba.", "in_proj_a.", 1),
|
||||||
]
|
]
|
||||||
|
|
||||||
|
num_experts = self.config.num_experts
|
||||||
|
|
||||||
# Params for weights, fp8 weight scales, fp8 activation scales
|
# Params for weights, fp8 weight scales, fp8 activation scales
|
||||||
# (param_name, weight_name, expert_id, shard_id)
|
# (param_name, weight_name, expert_id, shard_id)
|
||||||
expert_params_mapping = FusedMoE.make_expert_params_mapping(
|
expert_params_mapping = FusedMoE.make_expert_params_mapping(
|
||||||
ckpt_gate_proj_name="gate_proj",
|
ckpt_gate_proj_name="gate_proj",
|
||||||
ckpt_down_proj_name="down_proj",
|
ckpt_down_proj_name="down_proj",
|
||||||
ckpt_up_proj_name="up_proj",
|
ckpt_up_proj_name="up_proj",
|
||||||
num_experts=self.config.num_experts,
|
num_experts=(
|
||||||
|
num_experts
|
||||||
|
if not self.enable_shared_expert_fusion
|
||||||
|
else num_experts + self.num_fused_shared_experts
|
||||||
|
),
|
||||||
)
|
)
|
||||||
|
|
||||||
# Skip loading extra parameters for GPTQ/modelopt models.
|
# Skip loading extra parameters for GPTQ/modelopt models.
|
||||||
@@ -1552,7 +1578,40 @@ class Qwen3_5MoeForConditionalGeneration(Qwen3VLForConditionalGeneration):
|
|||||||
("experts.w2_weight", "experts.down_proj", 0, "w2"),
|
("experts.w2_weight", "experts.down_proj", 0, "w2"),
|
||||||
]
|
]
|
||||||
|
|
||||||
num_experts = self.config.num_experts
|
if self.enable_shared_expert_fusion:
|
||||||
|
"""
|
||||||
|
When shared experts are fused, we need to map the shared experts to routed experts.
|
||||||
|
|
||||||
|
mlp.share_expert.gate_up_proj.weight --> experts.512.gate_up_proj.weight -> experts.w13_weight, expert_id = 512
|
||||||
|
mlp.share_expert.down_proj.weight --> experts.512.down_proj.weight -> experts.w2_weight, expert_id = 512
|
||||||
|
"""
|
||||||
|
fused_expert_params_mapping += [
|
||||||
|
(
|
||||||
|
"experts.w13_",
|
||||||
|
f"experts.{num_experts}.gate_up_proj.",
|
||||||
|
num_experts,
|
||||||
|
"w1",
|
||||||
|
),
|
||||||
|
(
|
||||||
|
"experts.w2_",
|
||||||
|
f"experts.{num_experts}.down_proj.",
|
||||||
|
num_experts,
|
||||||
|
"w2",
|
||||||
|
),
|
||||||
|
## shared experts may contain gate_proj and up_proj instead of gate_up_proj
|
||||||
|
(
|
||||||
|
"experts.w13_",
|
||||||
|
f"experts.{num_experts}.gate_proj.",
|
||||||
|
num_experts,
|
||||||
|
"w1",
|
||||||
|
),
|
||||||
|
(
|
||||||
|
"experts.w13_",
|
||||||
|
f"experts.{num_experts}.up_proj.",
|
||||||
|
num_experts,
|
||||||
|
"w3",
|
||||||
|
),
|
||||||
|
]
|
||||||
|
|
||||||
def load_fused_expert_weights(
|
def load_fused_expert_weights(
|
||||||
name: str,
|
name: str,
|
||||||
@@ -1609,6 +1668,14 @@ class Qwen3_5MoeForConditionalGeneration(Qwen3VLForConditionalGeneration):
|
|||||||
):
|
):
|
||||||
continue
|
continue
|
||||||
|
|
||||||
|
if self.enable_shared_expert_fusion:
|
||||||
|
if "mlp.shared_expert." in name:
|
||||||
|
# Firstly map mlp.shared_expert.xx_proj to mlp.experts.512.xx_proj
|
||||||
|
name = name.replace(
|
||||||
|
"mlp.shared_expert.",
|
||||||
|
f"mlp.experts.{num_experts}.",
|
||||||
|
)
|
||||||
|
|
||||||
for param_name, weight_name, shard_id in stacked_params_mapping:
|
for param_name, weight_name, shard_id in stacked_params_mapping:
|
||||||
if name.endswith("experts.gate_up_proj") or name.endswith(
|
if name.endswith("experts.gate_up_proj") or name.endswith(
|
||||||
"experts.down_proj"
|
"experts.down_proj"
|
||||||
@@ -1657,7 +1724,10 @@ class Qwen3_5MoeForConditionalGeneration(Qwen3VLForConditionalGeneration):
|
|||||||
is_expert_weight = True
|
is_expert_weight = True
|
||||||
name_mapped = name.replace(weight_name, param_name)
|
name_mapped = name.replace(weight_name, param_name)
|
||||||
if is_fused_expert:
|
if is_fused_expert:
|
||||||
|
# is_fused_expert is True, the checkpoint contains gate_up_proj and down_proj for each expert
|
||||||
if "experts.gate_up_proj" in name:
|
if "experts.gate_up_proj" in name:
|
||||||
|
# experts.gate_up_proj contains all 512 routed experts, excluding shared experts
|
||||||
|
# split into w1 and w3
|
||||||
loaded_weight = loaded_weight.chunk(2, dim=-2)
|
loaded_weight = loaded_weight.chunk(2, dim=-2)
|
||||||
load_fused_expert_weights(
|
load_fused_expert_weights(
|
||||||
name_mapped,
|
name_mapped,
|
||||||
@@ -1673,7 +1743,8 @@ class Qwen3_5MoeForConditionalGeneration(Qwen3VLForConditionalGeneration):
|
|||||||
"w3",
|
"w3",
|
||||||
num_experts,
|
num_experts,
|
||||||
)
|
)
|
||||||
else:
|
elif "experts.down_proj" in name:
|
||||||
|
# experts.down_proj contains all 512 routed experts, excluding shared experts
|
||||||
load_fused_expert_weights(
|
load_fused_expert_weights(
|
||||||
name_mapped,
|
name_mapped,
|
||||||
params_dict,
|
params_dict,
|
||||||
@@ -1681,6 +1752,42 @@ class Qwen3_5MoeForConditionalGeneration(Qwen3VLForConditionalGeneration):
|
|||||||
shard_id,
|
shard_id,
|
||||||
num_experts,
|
num_experts,
|
||||||
)
|
)
|
||||||
|
elif self.enable_shared_expert_fusion:
|
||||||
|
# shared experts should be loaded to experts.w13_weight and experts.w2_weight
|
||||||
|
param = params_dict[name_mapped]
|
||||||
|
weight_loader = getattr(
|
||||||
|
param, "weight_loader", default_weight_loader
|
||||||
|
)
|
||||||
|
param = params_dict[name_mapped]
|
||||||
|
if f"{num_experts}.gate_up_proj" in name:
|
||||||
|
# split into w1 and w3
|
||||||
|
loaded_weight = loaded_weight.chunk(2, dim=-2)
|
||||||
|
# load to experts.w13_weight, shard_id = w1, expert_id = 512
|
||||||
|
weight_loader(
|
||||||
|
param,
|
||||||
|
loaded_weight[0],
|
||||||
|
name_mapped,
|
||||||
|
"w1",
|
||||||
|
expert_id,
|
||||||
|
)
|
||||||
|
# load to experts.w13_weight, shard_id = w3, expert_id = 512
|
||||||
|
weight_loader(
|
||||||
|
param,
|
||||||
|
loaded_weight[1],
|
||||||
|
name_mapped,
|
||||||
|
"w3",
|
||||||
|
expert_id,
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
# load down_proj to experts.w2_weight, shard_id = w2, expert_id = 512
|
||||||
|
# Or load gate_proj and up_proj to experts.w13_weight, shard_id = w1/w3, expert_id = 512
|
||||||
|
weight_loader(
|
||||||
|
param,
|
||||||
|
loaded_weight,
|
||||||
|
name_mapped,
|
||||||
|
shard_id,
|
||||||
|
expert_id,
|
||||||
|
)
|
||||||
else:
|
else:
|
||||||
# Skip loading extra parameters for GPTQ models.
|
# Skip loading extra parameters for GPTQ models.
|
||||||
if (
|
if (
|
||||||
|
|||||||
Reference in New Issue
Block a user