[CPU] bypass scoring_func argument in topk for cpu device (#31110)
This commit is contained in:
@@ -700,7 +700,25 @@ def fused_topk_cpu(
|
|||||||
renormalize: bool,
|
renormalize: bool,
|
||||||
correction_bias: torch.Tensor = None,
|
correction_bias: torch.Tensor = None,
|
||||||
scoring_func: str = "softmax",
|
scoring_func: str = "softmax",
|
||||||
|
routed_scaling_factor: Optional[float] = None,
|
||||||
|
apply_routed_scaling_factor_on_output: Optional[bool] = False,
|
||||||
|
num_fused_shared_experts: int = 0,
|
||||||
|
packed_out: Optional[torch.Tensor] = None,
|
||||||
|
num_token_non_padded: Optional[torch.Tensor] = None,
|
||||||
):
|
):
|
||||||
|
if num_fused_shared_experts != 0:
|
||||||
|
raise ValueError(
|
||||||
|
f"num_fused_shared_experts must be 0 for CPU fused topk, got: {num_fused_shared_experts}"
|
||||||
|
)
|
||||||
|
if apply_routed_scaling_factor_on_output:
|
||||||
|
raise ValueError(
|
||||||
|
"apply_routed_scaling_factor_on_output is not supported for CPU fused topk"
|
||||||
|
)
|
||||||
|
if packed_out is not None:
|
||||||
|
raise ValueError("packed_out is not supported for CPU fused topk")
|
||||||
|
if num_token_non_padded is not None:
|
||||||
|
raise ValueError("num_token_non_padded is not supported for CPU fused topk")
|
||||||
|
|
||||||
# TODO: add c++ kernel for cpu
|
# TODO: add c++ kernel for cpu
|
||||||
# The topk_softmax_cpu kernel only handles vanilla softmax scoring with no
|
# The topk_softmax_cpu kernel only handles vanilla softmax scoring with no
|
||||||
# correction bias. Fall back to the torch-native impl for the rest
|
# correction bias. Fall back to the torch-native impl for the rest
|
||||||
@@ -963,8 +981,12 @@ def grouped_topk_cpu(
|
|||||||
num_fused_shared_experts: int = 0,
|
num_fused_shared_experts: int = 0,
|
||||||
routed_scaling_factor: Optional[float] = None,
|
routed_scaling_factor: Optional[float] = None,
|
||||||
apply_routed_scaling_factor_on_output: Optional[bool] = False,
|
apply_routed_scaling_factor_on_output: Optional[bool] = False,
|
||||||
|
scoring_func: str = "softmax",
|
||||||
):
|
):
|
||||||
assert not apply_routed_scaling_factor_on_output
|
assert not apply_routed_scaling_factor_on_output
|
||||||
|
if scoring_func != "softmax":
|
||||||
|
raise ValueError(f"Unsupported scoring function: {scoring_func}")
|
||||||
|
|
||||||
return torch.ops.sgl_kernel.grouped_topk_cpu(
|
return torch.ops.sgl_kernel.grouped_topk_cpu(
|
||||||
hidden_states,
|
hidden_states,
|
||||||
gating_output,
|
gating_output,
|
||||||
|
|||||||
@@ -7,7 +7,11 @@ from sglang.test.server_fixtures.spec_eagle_fixture import Eagle3Base
|
|||||||
|
|
||||||
# Estimated: 2 sequential 8B server launches + one 4-prompt greedy method
|
# Estimated: 2 sequential 8B server launches + one 4-prompt greedy method
|
||||||
# (CUDA sibling: 360); tune from CI TIMINGS once it has run there.
|
# (CUDA sibling: 360); tune from CI TIMINGS once it has run there.
|
||||||
register_cpu_ci(est_time=480, suite="base-b-test-cpu")
|
register_cpu_ci(
|
||||||
|
est_time=480,
|
||||||
|
suite="base-b-test-cpu",
|
||||||
|
disabled="EAGLE3 numerical parity mismatches on CPU intel_amx",
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
class TestEagle3ParityCPU(SpecParityKit, Eagle3Base):
|
class TestEagle3ParityCPU(SpecParityKit, Eagle3Base):
|
||||||
|
|||||||
Reference in New Issue
Block a user