From 6d79c609954585c5e40d5f2b24dc5eb30d1fe41a Mon Sep 17 00:00:00 2001 From: "Ethan (Yusheng) Su" Date: Thu, 9 Apr 2026 22:31:53 -0700 Subject: [PATCH] [Lora] Lora kimi support (#22381) --- .../compressed_tensors/compressed_tensors.py | 13 ++ .../schemes/compressed_tensors_wNa16_moe.py | 23 +-- python/sglang/srt/lora/layers.py | 9 +- python/sglang/srt/lora/lora_manager.py | 5 +- .../lora/test_lora_kimi_k25_logprob_diff.py | 150 ++++++++++++++++++ 5 files changed, 188 insertions(+), 12 deletions(-) create mode 100644 test/registered/lora/test_lora_kimi_k25_logprob_diff.py diff --git a/python/sglang/srt/layers/quantization/compressed_tensors/compressed_tensors.py b/python/sglang/srt/layers/quantization/compressed_tensors/compressed_tensors.py index a13c53af4..5734b9dc0 100644 --- a/python/sglang/srt/layers/quantization/compressed_tensors/compressed_tensors.py +++ b/python/sglang/srt/layers/quantization/compressed_tensors/compressed_tensors.py @@ -682,6 +682,16 @@ class CompressedTensorsConfig(QuantizationConfig): logger.info_once("Using CompressedTensorsWNA16TritonMoE (ROCm)") return CompressedTensorsWNA16TritonMoE(self) else: + from sglang.srt.server_args import get_global_server_args + + server_args = get_global_server_args() + if server_args and server_args.enable_lora: + logger.info_once( + "Using CompressedTensorsWNA16TritonMoEMethod " + "(LoRA requires triton-compatible MoE weights)" + ) + return CompressedTensorsWNA16TritonMoE(self) + logger.info_once("Using CompressedTensorsWNA16MarlinMoEMethod") return CompressedTensorsWNA16MoE(self) else: @@ -997,6 +1007,9 @@ class CompressedTensorsFusedMoEMethod(FusedMoEMethodBase): ): return layer.scheme.create_moe_runner(layer, moe_runner_config) + def get_triton_quant_info(self, layer: torch.nn.Module): + return layer.scheme.get_triton_quant_info(layer) + def apply( self, layer: torch.nn.Module, diff --git a/python/sglang/srt/layers/quantization/compressed_tensors/schemes/compressed_tensors_wNa16_moe.py b/python/sglang/srt/layers/quantization/compressed_tensors/schemes/compressed_tensors_wNa16_moe.py index 7a8fb6542..072fedfae 100644 --- a/python/sglang/srt/layers/quantization/compressed_tensors/schemes/compressed_tensors_wNa16_moe.py +++ b/python/sglang/srt/layers/quantization/compressed_tensors/schemes/compressed_tensors_wNa16_moe.py @@ -448,18 +448,10 @@ class CompressedTensorsWNA16TritonMoE(CompressedTensorsWNA16MoE): self.moe_runner_config = moe_runner_config self.runner = MoeRunner(MoeRunnerBackend.TRITON, moe_runner_config) - def apply_weights( - self, - layer: torch.nn.Module, - dispatch_output: "StandardDispatchOutput", - ) -> "CombineInput": + def get_triton_quant_info(self, layer): from sglang.srt.layers.moe.moe_runner.triton import TritonMoeQuantInfo - assert ( - self.moe_runner_config.activation == "silu" - ), "Only SiLU activation is supported." - - quant_info = TritonMoeQuantInfo( + return TritonMoeQuantInfo( w13_weight=layer.w13_weight_packed, w2_weight=layer.w2_weight_packed, use_int4_w4a16=True, @@ -467,6 +459,17 @@ class CompressedTensorsWNA16TritonMoE(CompressedTensorsWNA16MoE): w2_scale=layer.w2_weight_scale, block_shape=[0, self.group_size], ) + + def apply_weights( + self, + layer: torch.nn.Module, + dispatch_output: "StandardDispatchOutput", + ) -> "CombineInput": + assert ( + self.moe_runner_config.activation == "silu" + ), "Only SiLU activation is supported." + + quant_info = self.get_triton_quant_info(layer) return self.runner.run(dispatch_output, quant_info) diff --git a/python/sglang/srt/lora/layers.py b/python/sglang/srt/lora/layers.py index 4a327f1e4..4128d0fda 100644 --- a/python/sglang/srt/lora/layers.py +++ b/python/sglang/srt/lora/layers.py @@ -809,10 +809,17 @@ class FusedMoEWithLoRA(BaseLayerWithLoRA): ) # initialize triton_lora moe runner for batches with lora enabled + from sglang.srt.layers.moe import MoeRunnerBackend from sglang.srt.layers.moe.moe_runner.runner import MoeRunner + qm = base_layer.quant_method + if hasattr(qm, "runner") and qm.runner is not None: + runner_backend = qm.runner.runner_backend + else: + runner_backend = MoeRunnerBackend.TRITON + self._lora_runner = MoeRunner( - base_layer.quant_method.runner.runner_backend, + runner_backend, base_layer.moe_runner_config, lora_enabled=True, ) diff --git a/python/sglang/srt/lora/lora_manager.py b/python/sglang/srt/lora/lora_manager.py index b14de3f4b..06f3b09dd 100644 --- a/python/sglang/srt/lora/lora_manager.py +++ b/python/sglang/srt/lora/lora_manager.py @@ -66,7 +66,10 @@ class LoRAManager: lora_paths: Optional[List[LoRARef]] = None, ): self.base_model: torch.nn.Module = base_model - self.base_hf_config: AutoConfig = base_hf_config + if hasattr(base_hf_config, "get_text_config"): + self.base_hf_config: AutoConfig = base_hf_config.get_text_config() + else: + self.base_hf_config: AutoConfig = base_hf_config self.max_loras_per_batch: int = max_loras_per_batch self.load_config: LoadConfig = load_config self.dtype: torch.dtype = dtype diff --git a/test/registered/lora/test_lora_kimi_k25_logprob_diff.py b/test/registered/lora/test_lora_kimi_k25_logprob_diff.py new file mode 100644 index 000000000..b64b8e4b0 --- /dev/null +++ b/test/registered/lora/test_lora_kimi_k25_logprob_diff.py @@ -0,0 +1,150 @@ +# Copyright 2023-2025 SGLang Team +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# ============================================================================== + +""" +Regression test for Kimi-K2.5 (VLM + MLA + MoE) LoRA logprob accuracy. + +Compares SGLang LoRA logprobs against reference training logprobs from a +pre-computed dataset. The LoRA adapter and reference data are downloaded from: +https://huggingface.co/datasets/yushengsu/lora-diff-Kimi-K2.5 + +Usage: + python -m unittest test_lora_kimi_k25_logprob_diff +""" + +import multiprocessing as mp +import os +import unittest + +import torch +from huggingface_hub import snapshot_download + +import sglang as sgl +from sglang.test.ci.ci_register import register_cuda_ci +from sglang.test.test_utils import CustomTestCase + +register_cuda_ci( + est_time=360, + suite="nightly-8-gpu-b200", +) + +BASE_MODEL = "moonshotai/Kimi-K2.5" +LORA_HF_REPO = "yushengsu/lora-diff-Kimi-K2.5" +LORA_BACKEND = "triton" +MAX_LORA_RANK = 32 +TP_SIZE = 8 +MOE_RUNNER_BACKEND = "triton" +EXPERTS_SHARED_OUTER_LORAS = True +PREFILL_ATTENTION_BACKEND = "fa4" +DECODE_ATTENTION_BACKEND = "flashinfer" + +KL_THRESHOLD = 1.5e-2 + + +def kl_v2(a, b): + a = torch.tensor(a) if not torch.is_tensor(a) else a + b = torch.tensor(b) if not torch.is_tensor(b) else b + return (((a - b) ** 2) * 0.5).mean().item() + + +def get_prompt_logprobs(engine, input_ids, lora_path): + out = engine.generate( + input_ids=input_ids, + sampling_params={"max_new_tokens": 0, "temperature": 0.0}, + return_logprob=True, + logprob_start_len=0, + lora_path=lora_path, + ) + return [logprob for logprob, _, _ in out["meta_info"]["input_token_logprobs"]][1:] + + +class TestLoRAKimiK25LogprobDiff(CustomTestCase): + + def test_lora_kimi_k25_logprob_accuracy(self): + adapter_path = snapshot_download( + LORA_HF_REPO, + repo_type="dataset", + ) + + engine = sgl.Engine( + model_path=BASE_MODEL, + tp_size=TP_SIZE, + enable_lora=True, + max_lora_rank=MAX_LORA_RANK, + lora_paths={"my_lora": adapter_path}, + lora_backend=LORA_BACKEND, + attention_backend="flashinfer", + moe_runner_backend=MOE_RUNNER_BACKEND, + experts_shared_outer_loras=EXPERTS_SHARED_OUTER_LORAS, + prefill_attention_backend=PREFILL_ATTENTION_BACKEND, + decode_attention_backend=DECODE_ATTENTION_BACKEND, + trust_remote_code=True, + ) + + try: + cdata = torch.load( + os.path.join(adapter_path, "compare_sample_train_data.pt"), + weights_only=False, + ) + + base_logprobs = get_prompt_logprobs(engine, cdata["tokens"], lora_path=None) + logprobs = get_prompt_logprobs(engine, cdata["tokens"], lora_path="my_lora") + + base_t = torch.tensor(base_logprobs) + lora_t = torch.tensor(logprobs) + diff = (base_t - lora_t).abs() + print( + f"[VERIFY] base vs lora: mean_diff={diff.mean().item():.6f}, " + f"max_diff={diff.max().item():.6f}, " + f"identical={torch.equal(base_t, lora_t)}" + ) + + self.assertFalse( + torch.equal(base_t, lora_t), + "LoRA logprobs should differ from base model logprobs", + ) + + kl_sglang_trainer = kl_v2(cdata["training_logprobs"], logprobs) + kl_orig_trainer = kl_v2( + cdata["training_logprobs"], cdata["sampling_logprobs"] + ) + kl_sglang_orig = kl_v2(logprobs, cdata["sampling_logprobs"]) + + print(f"KL(orig_sampler, trainer) = {kl_orig_trainer:.6e}") + print(f"KL(sglang, trainer) = {kl_sglang_trainer:.6e}") + print(f"KL(sglang, orig_sampler) = {kl_sglang_orig:.6e}") + + self.assertLessEqual( + kl_sglang_trainer, + KL_THRESHOLD, + f"KL(sglang, trainer) = {kl_sglang_trainer:.6e} exceeds " + f"threshold {KL_THRESHOLD}", + ) + + finally: + engine.shutdown() + + +if __name__ == "__main__": + try: + mp.set_start_method("spawn") + except RuntimeError: + pass + + try: + unittest.main(warnings="ignore", verbosity=2) + finally: + if torch.cuda.is_available(): + torch.cuda.empty_cache() + torch.cuda.synchronize()