[Intel GPU] DeepSeek V4 13/N: use sgl-kernel implementation of kernels in V2 Compressor to run on XPU (#28439)

Signed-off-by: P V R K Jyothendra Varma <polisettyvarma@gmail.com>
Signed-off-by: P V R K Jyothendra Varma <polisetty.v.r.k.jyothendra.varma@intel.com>
Co-authored-by: Rahul Vijayaraghavan <rahul.vijayaraghavan@intel.com>
This commit is contained in:
Polisetty V R K Jyothendra Varma
2026-07-17 09:16:04 +08:00
committed by GitHub
co-authored by Rahul Vijayaraghavan
parent 302c3b97d2
commit 37f94cb7a0
5 changed files with 225 additions and 98 deletions
+97 -33
View File
@@ -10,10 +10,30 @@ from sglang.jit_kernel.utils import (
load_jit,
make_cpp_args,
)
from sglang.srt.utils import is_hip
from sglang.srt.utils import is_hip, is_xpu
from .utils import make_name
_is_xpu = is_xpu()
if _is_xpu:
from sgl_kernel import compress_norm_rope_store as compress_norm_rope_store_xpu
from sgl_kernel import (
flash_compress4_decode,
flash_compress4_prefill,
flash_compress128_decode,
flash_compress128_prefill,
plan_compress_decode,
plan_compress_decode_legacy,
plan_compress_prefill,
plan_compress_prefill_legacy,
)
_XPU_COMPRESS_FNS = {
4: (flash_compress4_decode, flash_compress4_prefill),
128: (flash_compress128_decode, flash_compress128_prefill),
}
if TYPE_CHECKING:
from tvm_ffi.module import Module
@@ -133,8 +153,13 @@ class CompressorDecodePlan(NamedTuple):
swa_page_size: int,
ring_size: int,
) -> CompressorDecodePlan:
module = _jit_compress_plan_module()
plan_d = module.plan_decode(
if _is_xpu:
fn = plan_compress_decode
else:
module = _jit_compress_plan_module()
fn = module.plan_decode
plan_d = fn(
req_pool_indices,
req_to_token,
full_to_state,
@@ -151,8 +176,13 @@ class CompressorDecodePlan(NamedTuple):
req_pool_indices: torch.Tensor,
seq_lens: torch.Tensor,
) -> CompressorDecodePlan:
module = _jit_compress_plan_module()
plan_d = module.plan_decode_legacy(req_pool_indices, seq_lens, compress_ratio)
if _is_xpu:
fn = plan_compress_decode_legacy
else:
module = _jit_compress_plan_module()
fn = module.plan_decode_legacy
plan_d = fn(req_pool_indices, seq_lens, compress_ratio)
return CompressorDecodePlan(compress_ratio, torch.from_dlpack(plan_d))
@staticmethod
@@ -208,7 +238,7 @@ class CompressorPrefillPlan(NamedTuple):
num_q_tokens: int,
use_cuda_graph: bool = False,
) -> CompressorPrefillPlan:
is_gpu_input = seq_lens.device.type == "cuda"
is_gpu_input = seq_lens.device.type in ["cuda", "xpu"]
pin_buffer = torch.empty(
0 if is_gpu_input else num_q_tokens * _PREFILL_PLAN_BYTES,
dtype=torch.uint8,
@@ -229,7 +259,13 @@ class CompressorPrefillPlan(NamedTuple):
pin_buffer,
)
module = _jit_compress_plan_module()
plan_c, plan_w = module.plan_prefill(
if _is_xpu:
fn = plan_compress_prefill
else:
module = _jit_compress_plan_module()
fn = module.plan_prefill
plan_c, plan_w = fn(
req_pool_indices,
req_to_token,
full_to_state,
@@ -244,8 +280,8 @@ class CompressorPrefillPlan(NamedTuple):
)
return CompressorPrefillPlan(
compress_ratio,
torch.from_dlpack(plan_c),
torch.from_dlpack(plan_w),
torch.from_dlpack(plan_c) if not _is_xpu else plan_c,
torch.from_dlpack(plan_w) if not _is_xpu else plan_w,
pin_buffer,
)
@@ -264,8 +300,13 @@ class CompressorPrefillPlan(NamedTuple):
dtype=torch.uint8,
pin_memory=True,
)
module = _jit_compress_plan_module()
plan_c, plan_w = module.plan_prefill_legacy(
if _is_xpu:
fn = plan_compress_prefill_legacy
else:
module = _jit_compress_plan_module()
fn = module.plan_prefill_legacy
plan_c, plan_w = fn(
req_pool_indices,
seq_lens,
extend_lens,
@@ -276,8 +317,8 @@ class CompressorPrefillPlan(NamedTuple):
)
return CompressorPrefillPlan(
compress_ratio,
torch.from_dlpack(plan_c),
torch.from_dlpack(plan_w),
torch.from_dlpack(plan_c) if not _is_xpu else plan_c,
torch.from_dlpack(plan_w) if not _is_xpu else plan_w,
pin_buffer,
)
@@ -346,11 +387,19 @@ def compress_forward(
assert compress_ratio == 128 and head_dim == 512
module = _jit_compress_128_online_module(512, kv_score_buffer.dtype)
else:
dtype_in, dtype_out = kv_score_input.dtype, out.dtype
module = _jit_compress_module(
head_dim, kv_score_buffer.dtype, dtype_in, dtype_out, compress_ratio
)
fn = module.decode if plan.is_decode else module.prefill
if _is_xpu:
decode_fn, prefill_fn = _XPU_COMPRESS_FNS[compress_ratio]
else:
dtype_in, dtype_out = kv_score_input.dtype, out.dtype
module = _jit_compress_module(
head_dim, kv_score_buffer.dtype, dtype_in, dtype_out, compress_ratio
)
if _is_xpu:
fn = decode_fn if plan.is_decode else prefill_fn
else:
fn = module.decode if plan.is_decode else module.prefill
fn(kv_score_buffer, kv_score_input, out, ape, *plan[1:3])
return out
@@ -371,18 +420,33 @@ def compress_norm_rope_store(
if use_fp4:
assert kv.shape[-1] == 128
freq_cis = torch.view_as_real(freq_cis).flatten(-2)
module = _jit_compress_norm_rope_module(
kv.dtype, kv.shape[-1], freq_cis.shape[-1], page_size, bf16_store
)
fn = module.forward_fp4 if use_fp4 else module.forward
fn(
kv,
plan[1],
norm_weight,
norm_eps,
freq_cis,
out_loc,
kvcache,
plan.is_decode,
plan.compress_ratio,
)
if _is_xpu:
compress_norm_rope_store_xpu(
kv,
plan[1],
norm_weight,
norm_eps,
freq_cis,
out_loc,
kvcache,
plan.is_decode,
plan.compress_ratio,
page_size,
use_fp4,
)
else:
module = _jit_compress_norm_rope_module(
kv.dtype, kv.shape[-1], freq_cis.shape[-1], page_size, bf16_store
)
fn = module.forward_fp4 if use_fp4 else module.forward
fn(
kv,
plan[1],
norm_weight,
norm_eps,
freq_cis,
out_loc,
kvcache,
plan.is_decode,
plan.compress_ratio,
)
@@ -6,6 +6,7 @@ from typing import List, Literal, Optional, Tuple
import torch
from sglang.jit_kernel.dsv4 import CompressorDecodePlan, CompressorPrefillPlan
from sglang.srt.utils import get_device
@dataclass
@@ -46,7 +47,7 @@ class LegacyContext:
seq_lens=seq_lens_cpu,
extend_lens=extend_lens_cpu,
num_q_tokens=num_q_tokens,
device=torch.device("cuda"),
device=torch.device(get_device()),
)
def make_decode_plan(self, seq_lens_gpu: torch.Tensor) -> CompressorDecodePlan:
@@ -127,7 +128,7 @@ def make_legacy_context(
head_dim: int = 512,
) -> LegacyContext:
pages_per_req = 2 if compress_ratio == 4 else 1
req_pool_indices = torch.arange(bs, dtype=torch.int64, device="cuda")
req_pool_indices = torch.arange(bs, dtype=torch.int64, device=get_device())
return LegacyContext(
bs=bs,
head_dim=head_dim,
@@ -171,9 +172,9 @@ def make_paged_context(
swa_page_size=swa_page_size,
ring_size=ring_size,
num_swa_pages_per_req=num_swa_pages_per_req,
req_pool_indices=req_pool_indices.cuda(),
req_to_token=req_to_token.cuda(),
full_to_swa=full_to_swa.cuda(),
req_pool_indices=req_pool_indices.to(get_device()),
req_to_token=req_to_token.to(get_device()),
full_to_swa=full_to_swa.to(get_device()),
)
@@ -182,7 +183,7 @@ def make_state_pool(num_pages: int, compress_ratio: int, head_dim: int) -> torch
return torch.zeros(
(num_pages, compress_ratio, last_dim),
dtype=torch.float32,
device="cuda",
device=get_device(),
)