[LFM2-MoE] Support Transformers v5 packed MoE expert weights (#29659)
Co-authored-by: Piotr Mazurek <piotr.mazurek@liquid.ai> Co-authored-by: Changyi Yang <changyiyang2023@gmail.com>
This commit is contained in:
co-authored by
Piotr Mazurek
Changyi Yang
parent
b276a9acee
commit
26b15a0825
@@ -595,6 +595,82 @@ class Lfm2MoeForCausalLM(nn.Module):
|
|||||||
if "feed_forward.w2" in name and "experts" not in name:
|
if "feed_forward.w2" in name and "experts" not in name:
|
||||||
name = name.replace("feed_forward.w2", "feed_forward.down_proj")
|
name = name.replace("feed_forward.w2", "feed_forward.down_proj")
|
||||||
|
|
||||||
|
# Transformers >= v5.0 packs MoE expert weights into a single 3D tensor
|
||||||
|
# per projection (experts.gate_up_proj / experts.down_proj) instead of
|
||||||
|
# per-expert weights (experts.{i}.w{1,2,3}.weight). This is the layout an
|
||||||
|
# in-memory Transformers model exposes -- e.g. the update_weights_from_tensor
|
||||||
|
# / RLHF weight-sync path -- so map the packed tensors onto the fused
|
||||||
|
# FusedMoE params (w13_weight / w2_weight) per expert. LFM2-MoE packs
|
||||||
|
# out-features-major (gate_up_proj as [num_experts, 2 * intermediate,
|
||||||
|
# hidden], down_proj as [num_experts, hidden, intermediate]), matching the
|
||||||
|
# FusedMoE layout, so no transpose is needed.
|
||||||
|
if "feed_forward.experts.gate_up_proj" in name:
|
||||||
|
fused_name = name
|
||||||
|
if fused_name.endswith(".weight"):
|
||||||
|
fused_name = fused_name[: -len(".weight")]
|
||||||
|
fused_name = fused_name.replace(
|
||||||
|
"feed_forward.experts.gate_up_proj",
|
||||||
|
"feed_forward.experts.w13_weight",
|
||||||
|
)
|
||||||
|
if fused_name in params_dict:
|
||||||
|
if loaded_weight.dim() != 3:
|
||||||
|
raise ValueError(
|
||||||
|
f"Expected a 3D packed tensor for {name}, got "
|
||||||
|
f"{loaded_weight.dim()}D {tuple(loaded_weight.shape)}"
|
||||||
|
)
|
||||||
|
param = params_dict[fused_name]
|
||||||
|
weight_loader = param.weight_loader
|
||||||
|
if loaded_weight.shape[1] % 2 != 0:
|
||||||
|
raise ValueError(
|
||||||
|
f"Invalid gate_up_proj shape for {name}: "
|
||||||
|
f"{tuple(loaded_weight.shape)}"
|
||||||
|
)
|
||||||
|
w1, w3 = loaded_weight.chunk(2, dim=1)
|
||||||
|
for expert_id in range(w1.shape[0]):
|
||||||
|
weight_loader(
|
||||||
|
param,
|
||||||
|
w1[expert_id],
|
||||||
|
fused_name,
|
||||||
|
shard_id="w1",
|
||||||
|
expert_id=expert_id,
|
||||||
|
)
|
||||||
|
weight_loader(
|
||||||
|
param,
|
||||||
|
w3[expert_id],
|
||||||
|
fused_name,
|
||||||
|
shard_id="w3",
|
||||||
|
expert_id=expert_id,
|
||||||
|
)
|
||||||
|
loaded_params.add(fused_name)
|
||||||
|
continue
|
||||||
|
|
||||||
|
if "feed_forward.experts.down_proj" in name:
|
||||||
|
fused_name = name
|
||||||
|
if fused_name.endswith(".weight"):
|
||||||
|
fused_name = fused_name[: -len(".weight")]
|
||||||
|
fused_name = fused_name.replace(
|
||||||
|
"feed_forward.experts.down_proj",
|
||||||
|
"feed_forward.experts.w2_weight",
|
||||||
|
)
|
||||||
|
if fused_name in params_dict:
|
||||||
|
if loaded_weight.dim() != 3:
|
||||||
|
raise ValueError(
|
||||||
|
f"Expected a 3D packed tensor for {name}, got "
|
||||||
|
f"{loaded_weight.dim()}D {tuple(loaded_weight.shape)}"
|
||||||
|
)
|
||||||
|
param = params_dict[fused_name]
|
||||||
|
weight_loader = param.weight_loader
|
||||||
|
for expert_id in range(loaded_weight.shape[0]):
|
||||||
|
weight_loader(
|
||||||
|
param,
|
||||||
|
loaded_weight[expert_id],
|
||||||
|
fused_name,
|
||||||
|
shard_id="w2",
|
||||||
|
expert_id=expert_id,
|
||||||
|
)
|
||||||
|
loaded_params.add(fused_name)
|
||||||
|
continue
|
||||||
|
|
||||||
# Handle stacked params (QKV, dense MLP gate_up)
|
# Handle stacked params (QKV, dense MLP gate_up)
|
||||||
for param_name, weight_name, shard_id in stacked_params_mapping:
|
for param_name, weight_name, shard_id in stacked_params_mapping:
|
||||||
if weight_name not in name:
|
if weight_name not in name:
|
||||||
|
|||||||
Reference in New Issue
Block a user