fix ep
This commit is contained in:
@@ -154,7 +154,19 @@ def patch_fused_moe_for_ep():
|
|||||||
|
|
||||||
def _ep_weight_loader(self, param, loaded_weight, weight_name,
|
def _ep_weight_loader(self, param, loaded_weight, weight_name,
|
||||||
shard_id, expert_id):
|
shard_id, expert_id):
|
||||||
"""EP-aware weight_loader: skip experts not owned by this rank."""
|
"""EP-aware weight_loader: skip experts not owned by this rank.
|
||||||
|
|
||||||
|
Critical fix: the original weight_loader calls
|
||||||
|
get_tensor_model_parallel_rank() internally to compute tp_rank,
|
||||||
|
then uses it to narrow() the loaded weight for TP sharding.
|
||||||
|
Under EP we pass tp_size=1 (no TP split on MoE), so weights are
|
||||||
|
full-sized, but tp_rank still returns the real rank (0-3).
|
||||||
|
rank >= 1 causes narrow(dim, 512*rank, 512) to exceed dim size 512.
|
||||||
|
|
||||||
|
Solution: temporarily monkey-patch get_tensor_model_parallel_rank
|
||||||
|
to return 0 during EP weight loading, so narrow() always starts
|
||||||
|
at offset 0 and loads the full (unsplit) weight.
|
||||||
|
"""
|
||||||
if not getattr(self, '_ep_enabled', False):
|
if not getattr(self, '_ep_enabled', False):
|
||||||
return _orig_weight_loader(self, param, loaded_weight,
|
return _orig_weight_loader(self, param, loaded_weight,
|
||||||
weight_name, shard_id, expert_id)
|
weight_name, shard_id, expert_id)
|
||||||
@@ -169,10 +181,16 @@ def patch_fused_moe_for_ep():
|
|||||||
# Remap global expert_id to local index
|
# Remap global expert_id to local index
|
||||||
local_expert_id = expert_id - start
|
local_expert_id = expert_id - start
|
||||||
|
|
||||||
# Call original weight_loader with local expert_id.
|
# Monkey-patch tp rank to 0 during load so narrow() doesn't OOB.
|
||||||
# The param tensor is sized for local experts only.
|
# The weight is already full-sized (tp_size=1), no split needed.
|
||||||
_orig_weight_loader(self, param, loaded_weight, weight_name,
|
import vllm.model_executor.layers.fused_moe.layer as _fused_moe_mod
|
||||||
shard_id, local_expert_id)
|
_real_get_tp_rank = _fused_moe_mod.get_tensor_model_parallel_rank
|
||||||
|
_fused_moe_mod.get_tensor_model_parallel_rank = lambda: 0
|
||||||
|
try:
|
||||||
|
_orig_weight_loader(self, param, loaded_weight, weight_name,
|
||||||
|
shard_id, local_expert_id)
|
||||||
|
finally:
|
||||||
|
_fused_moe_mod.get_tensor_model_parallel_rank = _real_get_tp_rank
|
||||||
|
|
||||||
_orig_forward = FusedMoE.forward
|
_orig_forward = FusedMoE.forward
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user