[Intel XPU] Enable fused_moe_triton tuning on XPU and add tuned DeepSeek-OCR-2 configs (#28723)

This commit is contained in:
Rahul Vijayaraghavan
2026-09-16 10:20:22 +08:00
committed by GitHub
parent f11cd8ab0e
commit d4ad368ed9
7 changed files with 433 additions and 26 deletions
@@ -103,6 +103,7 @@ def get_model_config(
"DeepseekV3ForCausalLM",
"DeepseekV32ForCausalLM",
"DeepseekV4ForCausalLM",
"DeepseekOCRForCausalLM",
"Glm4MoeForCausalLM",
"Glm4MoeLiteForCausalLM",
"GlmMoeDsaForCausalLM",
@@ -206,6 +207,9 @@ def get_model_config(
# text_config may not carry torch_dtype; fall back to bf16.
torch_dtype = getattr(config, "torch_dtype", None) or torch.bfloat16
num_layers = getattr(config, "num_hidden_layers", 0)
# Only the DeepSeek family replaces leading MoE layers with dense ones.
dense_layers = getattr(config, "first_k_dense_replace", 0)
return {
"num_experts": E,
@@ -215,6 +219,8 @@ def get_model_config(
"dtype": torch_dtype,
"block_shape": block_shape,
"architecture": architecture,
"num_layers": num_layers,
"dense_layers": dense_layers,
}
@@ -24,6 +24,7 @@ from common_utils import (
)
from ray.experimental.tqdm_ray import tqdm
from sglang.kernels.ops.moe.fused_moe_triton_kernels import clear_b_tma_desc_cache
from sglang.srt.layers.moe.moe_runner import MoeRunnerConfig
from sglang.srt.layers.moe.moe_runner.triton_utils.fused_moe import (
get_config_dtype_str,
@@ -34,13 +35,18 @@ from sglang.srt.layers.moe.moe_runner.triton_utils.fused_moe_triton_config impor
get_config_file_name,
)
from sglang.srt.layers.moe.topk import TopKConfig, select_experts
from sglang.srt.server_args import (
ServerArgs,
set_global_server_args_for_scheduler,
from sglang.srt.runtime_context import get_model, get_parallel
from sglang.srt.server_args import ServerArgs, set_global_server_args_for_scheduler
from sglang.srt.utils import (
get_device,
get_device_module,
is_hip,
is_xpu,
)
from sglang.srt.utils import is_hip
_is_hip = is_hip()
_is_xpu = is_xpu()
device_module = get_device_module()
@dataclasses.dataclass
@@ -72,11 +78,12 @@ class KernelWrapper:
expert_ids=moe_input.expert_ids,
num_tokens_post_padded=moe_input.num_tokens_post_padded,
)
torch.cuda.synchronize()
device_module.synchronize()
# Capture 10 invocations with CUDA graph
graph = torch.cuda.CUDAGraph()
with torch.cuda.graph(graph):
# Capture inner_iter invocations into one replayable graph.
graph_cls = torch.xpu.XPUGraph if _is_xpu else torch.cuda.CUDAGraph
graph = graph_cls()
with device_module.graph(graph):
for k in range(self.inner_iter):
moe_input = self.moe_inputs[k]
self.func(
@@ -86,19 +93,19 @@ class KernelWrapper:
expert_ids=moe_input.expert_ids,
num_tokens_post_padded=moe_input.num_tokens_post_padded,
)
torch.cuda.synchronize()
device_module.synchronize()
# Warmup
for _ in range(5):
graph.replay()
torch.cuda.synchronize()
device_module.synchronize()
return graph
def forward_cost(self, try_cnt=2):
time_cost = float("inf")
for _ in range(try_cnt):
start_event = torch.cuda.Event(enable_timing=True)
end_event = torch.cuda.Event(enable_timing=True)
start_event = device_module.Event(enable_timing=True)
end_event = device_module.Event(enable_timing=True)
start_event.record()
if self.use_cuda_graph:
self.graph.replay()
@@ -113,14 +120,19 @@ class KernelWrapper:
num_tokens_post_padded=moe_input.num_tokens_post_padded,
)
end_event.record()
torch.cuda.synchronize()
device_module.synchronize()
time_cost = min(time_cost, start_event.elapsed_time(end_event))
return time_cost
def load_topk_ids(topk_ids_dir, i: int):
num_layers = 61
dense_layers = 3
model_config = get_model_config(
get_model().model_path,
tp_size=get_parallel().tp_size,
ep_size=get_parallel().ep_size,
)
num_layers = model_config["num_layers"]
dense_layers = model_config["dense_layers"]
moe_layers = num_layers - dense_layers
return torch.load(
f"{topk_ids_dir}/topk_ids_layer{i % moe_layers + dense_layers}_idx{i // moe_layers}.pt"
@@ -150,6 +162,8 @@ def benchmark_config(
ncu_enable = os.getenv("NCU_ENABLE", "0") == "1"
if ncu_enable:
num_iters = 1
# Weights here are per-call, so a cached descriptor only pins a dead w1/w2.
clear_b_tma_desc_cache()
init_dtype = torch.float16 if use_fp8_w8a8 else dtype
hidden_states = torch.randn(num_tokens, hidden_size, dtype=dtype)
if use_int8_w8a16 or use_int8_w8a8:
@@ -470,7 +484,7 @@ def benchmark_config(
if build_down:
ts1.append(kernel1.forward_cost()) # down no-tma
ts_tma1.append(kernel1_tma.forward_cost()) # down tma
torch.cuda.synchronize()
device_module.synchronize()
avg = sum(ts0) / (num_iters) * 1000 if ts0 else float("inf")
avg1 = sum(ts1) / (num_iters) * 1000 if ts1 else float("inf")
@@ -536,12 +550,12 @@ class BestConfigTrace:
class BenchmarkWorker:
def __init__(self, seed: int, server_args: ServerArgs) -> None:
torch.set_default_device("cuda")
torch.cuda.manual_seed_all(0)
torch.set_default_device(get_device())
device_module.manual_seed_all(0)
self.seed = seed
# Get the device ID to allocate tensors and kernels
# on the respective GPU.
self.device_id = 0 # int(ray.get_gpu_ids()[0])
self.device_id = 0 if not ray.is_initialized() else int(ray.get_gpu_ids()[0])
set_global_server_args_for_scheduler(server_args)
def benchmark(
@@ -562,9 +576,13 @@ class BenchmarkWorker:
ep_size: int = 1,
enable_up_tma: bool = False,
) -> Tuple[Dict[str, int], float]:
torch.cuda.manual_seed_all(0)
device_module.manual_seed_all(0)
topk_ids_list = [load_topk_ids(topk_ids_dir, i) for i in range(100)]
with torch.cuda.device(self.device_id) if is_hip() else nullcontext():
with (
device_module.device(self.device_id)
if _is_xpu or _is_hip
else nullcontext()
):
if enable_up_tma:
# Two-step: first measure down to determine c_sorted,
# then measure up with the correct c_sorted.
@@ -654,7 +672,11 @@ class BenchmarkWorker:
trace0 = BestConfigTrace("kernel0", down_moe=False, enable_up_tma=False)
trace1 = BestConfigTrace("kernel1", down_moe=True)
with torch.cuda.device(self.device_id) if is_hip() else nullcontext():
with (
device_module.device(self.device_id)
if _is_xpu or _is_hip
else nullcontext()
):
for config in tqdm(search_space):
try:
kt0_no_tma, kt0_tma, kt1_no_tma, kt1_tma = benchmark_config(
@@ -693,7 +715,11 @@ class BenchmarkWorker:
trace1 = BestConfigTrace("kernel1", down_moe=True)
# === Round 1: Down-only ===
with torch.cuda.device(self.device_id) if is_hip() else nullcontext():
with (
device_module.device(self.device_id)
if _is_xpu or _is_hip
else nullcontext()
):
for config in tqdm(search_space, desc="Round 1 (down)"):
try:
_, _, kt1_no_tma, kt1_tma = benchmark_config(
@@ -732,7 +758,11 @@ class BenchmarkWorker:
)
# === Round 2: Up with c_sorted from round 1 ===
with torch.cuda.device(self.device_id) if is_hip() else nullcontext():
with (
device_module.device(self.device_id)
if _is_xpu or _is_hip
else nullcontext()
):
for config in tqdm(search_space, desc="Round 2 (up)"):
try:
kt0_no_tma, kt0_tma, _, _ = benchmark_config(
@@ -804,8 +834,12 @@ class BenchmarkWorker:
print(f"config {i}: {file}")
topk_ids_list = [load_topk_ids(topk_ids_dir, i) for i in range(100)]
torch.cuda.manual_seed_all(0)
with torch.cuda.device(self.device_id) if is_hip() else nullcontext():
device_module.manual_seed_all(0)
with (
device_module.device(self.device_id)
if _is_xpu or _is_hip
else nullcontext()
):
for bs in num_tokens:
kernel_times = []
cfgs = []
+6
View File
@@ -93,6 +93,12 @@ diffusion = [
"xatlas",
]
ray = [
"ray[default]>=2.55.1",
# Ray registers no GPU resource on XPU without dpctl.
"dpctl",
]
tracing = [
"opentelemetry-api",
"opentelemetry-exporter-otlp",
@@ -744,6 +744,11 @@ _B_DESC_CACHE_MAX = 64
_B_DESC_CACHE: OrderedDict[tuple, TensorDescriptor] = OrderedDict()
def clear_b_tma_desc_cache() -> None:
"""Drop all cached B TensorDescriptors, releasing the weights they pin."""
_B_DESC_CACHE.clear()
def _get_b_tma_desc_cached(B: torch.Tensor, block_n: int, block_k: int):
"""
Cache TensorDescriptor for constant weight B.
@@ -0,0 +1,146 @@
{
"1": {
"BLOCK_SIZE_M": 16,
"BLOCK_SIZE_N": 64,
"BLOCK_SIZE_K": 64,
"GROUP_SIZE_M": 16,
"num_warps": 4,
"num_stages": 4
},
"2": {
"BLOCK_SIZE_M": 16,
"BLOCK_SIZE_N": 128,
"BLOCK_SIZE_K": 64,
"GROUP_SIZE_M": 64,
"num_warps": 4,
"num_stages": 3
},
"4": {
"BLOCK_SIZE_M": 16,
"BLOCK_SIZE_N": 64,
"BLOCK_SIZE_K": 64,
"GROUP_SIZE_M": 1,
"num_warps": 4,
"num_stages": 2
},
"8": {
"BLOCK_SIZE_M": 16,
"BLOCK_SIZE_N": 64,
"BLOCK_SIZE_K": 64,
"GROUP_SIZE_M": 1,
"num_warps": 4,
"num_stages": 2
},
"16": {
"BLOCK_SIZE_M": 16,
"BLOCK_SIZE_N": 64,
"BLOCK_SIZE_K": 64,
"GROUP_SIZE_M": 1,
"num_warps": 4,
"num_stages": 3
},
"24": {
"BLOCK_SIZE_M": 16,
"BLOCK_SIZE_N": 64,
"BLOCK_SIZE_K": 64,
"GROUP_SIZE_M": 64,
"num_warps": 4,
"num_stages": 4
},
"32": {
"BLOCK_SIZE_M": 16,
"BLOCK_SIZE_N": 64,
"BLOCK_SIZE_K": 64,
"GROUP_SIZE_M": 1,
"num_warps": 4,
"num_stages": 2
},
"48": {
"BLOCK_SIZE_M": 16,
"BLOCK_SIZE_N": 64,
"BLOCK_SIZE_K": 64,
"GROUP_SIZE_M": 1,
"num_warps": 4,
"num_stages": 4
},
"64": {
"BLOCK_SIZE_M": 16,
"BLOCK_SIZE_N": 64,
"BLOCK_SIZE_K": 64,
"GROUP_SIZE_M": 16,
"num_warps": 4,
"num_stages": 4
},
"96": {
"BLOCK_SIZE_M": 16,
"BLOCK_SIZE_N": 64,
"BLOCK_SIZE_K": 64,
"GROUP_SIZE_M": 64,
"num_warps": 4,
"num_stages": 2
},
"128": {
"BLOCK_SIZE_M": 32,
"BLOCK_SIZE_N": 128,
"BLOCK_SIZE_K": 64,
"GROUP_SIZE_M": 16,
"num_warps": 8,
"num_stages": 3
},
"256": {
"BLOCK_SIZE_M": 32,
"BLOCK_SIZE_N": 64,
"BLOCK_SIZE_K": 64,
"GROUP_SIZE_M": 16,
"num_warps": 4,
"num_stages": 5
},
"512": {
"BLOCK_SIZE_M": 64,
"BLOCK_SIZE_N": 64,
"BLOCK_SIZE_K": 64,
"GROUP_SIZE_M": 64,
"num_warps": 4,
"num_stages": 2
},
"1024": {
"BLOCK_SIZE_M": 64,
"BLOCK_SIZE_N": 128,
"BLOCK_SIZE_K": 64,
"GROUP_SIZE_M": 1,
"num_warps": 8,
"num_stages": 5
},
"1536": {
"BLOCK_SIZE_M": 64,
"BLOCK_SIZE_N": 128,
"BLOCK_SIZE_K": 64,
"GROUP_SIZE_M": 1,
"num_warps": 8,
"num_stages": 2
},
"2048": {
"BLOCK_SIZE_M": 64,
"BLOCK_SIZE_N": 128,
"BLOCK_SIZE_K": 64,
"GROUP_SIZE_M": 1,
"num_warps": 8,
"num_stages": 5
},
"3072": {
"BLOCK_SIZE_M": 64,
"BLOCK_SIZE_N": 128,
"BLOCK_SIZE_K": 64,
"GROUP_SIZE_M": 1,
"num_warps": 8,
"num_stages": 2
},
"4096": {
"BLOCK_SIZE_M": 64,
"BLOCK_SIZE_N": 128,
"BLOCK_SIZE_K": 64,
"GROUP_SIZE_M": 1,
"num_warps": 8,
"num_stages": 5
}
}
@@ -0,0 +1,164 @@
{
"1": {
"BLOCK_SIZE_M": 16,
"BLOCK_SIZE_N": 32,
"BLOCK_SIZE_K": 64,
"GROUP_SIZE_M": 1,
"num_warps": 4,
"num_stages": 3,
"USE_TMA": true
},
"2": {
"BLOCK_SIZE_M": 16,
"BLOCK_SIZE_N": 32,
"BLOCK_SIZE_K": 64,
"GROUP_SIZE_M": 64,
"num_warps": 4,
"num_stages": 2,
"USE_TMA": true
},
"4": {
"BLOCK_SIZE_M": 16,
"BLOCK_SIZE_N": 32,
"BLOCK_SIZE_K": 64,
"GROUP_SIZE_M": 1,
"num_warps": 4,
"num_stages": 2,
"USE_TMA": true
},
"8": {
"BLOCK_SIZE_M": 16,
"BLOCK_SIZE_N": 64,
"BLOCK_SIZE_K": 64,
"GROUP_SIZE_M": 1,
"num_warps": 8,
"num_stages": 2,
"USE_TMA": true
},
"16": {
"BLOCK_SIZE_M": 16,
"BLOCK_SIZE_N": 32,
"BLOCK_SIZE_K": 64,
"GROUP_SIZE_M": 64,
"num_warps": 8,
"num_stages": 3,
"USE_TMA": true
},
"24": {
"BLOCK_SIZE_M": 16,
"BLOCK_SIZE_N": 32,
"BLOCK_SIZE_K": 64,
"GROUP_SIZE_M": 64,
"num_warps": 8,
"num_stages": 3,
"USE_TMA": true
},
"32": {
"BLOCK_SIZE_M": 16,
"BLOCK_SIZE_N": 32,
"BLOCK_SIZE_K": 64,
"GROUP_SIZE_M": 64,
"num_warps": 8,
"num_stages": 2,
"USE_TMA": true
},
"48": {
"BLOCK_SIZE_M": 16,
"BLOCK_SIZE_N": 32,
"BLOCK_SIZE_K": 64,
"GROUP_SIZE_M": 64,
"num_warps": 8,
"num_stages": 4,
"USE_TMA": true
},
"64": {
"BLOCK_SIZE_M": 16,
"BLOCK_SIZE_N": 64,
"BLOCK_SIZE_K": 64,
"GROUP_SIZE_M": 1,
"num_warps": 8,
"num_stages": 2,
"USE_TMA": true
},
"96": {
"BLOCK_SIZE_M": 16,
"BLOCK_SIZE_N": 64,
"BLOCK_SIZE_K": 128,
"GROUP_SIZE_M": 1,
"num_warps": 8,
"num_stages": 2,
"USE_TMA": true
},
"128": {
"BLOCK_SIZE_M": 32,
"BLOCK_SIZE_N": 32,
"BLOCK_SIZE_K": 64,
"GROUP_SIZE_M": 64,
"num_warps": 8,
"num_stages": 3,
"USE_TMA": true
},
"256": {
"BLOCK_SIZE_M": 32,
"BLOCK_SIZE_N": 32,
"BLOCK_SIZE_K": 64,
"GROUP_SIZE_M": 1,
"num_warps": 4,
"num_stages": 3,
"USE_TMA": true
},
"512": {
"BLOCK_SIZE_M": 64,
"BLOCK_SIZE_N": 32,
"BLOCK_SIZE_K": 64,
"GROUP_SIZE_M": 1,
"num_warps": 8,
"num_stages": 4,
"USE_TMA": true
},
"1024": {
"BLOCK_SIZE_M": 64,
"BLOCK_SIZE_N": 32,
"BLOCK_SIZE_K": 64,
"GROUP_SIZE_M": 1,
"num_warps": 8,
"num_stages": 3,
"USE_TMA": true
},
"1536": {
"BLOCK_SIZE_M": 64,
"BLOCK_SIZE_N": 32,
"BLOCK_SIZE_K": 64,
"GROUP_SIZE_M": 1,
"num_warps": 8,
"num_stages": 2,
"USE_TMA": true
},
"2048": {
"BLOCK_SIZE_M": 64,
"BLOCK_SIZE_N": 32,
"BLOCK_SIZE_K": 64,
"GROUP_SIZE_M": 1,
"num_warps": 8,
"num_stages": 3,
"USE_TMA": true
},
"3072": {
"BLOCK_SIZE_M": 64,
"BLOCK_SIZE_N": 32,
"BLOCK_SIZE_K": 64,
"GROUP_SIZE_M": 1,
"num_warps": 8,
"num_stages": 2,
"USE_TMA": true
},
"4096": {
"BLOCK_SIZE_M": 64,
"BLOCK_SIZE_N": 32,
"BLOCK_SIZE_K": 64,
"GROUP_SIZE_M": 1,
"num_warps": 8,
"num_stages": 2,
"USE_TMA": true
}
}
@@ -29,6 +29,8 @@ def test_get_model_config_supports_kimi_vl():
n_routed_experts=64,
num_experts_per_tok=6,
moe_intermediate_size=1408,
num_hidden_layers=27,
first_k_dense_replace=1,
torch_dtype=torch.bfloat16,
)
model_config = SimpleNamespace(
@@ -42,6 +44,8 @@ def test_get_model_config_supports_kimi_vl():
"moonshotai/Kimi-VL-A3B-Instruct", tp_size=4, ep_size=4
)
# The layer counts come from the text config, not the outer one, which has
# neither field.
assert tuned_config == {
"num_experts": 16,
"topk": 6,
@@ -50,6 +54,48 @@ def test_get_model_config_supports_kimi_vl():
"dtype": torch.bfloat16,
"block_shape": None,
"architecture": "KimiVLForConditionalGeneration",
"num_layers": 27,
"dense_layers": 1,
}
def test_get_model_config_reports_deepseek_ocr_layer_layout():
"""``load_topk_ids`` indexes recorded topk_ids per MoE layer.
It derives the MoE layer count as ``num_layers - dense_layers``, so a config
that reports either as 0 makes the tuner read the wrong (or no) recording.
"""
common_utils = _load_common_utils()
text_config = SimpleNamespace(
hidden_size=1280,
n_routed_experts=64,
num_experts_per_tok=6,
moe_intermediate_size=896,
num_hidden_layers=12,
first_k_dense_replace=1,
torch_dtype=torch.bfloat16,
)
model_config = SimpleNamespace(
architectures=["DeepseekOCRForCausalLM"],
text_config=text_config,
get_text_config=lambda: text_config,
)
with patch.object(common_utils, "get_config", return_value=model_config):
tuned_config = common_utils.get_model_config(
"deepseek-ai/DeepSeek-OCR-2", tp_size=1, ep_size=1
)
assert tuned_config == {
"num_experts": 64,
"topk": 6,
"hidden_size": 1280,
"shard_intermediate_size": 1792,
"dtype": torch.bfloat16,
"block_shape": None,
"architecture": "DeepseekOCRForCausalLM",
"num_layers": 12,
"dense_layers": 1,
}