[BugFix] weight load bug when checkpoint expert.gate and exepert.up_proj are not fused (#13113)
Co-authored-by: yuechguo <yuechguo@amd.com>
This commit is contained in:
@@ -32,6 +32,8 @@ if _is_cuda:
|
|||||||
if _is_npu:
|
if _is_npu:
|
||||||
import torch_npu
|
import torch_npu
|
||||||
|
|
||||||
|
_use_aiter = get_bool_env_var("SGLANG_USE_AITER") and _is_hip
|
||||||
|
|
||||||
from sglang.srt.distributed import (
|
from sglang.srt.distributed import (
|
||||||
split_tensor_along_last_dim,
|
split_tensor_along_last_dim,
|
||||||
tensor_model_parallel_all_gather,
|
tensor_model_parallel_all_gather,
|
||||||
@@ -55,10 +57,6 @@ ROTARY_EMBED_CLASSES = {
|
|||||||
"normal": apply_rotary_pos_emb,
|
"normal": apply_rotary_pos_emb,
|
||||||
}
|
}
|
||||||
|
|
||||||
_use_aiter = get_bool_env_var("SGLANG_USE_AITER") and _is_hip
|
|
||||||
if _use_aiter:
|
|
||||||
from aiter import flash_attn_varlen_func as aiter_flash_attn_varlen_func
|
|
||||||
|
|
||||||
|
|
||||||
@dataclasses.dataclass
|
@dataclasses.dataclass
|
||||||
class SingletonCache:
|
class SingletonCache:
|
||||||
@@ -348,8 +346,16 @@ class VisionAiterAttention(nn.Module):
|
|||||||
self,
|
self,
|
||||||
**kwargs,
|
**kwargs,
|
||||||
):
|
):
|
||||||
if not _use_aiter:
|
if not _is_hip:
|
||||||
raise Exception("aiter_attn is only available for AMD")
|
raise Exception("aiter_attn is only available for AMD")
|
||||||
|
try:
|
||||||
|
from aiter import flash_attn_varlen_func as aiter_flash_attn_varlen_func
|
||||||
|
except ImportError as e:
|
||||||
|
raise ImportError(
|
||||||
|
"aiter is AMD specific kernel library. Please make sure aiter is installed on your AMD device."
|
||||||
|
) from e
|
||||||
|
|
||||||
|
self.flash_attn_varlen_func = aiter_flash_attn_varlen_func
|
||||||
super().__init__()
|
super().__init__()
|
||||||
|
|
||||||
def forward(
|
def forward(
|
||||||
@@ -375,7 +381,7 @@ class VisionAiterAttention(nn.Module):
|
|||||||
seq_lens = cu_seqlens[1:] - cu_seqlens[:-1]
|
seq_lens = cu_seqlens[1:] - cu_seqlens[:-1]
|
||||||
max_seqlen = seq_lens.max().item()
|
max_seqlen = seq_lens.max().item()
|
||||||
|
|
||||||
return aiter_flash_attn_varlen_func(
|
return self.flash_attn_varlen_func(
|
||||||
q=q,
|
q=q,
|
||||||
k=k,
|
k=k,
|
||||||
v=v,
|
v=v,
|
||||||
@@ -590,11 +596,11 @@ class VisionAttention(nn.Module):
|
|||||||
backend = "fa3"
|
backend = "fa3"
|
||||||
else:
|
else:
|
||||||
backend = "triton_attn"
|
backend = "triton_attn"
|
||||||
elif _use_aiter:
|
elif _is_hip:
|
||||||
if get_device_capability() < (9, 4):
|
if get_device_capability() >= (9, 4) and _use_aiter:
|
||||||
backend = "triton_attn"
|
|
||||||
else:
|
|
||||||
backend = "aiter_attn"
|
backend = "aiter_attn"
|
||||||
|
else:
|
||||||
|
backend = "triton_attn"
|
||||||
else:
|
else:
|
||||||
backend = "sdpa"
|
backend = "sdpa"
|
||||||
if backend == "fa3" and is_blackwell():
|
if backend == "fa3" and is_blackwell():
|
||||||
|
|||||||
@@ -521,6 +521,11 @@ class FusedMoE(torch.nn.Module):
|
|||||||
|
|
||||||
global_expert_location_metadata = get_global_expert_location_metadata()
|
global_expert_location_metadata = get_global_expert_location_metadata()
|
||||||
if global_expert_location_metadata is None:
|
if global_expert_location_metadata is None:
|
||||||
|
if not getattr(param, "_sglang_require_global_experts", False):
|
||||||
|
expert_id = self._map_global_expert_id_to_local_expert_id(expert_id)
|
||||||
|
if expert_id == -1:
|
||||||
|
return
|
||||||
|
|
||||||
self._weight_loader_impl(
|
self._weight_loader_impl(
|
||||||
param=param,
|
param=param,
|
||||||
loaded_weight=loaded_weight,
|
loaded_weight=loaded_weight,
|
||||||
|
|||||||
@@ -22,10 +22,6 @@ import torch
|
|||||||
import torch.nn as nn
|
import torch.nn as nn
|
||||||
|
|
||||||
from sglang.srt.configs.qwen3_vl import Qwen3VLMoeConfig, Qwen3VLMoeTextConfig
|
from sglang.srt.configs.qwen3_vl import Qwen3VLMoeConfig, Qwen3VLMoeTextConfig
|
||||||
from sglang.srt.distributed import (
|
|
||||||
get_moe_expert_parallel_world_size,
|
|
||||||
get_tensor_model_parallel_rank,
|
|
||||||
)
|
|
||||||
from sglang.srt.layers.moe.fused_moe_triton.layer import FusedMoE
|
from sglang.srt.layers.moe.fused_moe_triton.layer import FusedMoE
|
||||||
from sglang.srt.layers.quantization.base_config import QuantizationConfig
|
from sglang.srt.layers.quantization.base_config import QuantizationConfig
|
||||||
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, PPProxyTensors
|
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, PPProxyTensors
|
||||||
@@ -127,34 +123,16 @@ def load_fused_expert_weights(
|
|||||||
param = params_dict[name]
|
param = params_dict[name]
|
||||||
# weight_loader = typing.cast(Callable[..., bool], param.weight_loader)
|
# weight_loader = typing.cast(Callable[..., bool], param.weight_loader)
|
||||||
weight_loader = param.weight_loader
|
weight_loader = param.weight_loader
|
||||||
ep_rank = get_tensor_model_parallel_rank()
|
# let ep moe layer to gracefully handle expert_ids that do not belong to local moe rank
|
||||||
ep_size = get_moe_expert_parallel_world_size()
|
for expert_id in range(num_experts):
|
||||||
if ep_size == 1:
|
curr_expert_weight = loaded_weight[expert_id]
|
||||||
for expert_id in range(num_experts):
|
weight_loader(
|
||||||
curr_expert_weight = loaded_weight[expert_id]
|
param,
|
||||||
weight_loader(
|
curr_expert_weight,
|
||||||
param,
|
name,
|
||||||
curr_expert_weight,
|
shard_id,
|
||||||
name,
|
expert_id,
|
||||||
shard_id,
|
|
||||||
expert_id,
|
|
||||||
)
|
|
||||||
else:
|
|
||||||
experts_per_ep = num_experts // ep_size
|
|
||||||
start_expert = ep_rank * experts_per_ep
|
|
||||||
end_expert = (
|
|
||||||
(ep_rank + 1) * experts_per_ep if ep_rank != ep_size - 1 else num_experts
|
|
||||||
)
|
)
|
||||||
|
|
||||||
for idx, expert_id in enumerate(range(start_expert, end_expert)):
|
|
||||||
curr_expert_weight = loaded_weight[expert_id]
|
|
||||||
weight_loader(
|
|
||||||
param,
|
|
||||||
curr_expert_weight,
|
|
||||||
name,
|
|
||||||
shard_id,
|
|
||||||
idx,
|
|
||||||
)
|
|
||||||
return True
|
return True
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user