Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
22 changes: 19 additions & 3 deletions python/sglang/srt/layers/rotary_embedding/base.py
Original file line number Diff line number Diff line change
Expand Up @@ -69,7 +69,19 @@
)

if _is_xpu:
from sgl_kernel import fused_qk_rope_with_cos_sin_cache_inplace
try:
from sgl_kernel import fused_qk_rope_with_cos_sin_cache_inplace
except ImportError:
# Older/CUDA sgl_kernel builds don't ship this XPU fused kernel. Don't
# let a missing symbol break module import (it would deregister every
# model that imports get_rope and silently fall back to the generic
# Transformers backbone). forward_xpu() handles the None case below.
fused_qk_rope_with_cos_sin_cache_inplace = None
logger.warning(
"sgl_kernel.fused_qk_rope_with_cos_sin_cache_inplace is unavailable; "
"XPU rotary embedding will use the generic rotary_embedding kernel. "
"Upgrade sgl_kernel to enable the fused XPU kernel."
)


class RotaryEmbedding(MultiPlatformOp):
Expand Down Expand Up @@ -453,8 +465,12 @@ def forward_xpu(

self._match_cos_sin_cache_dtype(query)

# Fused_qk_rope only supports aligned head_size
if self.head_size in [128, 256, 512]:
# Fused_qk_rope only supports aligned head_size, and requires the fused
# XPU kernel to be available in sgl_kernel.
if (
fused_qk_rope_with_cos_sin_cache_inplace is not None
and self.head_size in [128, 256, 512]
):
num_tokens = positions.size(0)
q_rope = query.view(num_tokens, -1, self.head_size)
k_rope = key.view(num_tokens, -1, self.head_size)
Expand Down
15 changes: 13 additions & 2 deletions python/sglang/srt/layers/rotary_embedding/mrope.py
Original file line number Diff line number Diff line change
Expand Up @@ -40,7 +40,14 @@
import torch_npu

if _is_xpu:
from sgl_kernel import multimodal_rotary_embedding
try:
from sgl_kernel import multimodal_rotary_embedding
except ImportError:
# Older/CUDA sgl_kernel builds don't ship this XPU kernel. Avoid a hard
# import failure (it would deregister every model importing get_rope and
# silently fall back to the generic Transformers backbone). forward_xpu()
# falls back to forward_native when the kernel is None.
multimodal_rotary_embedding = None

import triton
import triton.language as tl
Expand Down Expand Up @@ -386,7 +393,11 @@ def forward_xpu(
fused_set_kv_buffer_arg=None,
) -> Tuple[torch.Tensor, torch.Tensor]:
assert positions.ndim in (1, 2)
if positions.ndim == 2 and self.mrope_section:
if (
multimodal_rotary_embedding is not None
and positions.ndim == 2
and self.mrope_section
):
multimodal_rotary_embedding(
query,
key,
Expand Down
2 changes: 1 addition & 1 deletion python/sglang/srt/lora/backend/chunked_backend.py
Original file line number Diff line number Diff line change
Expand Up @@ -224,7 +224,7 @@ def init_cuda_graph_batch_info(
(num_tokens_per_bs + MIN_CHUNK_SIZE - 1) // MIN_CHUNK_SIZE
) * max_bs_in_cuda_graph
max_num_tokens = max_bs_in_cuda_graph * num_tokens_per_bs
with torch.device("cuda"):
with torch.device(self.device):
self.cuda_graph_batch_info = LoRABatchInfo(
bs=max_bs_in_cuda_graph,
use_cuda_graph=True,
Expand Down
2 changes: 1 addition & 1 deletion python/sglang/srt/lora/backend/torch_backend.py
Original file line number Diff line number Diff line change
Expand Up @@ -166,7 +166,7 @@ def init_cuda_graph_batch_info(
max_bs_in_cuda_graph: int,
num_tokens_per_bs: int,
):
with torch.device("cuda"):
with torch.device(self.device):
self.cuda_graph_batch_info = TorchNativeLoRABatchInfo(
use_cuda_graph=True,
bs=max_bs_in_cuda_graph,
Expand Down
2 changes: 1 addition & 1 deletion python/sglang/srt/lora/backend/triton_backend.py
Original file line number Diff line number Diff line change
Expand Up @@ -146,7 +146,7 @@ def init_cuda_graph_batch_info(
):
max_tokens = max_bs_in_cuda_graph * num_tokens_per_bs
mlpb = self.max_loras_per_batch
with torch.device("cuda"):
with torch.device(self.device):
self.cuda_graph_batch_info = LoRABatchInfo(
bs=max_bs_in_cuda_graph,
use_cuda_graph=True,
Expand Down
2 changes: 1 addition & 1 deletion python/sglang/srt/lora/lora_moe_runners.py
Original file line number Diff line number Diff line change
Expand Up @@ -227,7 +227,7 @@ def _compute_lora_alignment(

device = topk_ids.device

use_naive = (
use_naive = _is_xpu or (
cg is None
and M * topk_ids.shape[1] * _SPARSITY_FACTOR
<= lora_info.num_experts * max_loras
Expand Down
2 changes: 1 addition & 1 deletion python/sglang/srt/lora/lora_overlap_loader.py
Original file line number Diff line number Diff line change
Expand Up @@ -72,7 +72,7 @@ def _drain_completed_overlap_loads(self) -> None:
if event.query()
]
for lora_id, event in completed_loads:
torch.cuda.current_stream().wait_event(event)
self.device_module.current_stream().wait_event(event)
del self.lora_to_overlap_load_event[lora_id]

def _try_start_overlap_load(
Expand Down
50 changes: 40 additions & 10 deletions python/sglang/test/lora_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,9 +4,34 @@

import torch

from sglang.srt.utils import is_xpu
from sglang.test.runners import HFRunner, SRTRunner
from sglang.test.test_utils import calculate_rouge_l

_IS_XPU = is_xpu()


def _assert_lora_output_match(srt_str: str, hf_str: str, rouge_tol: float, context: str):
"""Compare SRT vs HF greedy output strings.

Everywhere except XPU we keep the historical strict exact-match (SGLang and HF
kernels agree numerically enough for greedy argmax to pick identical tokens).
On XPU, small kernel-level fp differences can make greedy decoding diverge
after a shared prefix even when the LoRA math is correct, so we fall back to
the same ROUGE-L tolerance the per-adaptor comparison path uses.
"""
srt_str = srt_str.strip(" ")
hf_str = hf_str.strip(" ")
if not _IS_XPU:
assert srt_str == hf_str, (srt_str, hf_str)
return
rouge_score = calculate_rouge_l([srt_str], [hf_str])[0]
if rouge_score < rouge_tol:
raise AssertionError(
f"ROUGE-L score {rouge_score} below tolerance {rouge_tol} for {context}. "
f"SRT: {srt_str!r} HF: {hf_str!r}"
)


@dataclasses.dataclass
class LoRAAdaptor:
Expand Down Expand Up @@ -625,17 +650,22 @@ def run_lora_test_by_batch(
print("HF output:", hf_output_str)
print("SRT no lora output:", srt_no_lora_outputs.output_strs[i].strip())
print("HF no lora output:", hf_no_lora_outputs.output_strs[i].strip())
assert srt_outputs.output_strs[i].strip(" ") == hf_outputs.output_strs[i].strip(
" "
), (
srt_outputs.output_strs[i].strip(" "),
hf_outputs.output_strs[i].strip(" "),
rouge_tol = (
adaptors[i].rouge_l_tolerance
if adaptors[i].rouge_l_tolerance is not None
else model_case.rouge_l_tolerance
)
_assert_lora_output_match(
srt_outputs.output_strs[i],
hf_outputs.output_strs[i],
rouge_tol,
f"base '{base_path}', adaptor '{adaptor_names[i]}', backend '{backend}' (LoRA)",
)
assert srt_no_lora_outputs.output_strs[i].strip(
" "
) == hf_no_lora_outputs.output_strs[i].strip(" "), (
srt_no_lora_outputs.output_strs[i].strip(" "),
hf_no_lora_outputs.output_strs[i].strip(" "),
_assert_lora_output_match(
srt_no_lora_outputs.output_strs[i],
hf_no_lora_outputs.output_strs[i],
rouge_tol,
f"base '{base_path}', backend '{backend}' (no-LoRA baseline)",
)


Expand Down
25 changes: 21 additions & 4 deletions test/registered/lora/test_chunked_sgmv_backend.py
Original file line number Diff line number Diff line change
Expand Up @@ -20,7 +20,8 @@
from sglang.srt.lora.triton_ops.chunked_sgmv_shrink import _chunked_lora_shrink_kernel
from sglang.srt.lora.utils import LoRABatchInfo, get_lm_head_pruned_lens
from sglang.srt.model_executor.forward_batch_info import ForwardMode
from sglang.test.ci.ci_register import register_cuda_ci
from sglang.srt.utils import get_device
from sglang.test.ci.ci_register import register_cuda_ci, register_xpu_ci
from sglang.test.lora_utils import (
reference_embedding_lora_a_shrink,
reference_sgmv_expand,
Expand All @@ -30,11 +31,27 @@
CHUNK_SIZE = 16

register_cuda_ci(est_time=60, suite="nightly-1-gpu", nightly=True)
register_xpu_ci(est_time=60, suite="stage-a-test-1-gpu-xpu")


def _clear_jit_cache(kernel):
"""Clear a triton JITFunction's compiled-kernel cache across triton versions.

Older triton exposes ``_clear_cache()``; newer builds (e.g. Intel triton on
XPU) keep per-device caches in ``device_caches`` instead.
"""
clear = getattr(kernel, "_clear_cache", None)
if callable(clear):
clear()
return
device_caches = getattr(kernel, "device_caches", None)
if isinstance(device_caches, dict):
device_caches.clear()


def reset_kernel_cache():
_chunked_lora_shrink_kernel._clear_cache()
_chunked_lora_expand_kernel._clear_cache()
_clear_jit_cache(_chunked_lora_shrink_kernel)
_clear_jit_cache(_chunked_lora_expand_kernel)


class BatchComposition(Enum):
Expand Down Expand Up @@ -110,7 +127,7 @@ def setUp(self):
torch.manual_seed(42)
random.seed(42)

self.device = torch.device("cuda")
self.device = torch.device(get_device())
self.dtype = torch.float16
self.input_dim = 2560 # Hidden dimension
self.max_seq_len = 1024
Expand Down
59 changes: 39 additions & 20 deletions test/registered/lora/test_fused_moe_lora_kernel.py
Original file line number Diff line number Diff line change
Expand Up @@ -9,13 +9,15 @@
# IMPORT PREBUILT KERNEL
# ==============================================================================
from sglang.jit_kernel.moe_lora_align import moe_lora_align_block_size
from sglang.srt.lora.lora_moe_runners import _naive_moe_lora_align_block_size
from sglang.srt.lora.triton_ops import fused_moe_lora
from sglang.srt.utils import set_random_seed
from sglang.test.ci.ci_register import register_cuda_ci
from sglang.srt.utils import get_device, is_xpu, set_random_seed
from sglang.test.ci.ci_register import register_cuda_ci, register_xpu_ci

# ==============================================================================

register_cuda_ci(est_time=28, stage="base-b", runner_config="1-gpu-large")
register_xpu_ci(est_time=28, suite="stage-a-test-1-gpu-xpu")


def round_up(x, base):
Expand Down Expand Up @@ -159,23 +161,40 @@ def use_fused_moe_lora_kernel(
adapter_enabled = torch.ones(max_loras + 1, dtype=torch.int32, device=device)
lora_ids = torch.arange(max_loras, dtype=torch.int32, device=device)

# call kernel
moe_lora_align_block_size(
topk_ids,
seg_indptr,
req_to_lora,
num_experts,
block_size,
max_loras,
max_num_tokens_padded,
max_num_m_blocks,
sorted_token_ids,
expert_ids,
num_tokens_post_padded,
adapter_enabled,
lora_ids,
None, # maybe_expert_map
)
# call kernel — the fused CUDA align kernel exists only for CUDA; on XPU
# use the pure-torch native alignment.
if is_xpu():
sorted_token_ids, expert_ids, num_tokens_post_padded = (
_naive_moe_lora_align_block_size(
topk_ids,
seg_indptr,
req_to_lora,
num_experts,
block_size,
max_loras,
max_num_tokens_padded,
max_num_m_blocks,
adapter_enabled,
device,
)
)
else:
moe_lora_align_block_size(
topk_ids,
seg_indptr,
req_to_lora,
num_experts,
block_size,
max_loras,
max_num_tokens_padded,
max_num_m_blocks,
sorted_token_ids,
expert_ids,
num_tokens_post_padded,
adapter_enabled,
lora_ids,
None, # maybe_expert_map
)

config = {
"BLOCK_SIZE_M": 16,
Expand Down Expand Up @@ -263,7 +282,7 @@ def use_torch(


DTYPES = [torch.float32, torch.float16, torch.bfloat16]
DEVICES = [f"cuda:{0}"]
DEVICES = [get_device(0)]
SEED = [42]


Expand Down
7 changes: 6 additions & 1 deletion test/registered/lora/test_lora_hf_sgl_logprob_diff.py
Original file line number Diff line number Diff line change
Expand Up @@ -35,7 +35,11 @@
import numpy as np
import torch

from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci
from sglang.test.ci.ci_register import (
register_amd_ci,
register_cuda_ci,
register_xpu_ci,
)
from sglang.test.runners import HFRunner, SRTRunner
from sglang.test.test_utils import DEFAULT_PORT_FOR_SRT_TEST_RUNNER, CustomTestCase

Expand All @@ -48,6 +52,7 @@
est_time=250,
suite="stage-b-test-1-gpu-small-amd",
)
register_xpu_ci(est_time=250, suite="stage-a-test-1-gpu-xpu")
# Test configuration constants
BASE_MODEL = "meta-llama/Llama-2-7b-hf"
LORA_PATHS = ["yushengsu/sglang_lora_logprob_diff_without_tuning"]
Expand Down
16 changes: 14 additions & 2 deletions test/registered/lora/test_lora_moe_vllm_sgl_logprob_diff.py
Original file line number Diff line number Diff line change
Expand Up @@ -16,7 +16,7 @@

import torch

from sglang.test.ci.ci_register import register_cuda_ci
from sglang.test.ci.ci_register import register_cuda_ci, register_xpu_ci
from sglang.test.lora_utils import (
MOE_BASE_MODEL_PATH,
MOE_LORA_PATH,
Expand All @@ -29,6 +29,7 @@
stage="base-b",
runner_config="1-gpu-large",
)
register_xpu_ci(est_time=60, suite="stage-a-test-1-gpu-xpu")

# Format: [{"text": "result string", "lps": [0.1, 0.2, ...]}, ...]
VLLM_CACHED_RESULTS = [
Expand Down Expand Up @@ -287,6 +288,17 @@ class TestMoELoraRegression(unittest.TestCase):

def test_sglang_moe_parity_strict(self):

# flashinfer is CUDA-only; on other devices (e.g. XPU) use the
# platform-appropriate attention backend.
from sglang.srt.utils import is_cuda, is_xpu

if is_cuda():
attention_backend = "flashinfer"
elif is_xpu():
attention_backend = "intel_xpu"
else:
attention_backend = "triton"

with SRTRunner(
model_path=MOE_BASE_MODEL_PATH,
torch_dtype=torch.bfloat16,
Expand All @@ -296,7 +308,7 @@ def test_sglang_moe_parity_strict(self):
tp_size=1,
trust_remote_code=True,
disable_radix_cache=True,
attention_backend="flashinfer",
attention_backend=attention_backend,
mem_fraction_static=0.80,
) as srt_runner:

Expand Down
Loading
Loading