[AMD]Integrate aiter's fused_topk for softmax scoring in topk function (#21421)
Co-authored-by: Chen, Todd <zhenchen@amd.com>
This commit is contained in:
co-authored by
Chen, Todd
parent
a305964159
commit
fd535942ac
@@ -129,6 +129,7 @@ if _is_cuda or _is_hip or _is_xpu:
|
|||||||
if _use_aiter:
|
if _use_aiter:
|
||||||
try:
|
try:
|
||||||
from aiter import biased_grouped_topk as aiter_biased_grouped_topk
|
from aiter import biased_grouped_topk as aiter_biased_grouped_topk
|
||||||
|
from aiter.fused_moe import fused_topk as aiter_fused_topk
|
||||||
except ImportError:
|
except ImportError:
|
||||||
raise ImportError("aiter is required when SGLANG_USE_AITER is set to True")
|
raise ImportError("aiter is required when SGLANG_USE_AITER is set to True")
|
||||||
|
|
||||||
@@ -511,12 +512,24 @@ def fused_topk(
|
|||||||
topk_ids = torch.empty(M, topk, dtype=torch.int32, device=hidden_states.device)
|
topk_ids = torch.empty(M, topk, dtype=torch.int32, device=hidden_states.device)
|
||||||
|
|
||||||
if scoring_func == "softmax":
|
if scoring_func == "softmax":
|
||||||
topk_softmax(
|
if _use_aiter:
|
||||||
topk_weights,
|
|
||||||
topk_ids,
|
# Use fused_topk instead of topk_softmax to auto dispatch to the correct kernel
|
||||||
gating_output,
|
topk_weights, topk_ids = aiter_fused_topk(
|
||||||
renormalize,
|
hidden_states,
|
||||||
)
|
gating_output,
|
||||||
|
topk,
|
||||||
|
renormalize,
|
||||||
|
topk_ids=topk_ids,
|
||||||
|
topk_weights=topk_weights,
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
topk_softmax(
|
||||||
|
topk_weights,
|
||||||
|
topk_ids,
|
||||||
|
gating_output,
|
||||||
|
renormalize,
|
||||||
|
)
|
||||||
elif scoring_func == "sigmoid":
|
elif scoring_func == "sigmoid":
|
||||||
topk_sigmoid(
|
topk_sigmoid(
|
||||||
topk_weights,
|
topk_weights,
|
||||||
|
|||||||
Reference in New Issue
Block a user