From 7d893255c359bb8ab74d2870c8ac865fb57230d6 Mon Sep 17 00:00:00 2001 From: Yanbin Jiang Date: Fri, 21 Aug 2026 14:21:53 -0700 Subject: [PATCH] [Spec][LoRA] Support multi-adapter LoRA with EAGLE/NEXTN/DFLASH/DSPARK speculative decoding (#34337) --- .../sglang/srt/lora/backend/base_backend.py | 17 +- .../srt/lora/backend/chunked_backend.py | 7 +- .../sglang/srt/lora/backend/triton_backend.py | 25 +- python/sglang/srt/lora/layers.py | 5 + python/sglang/srt/lora/lora_manager.py | 25 +- python/sglang/srt/lora/utils.py | 37 + .../srt/model_executor/forward_batch_info.py | 6 +- .../sglang/srt/model_executor/model_runner.py | 6 +- .../cuda_graph_setup.py | 2 +- .../srt/model_executor/runner/base_runner.py | 2 +- .../runner/decode_cuda_graph_runner.py | 6 +- .../runner/prefill_cuda_graph_runner.py | 2 +- python/sglang/srt/server_args.py | 68 +- .../srt/speculative/dflash_worker_v2.py | 7 +- .../dspark_components/dspark_draft.py | 3 +- .../dspark_components/dspark_worker_v2.py | 7 +- .../sglang/srt/speculative/eagle_worker_v2.py | 6 +- .../lora/check_spec_baseline_divergence.py | 125 +++ test/manual/lora/run_spec_lora_matrix.py | 876 ++++++++++++++++++ .../lora/test_lora_spec_decoding.py | 142 +++ .../lora/test_lora_spec_verify_batch_info.py | 109 +++ .../test_draft_runner_skips_lora.py | 51 + .../test_prefill_cuda_graph_runner.py | 4 +- 23 files changed, 1487 insertions(+), 51 deletions(-) create mode 100644 test/manual/lora/check_spec_baseline_divergence.py create mode 100644 test/manual/lora/run_spec_lora_matrix.py create mode 100644 test/registered/lora/test_lora_spec_decoding.py create mode 100644 test/registered/unit/lora/test_lora_spec_verify_batch_info.py create mode 100644 test/registered/unit/model_executor/test_draft_runner_skips_lora.py diff --git a/python/sglang/srt/lora/backend/base_backend.py b/python/sglang/srt/lora/backend/base_backend.py index 46d5e47cb..5f70a02f5 100644 --- a/python/sglang/srt/lora/backend/base_backend.py +++ b/python/sglang/srt/lora/backend/base_backend.py @@ -5,7 +5,11 @@ import triton import triton.language as tl from sglang.srt.lora.backend.lmhead_mixing import LoRABackendLmHeadMixing -from sglang.srt.lora.utils import LoRABatchInfo, MoELoRABatchInfo +from sglang.srt.lora.utils import ( + LoRABatchInfo, + MoELoRABatchInfo, + get_batch_token_counts, +) from sglang.srt.model_executor.forward_batch_info import ForwardBatch @@ -296,16 +300,7 @@ class BaseLoRABackend(LoRABackendLmHeadMixing): adapter_enabled = None token_lora_mapping = None - num_tokens = ( - sum(forward_batch.extend_seq_lens_cpu) - if forward_batch.forward_mode.is_extend() - else forward_batch.batch_size - ) - max_len = ( - max(forward_batch.extend_seq_lens_cpu) - if forward_batch.forward_mode.is_extend() - else 1 - ) + num_tokens, max_len = get_batch_token_counts(forward_batch) if ( batch_info.req_seg_indptr is not None diff --git a/python/sglang/srt/lora/backend/chunked_backend.py b/python/sglang/srt/lora/backend/chunked_backend.py index 8e525cfc1..92a77b8d7 100644 --- a/python/sglang/srt/lora/backend/chunked_backend.py +++ b/python/sglang/srt/lora/backend/chunked_backend.py @@ -12,6 +12,7 @@ from sglang.srt.lora.backend.base_backend import BaseLoRABackend from sglang.srt.lora.utils import ( LoRABatchInfo, generate_sequence_lengths, + get_batch_token_counts, get_lm_head_pruned_lens, merge_and_chunk_segments, ) @@ -184,11 +185,7 @@ class ChunkedSgmvLoRABackend(BaseLoRABackend): Returns: The determined chunk size """ - num_tokens = ( - forward_batch.extend_num_tokens - if forward_batch.forward_mode.is_extend() - else forward_batch.batch_size - ) + num_tokens, _ = get_batch_token_counts(forward_batch) return self._determine_chunk_size_for_tokens(num_tokens) def _determine_chunk_size_for_tokens(self, num_tokens: int) -> int: diff --git a/python/sglang/srt/lora/backend/triton_backend.py b/python/sglang/srt/lora/backend/triton_backend.py index fbf8edb00..f8f0e6d12 100644 --- a/python/sglang/srt/lora/backend/triton_backend.py +++ b/python/sglang/srt/lora/backend/triton_backend.py @@ -11,6 +11,8 @@ from sglang.kernels.ops.gemm.sgemm_lora_b import sgemm_lora_b_fwd from sglang.srt.lora.backend.base_backend import BaseLoRABackend from sglang.srt.lora.utils import ( LoRABatchInfo, + generate_sequence_lengths, + get_batch_token_counts, get_lm_head_pruned_lens, merge_and_chunk_segments, ) @@ -289,6 +291,15 @@ class TritonLoRABackend(BaseLoRABackend): self.cuda_graph_batch_info is not None ), "CUDA Graph batch info is not initialized." batch_info = self.cuda_graph_batch_info + if forward_batch.forward_mode.is_target_verify(): + # seg_lens were pre-filled at the captured per-request width + # (stored as max_len); another width would silently + # mis-segment adapters onto the wrong token rows. + assert forward_batch.spec_info.draft_token_num == batch_info.max_len, ( + "target-verify width " + f"{forward_batch.spec_info.draft_token_num} does not match " + f"the captured LoRA cuda-graph width {batch_info.max_len}" + ) batch_info.bs = forward_batch.batch_size batch_info.num_segments = forward_batch.batch_size elif use_prefill_cuda_graph: @@ -303,17 +314,9 @@ class TritonLoRABackend(BaseLoRABackend): batch_info.seg_lens[bs:].zero_() torch.cumsum(batch_info.seg_lens, dim=0, out=batch_info.seg_indptr[1:]) else: - max_len = ( - # Calculate max_len from the CPU copy to avoid D2H transfer. - max(forward_batch.extend_seq_lens_cpu) - if forward_batch.forward_mode.is_extend() - else 1 - ) - seg_lens = ( - forward_batch.extend_seq_lens - if forward_batch.forward_mode.is_extend() - else torch.ones(bs, dtype=torch.int32, device=self.device) - ) + # max_len comes from the CPU-side counts to avoid a D2H transfer. + _, max_len = get_batch_token_counts(forward_batch) + seg_lens = generate_sequence_lengths(forward_batch, device=self.device) seg_indptr = torch.zeros((bs + 1,), dtype=torch.int32, device=self.device) seg_indptr[1:] = torch.cumsum(seg_lens, dim=0) diff --git a/python/sglang/srt/lora/layers.py b/python/sglang/srt/lora/layers.py index 87b2a3342..f92342f34 100644 --- a/python/sglang/srt/lora/layers.py +++ b/python/sglang/srt/lora/layers.py @@ -31,6 +31,11 @@ from sglang.srt.runtime_context import get_parallel _SGLANG_EXPERIMENTAL_LORA_OPTI = envs.SGLANG_EXPERIMENTAL_LORA_OPTI.get() +def unwrap_lora_layer(module: nn.Module) -> nn.Module: + """Return the plain module behind a LoRA wrapper, or the module itself.""" + return module.base_layer if isinstance(module, BaseLayerWithLoRA) else module + + class BaseLayerWithLoRA(nn.Module): def __init__( self, diff --git a/python/sglang/srt/lora/lora_manager.py b/python/sglang/srt/lora/lora_manager.py index aae2afffc..d19f3a6c7 100644 --- a/python/sglang/srt/lora/lora_manager.py +++ b/python/sglang/srt/lora/lora_manager.py @@ -43,6 +43,7 @@ from sglang.srt.lora.utils import ( auto_detect_lora_target_modules, get_normalized_target_modules, get_target_module_name, + warn_if_adapter_targets_embeddings, ) from sglang.srt.managers.io_struct import LoRAUpdateOutput from sglang.srt.model_executor.forward_batch_info import ForwardBatch @@ -50,6 +51,7 @@ from sglang.srt.runtime_context import ( get_exec, get_lora, get_parallel, + get_spec, ) from sglang.srt.server_args import ServerArgs from sglang.srt.utils import get_available_gpu_memory, replace_submodule @@ -105,6 +107,7 @@ class LoRAManager: self.lora_strict_loading: bool = getattr( server_args, "lora_strict_loading", False ) + self.speculative_algorithm: Optional[str] = get_spec().speculative_algorithm # LoRA backend for running sgemm kernels logger.info(f"Using {lora_backend} as backend of LoRA kernels.") @@ -782,6 +785,12 @@ class LoRAManager: ) lora_adapter.initialize_weights() + warn_if_adapter_targets_embeddings( + lora_name=lora_ref.lora_name, + embedding_layer_names=lora_adapter.embedding_layers.keys(), + speculative_algorithm=self.speculative_algorithm, + ) + self.loras[lora_ref.lora_id] = lora_adapter def load_lora_weights_from_tensors( @@ -799,6 +808,12 @@ class LoRAManager: base_model=self.base_model, ) lora_adapter.initialize_weights_from_tensors(tensors) + + warn_if_adapter_targets_embeddings( + lora_name=lora_ref.lora_name, + embedding_layer_names=lora_adapter.embedding_layers.keys(), + speculative_algorithm=self.speculative_algorithm, + ) self.loras[lora_ref.lora_id] = lora_adapter def load_lora_adapter_from_tensors( @@ -1035,12 +1050,18 @@ def init_lora_cuda_graph_moe_buffers( from sglang.srt.lora.layers import FusedMoEWithLoRA max_bs = get_exec().graph.cuda_graph_config.decode.max_bs + # With spec on, the decode graph captures TARGET_VERIFY batches of + # num_draft_tokens per request, and the buffers below are per-token, so + # they must be sized in tokens rather than requests. + max_tokens = max_bs * (get_spec().speculative_num_draft_tokens or 1) max_loras = get_lora().max_loras_per_batch for module in model.modules(): if isinstance(module, FusedMoEWithLoRA): - lora_manager.init_cuda_graph_moe_buffers(max_bs, max_loras, dtype, module) + lora_manager.init_cuda_graph_moe_buffers( + max_tokens, max_loras, dtype, module + ) logger.info( f"Pre-allocated shared MoE LoRA CUDA graph buffers " - f"(max_bs={max_bs}, max_loras={max_loras})" + f"(max_tokens={max_tokens}, max_loras={max_loras})" ) break diff --git a/python/sglang/srt/lora/utils.py b/python/sglang/srt/lora/utils.py index 213e40a1a..ce764275a 100644 --- a/python/sglang/srt/lora/utils.py +++ b/python/sglang/srt/lora/utils.py @@ -1,3 +1,4 @@ +import logging from dataclasses import dataclass from enum import Enum from typing import Iterable, List, Optional, Set, Tuple, Union @@ -7,6 +8,29 @@ import torch from sglang.srt.model_executor.forward_batch_info import ForwardBatch from sglang.srt.utils.hf_transformers_utils import AutoConfig +logger = logging.getLogger(__name__) + + +def warn_if_adapter_targets_embeddings( + lora_name: str, + embedding_layer_names: Iterable[str], + speculative_algorithm: Optional[str], +) -> None: + """Warn once when an adapter carries embedding weights under EAGLE spec.""" + if speculative_algorithm not in ("EAGLE", "EAGLE3"): + return + modules = sorted(embedding_layer_names) + if not modules: + return + logger.warning( + "LoRA adapter '%s' targets embedding modules (%s) while EAGLE-family " + "speculative decoding is enabled. The shared draft consumes their " + "base weights, so those deltas do not influence drafting and may " + "reduce the accept rate. Outputs are unaffected.", + lora_name, + ", ".join(modules), + ) + @dataclass class MoELoRABatchInfo: @@ -465,6 +489,19 @@ def get_lm_head_lora_b_shard_size(output_dim: int, shard_indices=None) -> int: return output_dim +def get_batch_token_counts(forward_batch: ForwardBatch) -> Tuple[int, int]: + """(total tokens, max tokens per request) for LoRA segment math.""" + mode = forward_batch.forward_mode + if mode.is_decode(): + return forward_batch.batch_size, 1 + if mode.is_target_verify(): + num_tokens_per_req = forward_batch.spec_info.draft_token_num + return forward_batch.batch_size * num_tokens_per_req, num_tokens_per_req + if mode.is_extend(): + return forward_batch.extend_num_tokens, max(forward_batch.extend_seq_lens_cpu) + raise ValueError(f"Unsupported forward mode: {mode}") + + def generate_sequence_lengths( forward_batch: ForwardBatch, device: Optional[torch.device] = None ) -> torch.Tensor: diff --git a/python/sglang/srt/model_executor/forward_batch_info.py b/python/sglang/srt/model_executor/forward_batch_info.py index 04883e376..a93c5d3d4 100644 --- a/python/sglang/srt/model_executor/forward_batch_info.py +++ b/python/sglang/srt/model_executor/forward_batch_info.py @@ -845,7 +845,7 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin): if ret.forward_mode.is_idle(): ret.positions = torch.empty((0,), dtype=torch.int64, device=device) - if model_runner.server_args.enable_lora: + if model_runner.lora_manager is not None: model_runner.lora_manager.reset_lora_batch() return ret @@ -919,8 +919,8 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin): else: ret._compute_mrope_positions(model_runner, batch) - # Init lora information - if model_runner.server_args.enable_lora: + # Init lora information (None on a draft runner: it is unadapted) + if model_runner.lora_manager is not None: # In the non-LoRA overlap loading case, we fetch LoRA adapters into the memory pool # as a batch, right before running the batch if not model_runner.server_args.enable_lora_overlap_loading: diff --git a/python/sglang/srt/model_executor/model_runner.py b/python/sglang/srt/model_executor/model_runner.py index 017198c76..e9ea1f666 100644 --- a/python/sglang/srt/model_executor/model_runner.py +++ b/python/sglang/srt/model_executor/model_runner.py @@ -319,6 +319,9 @@ class ModelRunner: # earlier publish. if not is_draft_worker: set_global_server_args_for_scheduler(server_args) + # Set by maybe_init_lora_manager; stays None when LoRA is off and on + # draft runners, which serve adapters' target model unadapted. + self.lora_manager: Optional[LoRAManager] = None self.draft_attention_backend = resolve_draft_attention_backend( draft_attention_backend=draft_attention_backend, server_args=server_args, @@ -751,7 +754,8 @@ class ModelRunner: self.apply_torch_tp() def maybe_init_lora_manager(self): - if get_lora().enable_lora: + # Adapters apply to the target model only; the draft runs unadapted. + if get_lora().enable_lora and not self.is_draft_worker: self.init_lora_manager() def maybe_enable_batch_invariant_mode(self): diff --git a/python/sglang/srt/model_executor/model_runner_components/cuda_graph_setup.py b/python/sglang/srt/model_executor/model_runner_components/cuda_graph_setup.py index 0fdd4e5ef..57f365d2c 100644 --- a/python/sglang/srt/model_executor/model_runner_components/cuda_graph_setup.py +++ b/python/sglang/srt/model_executor/model_runner_components/cuda_graph_setup.py @@ -271,7 +271,7 @@ def capture_prefill_graph( return result(eager_runner) if ( - model_runner.server_args.enable_lora + model_runner.lora_manager is not None and not model_runner.lora_manager.supports_prefill_cuda_graph ): logger.warning( diff --git a/python/sglang/srt/model_executor/runner/base_runner.py b/python/sglang/srt/model_executor/runner/base_runner.py index 20426de03..792380547 100644 --- a/python/sglang/srt/model_executor/runner/base_runner.py +++ b/python/sglang/srt/model_executor/runner/base_runner.py @@ -574,7 +574,7 @@ class BaseRunner(ABC): ) # Optional LoRA metadata. - if mr.server_args.enable_lora: + if mr.lora_manager is not None: lora_ids = [None] * batch_size else: lora_ids = None diff --git a/python/sglang/srt/model_executor/runner/decode_cuda_graph_runner.py b/python/sglang/srt/model_executor/runner/decode_cuda_graph_runner.py index 6f27a0594..fedc58d3c 100644 --- a/python/sglang/srt/model_executor/runner/decode_cuda_graph_runner.py +++ b/python/sglang/srt/model_executor/runner/decode_cuda_graph_runner.py @@ -348,7 +348,7 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner): self._captured_ragged_layouts: dict[int, object] = {} if self.ragged_verify_mode and ( self.enable_two_batch_overlap - or model_runner.server_args.enable_lora + or model_runner.lora_manager is not None or self.disable_padding ): raise ValueError( @@ -381,7 +381,7 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner): if self.enable_torch_compile: set_torch_compile_config() - if self.model_runner.server_args.enable_lora: + if self.model_runner.lora_manager is not None: # Phase 2 of LoRA CUDA graph init: dense LoRA batch metadata. # Phase 1 (MoE buffers) was handled earlier in ModelRunner via # lora_manager.init_cuda_graph_moe_buffers(). @@ -927,7 +927,7 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner): spec_info, ) - if self.model_runner.server_args.enable_lora: + if self.model_runner.lora_manager is not None: # It is safe to capture CUDA graph using empty LoRA id, as the LoRA kernels will always be launched whenever # `--enable-lora` is set to True (and return immediately if the LoRA id is empty for perf optimization). lora_ids = [None] * bs diff --git a/python/sglang/srt/model_executor/runner/prefill_cuda_graph_runner.py b/python/sglang/srt/model_executor/runner/prefill_cuda_graph_runner.py index a97adb875..8803a3f2f 100644 --- a/python/sglang/srt/model_executor/runner/prefill_cuda_graph_runner.py +++ b/python/sglang/srt/model_executor/runner/prefill_cuda_graph_runner.py @@ -256,7 +256,7 @@ class PrefillCudaGraphRunner(BaseCudaGraphRunner): # --- model flags ---------------------------------------------- self.quant_config = getattr(model_runner.model, "quant_config", None) self.is_multimodal = model_runner.model_config.is_multimodal - self.enable_lora = model_runner.server_args.enable_lora + self.enable_lora = model_runner.lora_manager is not None # Classification/reward forwards branch on return_pooled_hidden_states; # capture must use the same flag value as replay for those models. self.capture_return_pooled_hidden_states = not model_runner.is_generation diff --git a/python/sglang/srt/server_args.py b/python/sglang/srt/server_args.py index a97ca3cac..5c35a6eb6 100644 --- a/python/sglang/srt/server_args.py +++ b/python/sglang/srt/server_args.py @@ -319,6 +319,10 @@ RETRACTION_POLICY_CHOICES = ["length", "priority"] RL_ON_POLICY_TARGET_CHOICES = ["fsdp"] +# Speculative algorithms whose verify forward presents a uniform per-request +# token width, which is what the LoRA segment layout assumes. +_LORA_SPEC_ALGORITHMS = ("EAGLE", "EAGLE3", "DFLASH", "DSPARK") + LORA_BACKEND_CHOICES = ["triton", "csgmv", "ascend", "torch_native"] ENCODER_TRANSFER_BACKEND_CHOICES = [ @@ -9451,10 +9455,7 @@ class ServerArgs: ) # Validate compatibility with speculative decoding - if self.speculative_algorithm not in ["NGRAM", None]: - raise ValueError( - "Currently LoRA is only compatible with NGRAM speculative decoding." - ) + self._check_lora_speculative_compatibility() # Parse lora_paths if isinstance(self.lora_paths, list): @@ -9559,6 +9560,65 @@ class ServerArgs: self.lora_drain_wait_threshold >= 0.0 ), "--lora-drain-wait-threshold must be non-negative." + def _check_lora_speculative_compatibility(self): + """Validate LoRA + speculative decoding combinations. + + Adapters apply to the target only; a shared draft runs unadapted. + Matches resolved algorithm names (NEXTN has collapsed to EAGLE). + """ + if self.speculative_algorithm in ["NGRAM", None]: + return + + if self.speculative_algorithm not in _LORA_SPEC_ALGORITHMS: + promoted = ( + " (NEXTN/EAGLE with a Gemma4 assistant draft is automatically " + "promoted to FROZEN_KV_MTP, which does not support LoRA)" + if self.speculative_algorithm == "FROZEN_KV_MTP" + else "" + ) + raise ValueError( + "LoRA is only compatible with NGRAM, EAGLE, NEXTN, EAGLE3, " + "DFLASH, or DSPARK speculative decoding, not " + f"{self.speculative_algorithm}{promoted}." + ) + + ragged_mode = envs.SGLANG_RAGGED_VERIFY_MODE.get() + + # Each entry: (is unsupported, why). Reasons are appended to a shared + # prefix so the message names the combination, not just the flag. + unsupported = [ + ( + self.speculative_algorithm == "DSPARK" and ragged_mode != "static", + f"does not support SGLANG_RAGGED_VERIFY_MODE={ragged_mode!r}: " + "the per-request verify lengths it schedules break the " + "uniform-width LoRA segment layout", + ), + ( + self.speculative_adaptive, + "does not support --speculative-adaptive: the draft is built " + "from a static ServerArgs snapshot, and the runtime-state " + "swap does not rebuild LoRA cuda-graph metadata", + ), + ( + "experimental_sgl_trtllm" + in (self.moe_runner_backend, self.speculative_moe_runner_backend), + "does not support the experimental_sgl_trtllm MoE runner: its " + "TopK reads the LoRA config per forward, which the draft " + "resolves against the target's after its own publish ended", + ), + ( + envs.SGLANG_ENABLE_OVERLAP_PLAN_STREAM.get(), + "does not support SGLANG_ENABLE_OVERLAP_PLAN_STREAM=1: LoRA " + "batch preparation would run on the plan stream, unordered " + "against in-flight forwards", + ), + ] + for is_unsupported, reason in unsupported: + if is_unsupported: + raise ValueError( + f"LoRA with EAGLE/NEXTN/EAGLE3 speculative decoding {reason}." + ) + def validate_buckets_rule(self, arg_name: str, buckets_rule: List[str]): if not buckets_rule: return diff --git a/python/sglang/srt/speculative/dflash_worker_v2.py b/python/sglang/srt/speculative/dflash_worker_v2.py index 98bee8b5a..59aedd3bb 100644 --- a/python/sglang/srt/speculative/dflash_worker_v2.py +++ b/python/sglang/srt/speculative/dflash_worker_v2.py @@ -22,6 +22,7 @@ from sglang.srt.distributed.parallel_state_wrapper import ParallelState from sglang.srt.environ import envs from sglang.srt.layers.logits_processor import should_apply_lm_head_quant_method from sglang.srt.layers.logprob_processor import compute_spec_logprobs +from sglang.srt.lora.layers import unwrap_lora_layer from sglang.srt.managers.schedule_batch import ScheduleBatch from sglang.srt.managers.scheduler import GenerationBatchResult from sglang.srt.managers.tp_worker import TpModelWorker @@ -480,7 +481,7 @@ class DFlashWorkerV2(BaseSpecWorker): if self.block_size <= 1: return _eager("block_size<=1") target_model = self._target_worker.model_runner.model - lm_head = getattr(target_model, "lm_head", None) + lm_head = unwrap_lora_layer(getattr(target_model, "lm_head", None)) if lm_head is None: return _eager("no target lm_head") @@ -1749,8 +1750,8 @@ class DFlashWorkerV2(BaseSpecWorker): # --- 1) Draft a fixed block with the draft model. target_model = self.target_worker.model_runner.model - embed_module = target_model.get_input_embeddings() - lm_head = getattr(target_model, "lm_head", None) + embed_module = unwrap_lora_layer(target_model.get_input_embeddings()) + lm_head = unwrap_lora_layer(getattr(target_model, "lm_head", None)) if lm_head is None or not ( hasattr(lm_head, "weight") or callable(getattr(getattr(lm_head, "quant_method", None), "apply", None)) diff --git a/python/sglang/srt/speculative/dspark_components/dspark_draft.py b/python/sglang/srt/speculative/dspark_components/dspark_draft.py index da261b9e8..5514c7d49 100644 --- a/python/sglang/srt/speculative/dspark_components/dspark_draft.py +++ b/python/sglang/srt/speculative/dspark_components/dspark_draft.py @@ -11,6 +11,7 @@ from sglang.kernels.ops.speculative.dspark.dspark_draft_model import ( SampleStepTokens, ) from sglang.srt.environ import envs +from sglang.srt.lora.layers import unwrap_lora_layer from sglang.srt.managers.schedule_batch import ScheduleBatch from sglang.srt.model_executor.forward_batch_info import ( CaptureHiddenMode, @@ -201,7 +202,7 @@ class DraftBlockProposer: target_model, sampling_info, ) -> DraftProposal: - embed_module = target_model.get_input_embeddings() + embed_module = unwrap_lora_layer(target_model.get_input_embeddings()) draft_sampler = self._draft_sampler all_greedy = sampling_info is None or sampling_info.is_all_greedy fwd = self._run_forward( diff --git a/python/sglang/srt/speculative/dspark_components/dspark_worker_v2.py b/python/sglang/srt/speculative/dspark_components/dspark_worker_v2.py index 4788947a3..462db6d61 100644 --- a/python/sglang/srt/speculative/dspark_components/dspark_worker_v2.py +++ b/python/sglang/srt/speculative/dspark_components/dspark_worker_v2.py @@ -12,6 +12,7 @@ from sglang.srt.configs.hybrid_arch import mambaish_config from sglang.srt.distributed.parallel_state_wrapper import ParallelState from sglang.srt.environ import envs from sglang.srt.layers.logprob_processor import compute_spec_logprobs +from sglang.srt.lora.layers import unwrap_lora_layer from sglang.srt.managers.schedule_batch import ScheduleBatch from sglang.srt.managers.scheduler import GenerationBatchResult from sglang.srt.managers.tp_worker import TpModelWorker @@ -190,13 +191,15 @@ class DSparkWorkerV2(BaseSpecWorker): ) else: target_model = self.target_worker.model_runner.model - lm_head = getattr(target_model, "lm_head", None) + lm_head = unwrap_lora_layer(getattr(target_model, "lm_head", None)) if lm_head is None or not hasattr(lm_head, "weight"): raise RuntimeError( "DSpark requires the target model to expose `lm_head` with `weight`." ) self.draft_model.attach_shared_modules( - embed_tokens=self._resolve_target_embed_tokens(target_model), + embed_tokens=unwrap_lora_layer( + self._resolve_target_embed_tokens(target_model) + ), lm_head=lm_head, ) diff --git a/python/sglang/srt/speculative/eagle_worker_v2.py b/python/sglang/srt/speculative/eagle_worker_v2.py index ad6ceb23e..03a595ee4 100644 --- a/python/sglang/srt/speculative/eagle_worker_v2.py +++ b/python/sglang/srt/speculative/eagle_worker_v2.py @@ -278,8 +278,12 @@ class EagleDraftWorker(EagleDraftWorkerBase): self.hot_token_id = None def init_lm_head(self): + from sglang.srt.lora.layers import unwrap_lora_layer + embed, head = self.target_worker.model_runner.model.get_embed_and_head() - target_lm_head = getattr(self.target_worker.model_runner.model, "lm_head", None) + target_lm_head = unwrap_lora_layer( + getattr(self.target_worker.model_runner.model, "lm_head", None) + ) def maybe_share_target_lm_head(): if ( diff --git a/test/manual/lora/check_spec_baseline_divergence.py b/test/manual/lora/check_spec_baseline_divergence.py new file mode 100644 index 000000000..a7a2c7986 --- /dev/null +++ b/test/manual/lora/check_spec_baseline_divergence.py @@ -0,0 +1,125 @@ +"""Does a spec-on vs spec-off divergence need LoRA at all? + +Run this when the spec+LoRA harness +(``test/manual/lora/run_spec_lora_matrix.py``) reports a mismatch that +survives its filters. It runs the *same* speculative config with **no +adapters loaded**, comparing spec-off against spec-on greedy outputs and +sampling each side twice so run-to-run instability is visible. + +If the base model diverges on the same prompt, the cause is the model and +the speculative path (EAGLE3 topk>1 tree verify in particular), not the LoRA +integration -- which is exactly what it showed for +Qwen3-30B-A3B-Instruct-2507 on prompt #2: DIFFERS, and UNSTABLE-self, with +zero adapters involved. + +The harness's own filters cannot answer this: they can tell whether a pair +reproduces within one server config, but not whether the divergence depends +on LoRA being present. That needs this second config. + +Usage (from a checkout, with a GPU): + python test/manual/lora/check_spec_baseline_divergence.py +""" + +import os +import subprocess +import sys +import time + +import requests + +PROMPTS = [ + "What is the capital of France? Answer in one sentence.", + "List three primary colors.", + "Write a one-sentence story about a brave detective on Mars.", + "Explain what a hash table is in two sentences.", +] +BASE = "http://127.0.0.1:31000" +COMMON = [ + "--tp", + "4", + "--moe-runner-backend", + "triton", + "--attention-backend", + "flashinfer", + "--prefill-attention-backend", + "fa4", + "--decode-attention-backend", + "fa4", + "--mem-fraction-static", + "0.8", +] +SPEC = [ + "--speculative-algorithm", + "EAGLE3", + "--speculative-draft-model-path", + "lmsys/SGLang-EAGLE3-Qwen3-30B-A3B-Instruct-2507-SpecForge-Nex", + "--speculative-num-steps", + "3", + "--speculative-eagle-topk", + "4", + "--speculative-num-draft-tokens", + "8", +] + + +def launch(extra): + env = dict(os.environ, SGLANG_ALLOW_OVERWRITE_LONGER_CONTEXT_LEN="1") + p = subprocess.Popen( + [ + sys.executable, + "-m", + "sglang.launch_server", + "--model-path", + "Qwen/Qwen3-30B-A3B-Instruct-2507", + "--port", + "31000", + ] + + COMMON + + extra, + stdout=open("/scratch/loraspec/logs/attrib_server.log", "a"), + stderr=subprocess.STDOUT, + env=env, + ) + for _ in range(120): + time.sleep(10) + try: + if requests.get(BASE + "/health", timeout=3).ok: + return p + except Exception: + pass + raise SystemExit("server did not come up") + + +def gen(): + r = requests.post( + BASE + "/generate", + json={ + "text": PROMPTS, + "sampling_params": {"temperature": 0, "max_new_tokens": 32}, + }, + timeout=900, + ) + r.raise_for_status() + return [x["text"] for x in r.json()] + + +results = {} +for label, extra in [("nospec", []), ("spec", SPEC)]: + p = launch(extra) + try: + results[label] = [gen(), gen()] + finally: + p.terminate() + p.wait(timeout=120) + time.sleep(10) + +print("=" * 70) +for i, prompt in enumerate(PROMPTS): + a1, a2 = results["nospec"][0][i], results["nospec"][1][i] + b1, b2 = results["spec"][0][i], results["spec"][1][i] + stable = "stable" if (a1 == a2 and b1 == b2) else "UNSTABLE-self" + verdict = "same" if a1 == b1 else "DIFFERS" + print(f"prompt#{i}: nospec-vs-spec={verdict} ({stable})") + if a1 != b1: + print(f" nospec: {a1!r}") + print(f" spec : {b1!r}") diff --git a/test/manual/lora/run_spec_lora_matrix.py b/test/manual/lora/run_spec_lora_matrix.py new file mode 100644 index 000000000..54c32403b --- /dev/null +++ b/test/manual/lora/run_spec_lora_matrix.py @@ -0,0 +1,876 @@ +"""Manual validation harness: multi-adapter LoRA + EAGLE-family speculative decoding. + +The oracle is spec-on vs spec-off *per adapter*, not adapter vs base. +Speculative decoding is lossless with respect to the target it verifies +against, so for greedy sampling: + + output(adapter=X, spec=ON) == output(adapter=X, spec=OFF) + +Comparing an adapter's output to the base model's only proves LoRA was +applied at all. This harness launches the same weights twice (reference +without spec, then with spec) and runs four checks: + + 1. parity — per adapter, spec-on output == spec-off output (the oracle) + 2. distinct — each adapter's output differs from base (LoRA really applied) + 3. mixed — a batch interleaving every adapter matches the solo outputs + (crossed verify segments serve a request the wrong adapter) + 4. eager — a batch wider than --cuda-graph-max-bs still matches + (exercises the non-cuda-graph target-verify path) + +It also reports per-adapter accept length from each response's +meta_info["spec_accept_length"], which is how much speculation the shared +unadapted draft actually buys for that adapter. + +Usage: + python test/manual/lora/run_spec_lora_matrix.py --config eagle3-llama31 + python test/manual/lora/run_spec_lora_matrix.py --config nextn-qwen35-35b-a3b + python test/manual/lora/run_spec_lora_matrix.py --list + +Add --keep-derived to reuse a previously written derived adapter, and +--max-new-tokens / --port to adjust the run. +""" + +import argparse +import hashlib +import json +import os +import random +import shutil +import sys +import time +from concurrent import futures + +import requests + +from sglang.srt.utils import kill_process_tree +from sglang.test.test_utils import ( + DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH, + popen_launch_server, +) + +PROMPTS = [ + "What is the capital of France? Answer in one sentence.", + "List three primary colors.", + "Write a one-sentence story about a brave detective on Mars.", + "Explain what a hash table is in two sentences.", +] + +DERIVED_SUFFIX = "-derived" + + +# Each config is a runnable server shape. `adapters` are (name, hf_path) +# pairs; `derive_from` names an adapter to synthesize a second, distinct +# adapter from when the base has only one public adapter (needed to make +# mixed-adapter batches meaningful). +CONFIGS = { + # Cheapest: 5 genuinely distinct public adapters, ranks 8-64, no synthesis. + "eagle3-llama31": dict( + model="meta-llama/Llama-3.1-8B-Instruct", + spec_args=[ + "--speculative-algorithm=EAGLE3", + "--speculative-draft-model-path=lmsys/sglang-EAGLE3-LLaMA3.1-Instruct-8B", + "--speculative-num-steps=5", + "--speculative-eagle-topk=8", + "--speculative-num-draft-tokens=32", + ], + adapters=[ + ("fact", "algoprog/fact-generation-llama-3.1-8b-instruct-lora"), + ("guard", "nvidia/llama-3.1-nemoguard-8b-topic-control"), + ("sql", "philschmid/code-llama-3-1-8b-text-to-sql-lora"), + ("ocr", "pbevan11/llama-3.1-8b-ocr-correction"), + ("zh", "faridlazuarda/valadapt-llama-3.1-8B-it-chinese"), + ], + common_args=["--mem-fraction-static=0.7", "--max-lora-rank=64"], + tp=1, + ), + # EAGLE topk=1 is the same chain-shaped drafting NEXTN/MTP uses, on 1 GPU + # without needing an MTP checkpoint. cuda-graph-max-bs forces the eager path. + "eagle-topk1-llama2": dict( + model="meta-llama/Llama-2-7b-chat-hf", + spec_args=[ + "--speculative-algorithm=EAGLE", + "--speculative-draft-model-path=lmsys/sglang-EAGLE-llama2-chat-7B", + "--speculative-num-steps=3", + "--speculative-eagle-topk=1", + "--speculative-num-draft-tokens=4", + ], + adapters=[("norwegian", "RuterNorway/Llama-2-7b-chat-norwegian-LoRa")], + derive_from="norwegian", + common_args=[ + "--mem-fraction-static=0.7", + "--max-lora-rank=128", + "--cuda-graph-max-bs=2", + ], + tp=1, + ), + # Real NEXTN (self-bundled MTP head, no draft path) on a MoE base: the + # densest coverage per GPU-hour — MoE-LoRA cuda-graph buffers, virtual + # experts, and an arch whose MTP head shares the target lm_head module. + "nextn-qwen35-35b-a3b": dict( + model="Qwen/Qwen3.5-35B-A3B", + spec_args=[ + "--speculative-algorithm=NEXTN", + "--speculative-num-steps=3", + "--speculative-eagle-topk=1", + "--speculative-num-draft-tokens=4", + ], + adapters=[("case", "opherlie/lora-test-case-Qwen3.5-35B-A3B", "dataset")], + derive_from="case", + common_args=[ + "--max-lora-rank=64", + "--moe-runner-backend=triton", + "--experts-shared-outer-loras", + "--lora-use-virtual-experts", + "--disable-shared-experts-fusion", + "--mem-fraction-static=0.8", + ], + tp=4, + ), + # MoE + EAGLE3 with a draft checkpoint matched to the same base variant. + "eagle3-qwen3-30b-a3b": dict( + model="Qwen/Qwen3-30B-A3B-Instruct-2507", + spec_args=[ + "--speculative-algorithm=EAGLE3", + "--speculative-draft-model-path=lmsys/SGLang-EAGLE3-Qwen3-30B-A3B-Instruct-2507-SpecForge-Nex", + "--speculative-num-steps=3", + "--speculative-eagle-topk=4", + "--speculative-num-draft-tokens=8", + ], + adapters=[ + ("case", "yushengsu/lora-diff-Qwen3-30B-A3B-Instruct-2507", "dataset") + ], + derive_from="case", + common_args=[ + "--max-lora-rank=32", + "--moe-runner-backend=triton", + "--experts-shared-outer-loras", + "--attention-backend=flashinfer", + "--prefill-attention-backend=fa4", + "--decode-attention-backend=fa4", + "--mem-fraction-static=0.8", + ], + # The draft checkpoint's config derives a 2048 context; the draft + # follows the target's context_length (it reads target KV), so the + # longer-context guard has to be waived for this pairing. + env={"SGLANG_ALLOW_OVERWRITE_LONGER_CONTEXT_LEN": "1"}, + tp=4, + ), + # DFLASH: fixed-size block drafting, verified in one TARGET_VERIFY pass. + # Its target is the one base with several genuinely distinct public + # adapters, so no synthesis is needed. + "dflash-llama31": dict( + model="meta-llama/Llama-3.1-8B-Instruct", + spec_args=[ + "--speculative-algorithm=DFLASH", + "--speculative-draft-model-path=z-lab/LLaMA3.1-8B-Instruct-DFlash-UltraChat", + "--speculative-num-draft-tokens=4", + ], + adapters=[ + ("fact", "algoprog/fact-generation-llama-3.1-8b-instruct-lora"), + ("guard", "nvidia/llama-3.1-nemoguard-8b-topic-control"), + ], + common_args=["--mem-fraction-static=0.7", "--max-lora-rank=64"], + # The draft checkpoint derives a 40960 context against the target's + # 131072; the draft follows the target, so waive the guard. + env={"SGLANG_ALLOW_OVERWRITE_LONGER_CONTEXT_LEN": "1"}, + tp=1, + ), + # DSPARK with a separate DSpark draft head, on the one base that has two + # distinct public adapters. Only the default static ragged-verify mode is + # supported: the others schedule per-request verify lengths, which the + # uniform-width LoRA segment layout cannot express. + "dspark-inkling": dict( + model="thinkingmachines/Inkling-Small", + spec_args=[ + "--speculative-algorithm=DSPARK", + "--speculative-draft-model-path=RadixArk/Inkling-Small-DSpark", + "--speculative-dspark-block-size=5", + ], + adapters=[ + ("gutenberg", "nbeerbower/Inkling-Small-Gutenberg-DPO-LoRA"), + ("hemlock", "hemlang/Inkling-Small-Hemlock-SFT-LoRA"), + ], + common_args=[ + "--max-lora-rank=32", + "--moe-runner-backend=triton", + "--experts-shared-outer-loras", + "--mem-fraction-static=0.8", + ], + env={"SGLANG_ALLOW_OVERWRITE_LONGER_CONTEXT_LEN": "1"}, + # A ~500GB checkpoint over network storage does not load inside the + # default launch budget. + launch_timeout=3600, + tp=8, + ), + # Flagship MLA + MoE NextN. 8 GPUs. + "nextn-deepseek-v31": dict( + model="deepseek-ai/DeepSeek-V3.1-Base", + spec_args=[ + "--speculative-algorithm=NEXTN", + "--speculative-num-steps=3", + "--speculative-eagle-topk=1", + "--speculative-num-draft-tokens=4", + ], + adapters=[("case", "yushengsu/lora-diff-DeepSeek-V3.1-Base", "dataset")], + derive_from="case", + common_args=[ + "--max-lora-rank=32", + "--moe-runner-backend=triton", + "--experts-shared-outer-loras", + "--attention-backend=flashinfer", + "--prefill-attention-backend=fa4", + "--decode-attention-backend=flashinfer", + "--disable-shared-experts-fusion", + "--mem-fraction-static=0.8", + ], + tp=8, + ), +} + + +def log(msg: str) -> None: + print(f"[harness] {msg}", flush=True) + + +def materialize_adapter(hf_path: str, repo_type: str) -> str: + """Download an adapter and return its local directory. + + Several adapters used here live in *dataset* repos (that is how the + repo's own LoRA logprob-diff tests host them), which the server cannot + resolve as a model id — so every adapter is materialized locally and + passed to --lora-paths by path. + """ + from huggingface_hub import snapshot_download + + return snapshot_download(repo_id=hf_path, repo_type=repo_type) + + +def check_adapter_supported(name: str, local_path: str) -> None: + """Report what each adapter exercises. Nothing here is a rejection: a + draft sharing the target's lm_head gets the unwrapped base layer, so + lm_head and embedding adapters cost accept rate, not correctness.""" + config_path = os.path.join(local_path, "adapter_config.json") + try: + with open(config_path) as f: + config = json.load(f) + except Exception as e: # noqa: BLE001 - not every layout ships one + log(f" {name}: could not read adapter_config.json ({e}); skipping precheck") + return + modules = config.get("target_modules") + log(f" {name}: r={config.get('r')} target_modules={modules}") + if isinstance(modules, list): + shared = {"lm_head", "output", "unembed_tokens"} & set(modules) + if shared: + log(f" {name}: targets {sorted(shared)} — shared-lm_head path") + if {"embed_tokens", "vocab_emb", "word_embeddings"} & set(modules): + log(f" {name}: targets embeddings — expect a reduced accept rate") + if config.get("modules_to_save"): + log(f" {name}: WARNING modules_to_save={config['modules_to_save']}") + + +LM_HEAD_MARKERS = ("lm_head", "unembed_tokens") + + +def adapter_lm_head_keys(local_path: str): + """Tensor keys that would land on lm_head (PEFT unembed_tokens keys are + rewritten to lm_head during load).""" + from safetensors import safe_open + + weights = os.path.join(local_path, "adapter_model.safetensors") + if not os.path.isfile(weights): + return [] + with safe_open(weights, "pt") as f: + return [k for k in f.keys() if any(m in k for m in LM_HEAD_MARKERS)] + + +def write_adapter_variant( + source_dir: str, out_dir: str, keep: bool, *, negate_b: bool, drop_lm_head: bool +) -> str: + """Write a modified copy of an adapter. + + ``drop_lm_head`` removes output-layer tensors, isolating an adapter's + other modules from the shared-lm_head path. The server accepts lm_head + adapters, so this is for narrowing a failure, not a requirement. + + ``negate_b`` flips the sign of every lora_B tensor, producing a second + adapter of identical shape whose outputs genuinely diverge from the + original — without which a "mixed-adapter" batch would compare an adapter + against itself and could not detect crossed verify segments. + """ + if keep and os.path.isdir(out_dir): + log(f"reusing adapter variant at {out_dir}") + return out_dir + + from safetensors.torch import load_file, save_file + + weights_name = "adapter_model.safetensors" + src_weights = os.path.join(source_dir, weights_name) + if not os.path.isfile(src_weights): + raise SystemExit(f"{source_dir} has no {weights_name}") + + if os.path.isdir(out_dir): + shutil.rmtree(out_dir) + os.makedirs(out_dir) + + tensors = load_file(src_weights) + dropped = flipped = 0 + out = {} + for key, value in tensors.items(): + if drop_lm_head and any(m in key for m in LM_HEAD_MARKERS): + dropped += 1 + continue + if negate_b and "lora_B" in key: + value = -value + flipped += 1 + out[key] = value + save_file(out, os.path.join(out_dir, weights_name)) + shutil.copy( + os.path.join(source_dir, "adapter_config.json"), + os.path.join(out_dir, "adapter_config.json"), + ) + log( + f"wrote {out_dir} (dropped {dropped} lm_head tensors, " + f"negated {flipped} lora_B tensors)" + ) + return out_dir + + +def launch(config: dict, adapters, base_url: str, with_spec: bool): + other_args = list(config["common_args"]) + env = config.get("env") + if config["tp"] > 1: + other_args += [f"--tp={config['tp']}"] + if adapters: + other_args += [ + "--enable-lora", + f"--lora-backend={config.get('lora_backend', 'triton')}", + # +1: the base model occupies a pool slot too, and every mixed + # batch below co-batches base requests. + f"--max-loras-per-batch={len(adapters) + 1}", + "--lora-paths", + ] + [f"{name}={path}" for name, path in adapters] + else: + # Base arms: --max-lora-rank is a LoRA-only flag and is rejected + # without --enable-lora, so drop it from the common set. + other_args = [a for a in other_args if not a.startswith("--max-lora-rank")] + if with_spec: + other_args += config["spec_args"] + other_args += config.get("extra_args", []) + log(f"launching {'spec' if with_spec else 'reference'} server: {other_args}") + return popen_launch_server( + config["model"], + base_url, + timeout=config.get("launch_timeout", DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH), + other_args=other_args, + env=env, + ) + + +def generate(base_url: str, texts, routes, max_new_tokens: int): + """One /generate call; returns (texts, per-request accept lengths).""" + response = requests.post( + base_url + "/generate", + json={ + "text": texts, + "lora_path": routes, + "sampling_params": {"temperature": 0, "max_new_tokens": max_new_tokens}, + }, + timeout=1800, + ) + response.raise_for_status() + results = response.json() + outputs = [item["text"] for item in results] + accepts = [item["meta_info"].get("spec_accept_length") for item in results] + return outputs, accepts + + +def min_logprob_gap(base_url: str, route, prompt: str, max_new_tokens: int): + """Smallest top-2 logprob gap along the greedy path for one request. + + A gap of ~0 means the argmax at that position is decided by kernel scan + order, so any change in reduction order (a verify forward computes logits + at a different batch shape than a plain decode) flips the token with no + correctness implication. Used to classify surviving mismatches. + """ + response = requests.post( + base_url + "/generate", + json={ + "text": prompt, + "lora_path": route, + "sampling_params": {"temperature": 0, "max_new_tokens": max_new_tokens}, + "return_logprob": True, + "top_logprobs_num": 2, + }, + timeout=600, + ) + response.raise_for_status() + tops = response.json()["meta_info"].get("output_top_logprobs") or [] + gaps = [c[0][0] - c[1][0] for c in tops if c and len(c) >= 2] + return min(gaps) if gaps else None + + +def collect_solo(base_url: str, routes, max_new_tokens: int): + """Per route: outputs and mean accept length over PROMPTS.""" + solo, accept = {}, {} + for route in routes: + outputs, accepts = generate( + base_url, list(PROMPTS), [route] * len(PROMPTS), max_new_tokens + ) + solo[route] = outputs + seen = [a for a in accepts if a] + accept[route] = sum(seen) / len(seen) if seen else None + return solo, accept + + +def report_parity(label: str, reference: dict, actual: dict) -> int: + """Compare per-route output lists; returns the mismatch count.""" + mismatches = 0 + for route, expected in reference.items(): + for i, (want, got) in enumerate(zip(expected, actual[route])): + if want != got: + mismatches += 1 + log(f" MISMATCH [{label}] route={route} prompt#{i}") + log(f" spec-off: {want!r}") + log(f" spec-on : {got!r}") + return mismatches + + +def build_load(num_prompts: int, input_len: int, seed: int): + """Fixed synthetic prompts as raw token ids, so every arm sends identical + work and the input length is exact rather than tokenizer-dependent.""" + rng = random.Random(seed) + return [ + [rng.randint(1000, 10000) for _ in range(input_len)] for _ in range(num_prompts) + ] + + +def one_request(base_url: str, input_ids, route, output_len: int): + started = time.perf_counter() + response = requests.post( + base_url + "/generate", + json={ + "input_ids": input_ids, + "lora_path": route, + "sampling_params": { + "temperature": 0, + "max_new_tokens": output_len, + # Fixed-length output: every arm emits exactly output_len + # tokens, so throughput is comparable without length drift. + "ignore_eos": True, + }, + }, + timeout=3600, + ) + response.raise_for_status() + meta = response.json()["meta_info"] + return dict( + latency=time.perf_counter() - started, + output_tokens=meta.get("completion_tokens", output_len), + accept_length=meta.get("spec_accept_length"), + ) + + +def measure(base_url: str, loads, routes, output_len: int, concurrency: int): + """Drive `concurrency` in-flight requests over the load; return throughput, + accept length, and per-request latency.""" + pairs = [(loads[i], routes[i % len(routes)]) for i in range(len(loads))] + started = time.perf_counter() + with futures.ThreadPoolExecutor(max_workers=concurrency) as pool: + results = list( + pool.map( + lambda pair: one_request(base_url, pair[0], pair[1], output_len), pairs + ) + ) + elapsed = time.perf_counter() - started + accepts = [r["accept_length"] for r in results if r["accept_length"]] + total_out = sum(r["output_tokens"] for r in results) + latencies = sorted(r["latency"] for r in results) + return dict( + output_tps=total_out / elapsed, + accept_length=(sum(accepts) / len(accepts)) if accepts else None, + median_latency=latencies[len(latencies) // 2], + elapsed=elapsed, + requests=len(results), + ) + + +def run_perf(config: dict, args) -> int: + """Four arms answering two questions: does speculation still pay once LoRA + is on (lora_spec vs lora_nospec), and how much does serving adapters cost + relative to speculating on the plain base model (lora_spec vs base_spec)? + """ + base_url = f"http://127.0.0.1:{args.port}" + adapters = resolve_adapters(config, args) + concurrencies = [int(c) for c in args.concurrency.split(",")] + loads = build_load(args.num_prompts, args.input_len, args.seed) + warmup = build_load(max(concurrencies), args.input_len, args.seed + 1) + + # (label, serve LoRA?, speculate?) + arms = [ + ("base_nospec", False, False), + ("base_spec", False, True), + ("lora_nospec", True, False), + ("lora_spec", True, True), + ] + table = {} + + for label, with_lora, with_spec in arms: + log("=" * 68) + log(f"arm {label}: lora={with_lora} spec={with_spec}") + process = None + try: + process = launch( + config, + adapters if with_lora else [], + base_url, + with_spec=with_spec, + ) + # Multi-adapter arms spread requests round-robin across every + # adapter (the multi-tenant shape); base arms send no adapter. + routes = [name for name, _ in adapters] if with_lora else [None] + measure(base_url, warmup, routes, args.output_len, len(warmup)) + for concurrency in concurrencies: + stats = measure(base_url, loads, routes, args.output_len, concurrency) + table[(label, concurrency)] = stats + log( + f" c={concurrency:<4} {stats['output_tps']:8.1f} tok/s " + f"accept={stats['accept_length']} " + f"p50={stats['median_latency']:.2f}s" + ) + finally: + if process is not None: + kill_process_tree(process.pid) + + report_perf(table, concurrencies, args) + return 0 + + +def report_perf(table, concurrencies, args) -> None: + log("=" * 78) + log( + f"PERF model-shape: {args.num_prompts} reqs x {args.input_len} in / " + f"{args.output_len} out" + ) + log("=" * 78) + labels = ["base_nospec", "base_spec", "lora_nospec", "lora_spec"] + log(f"{'concurrency':<12}" + "".join(f"{label:>16}" for label in labels)) + for concurrency in concurrencies: + cells = [] + for label in labels: + stats = table.get((label, concurrency)) + cells.append(f"{stats['output_tps']:.1f}" if stats else "-") + log( + f"{concurrency:<12}" + "".join(f"{cell:>16}" for cell in cells) + " tok/s" + ) + log("") + log(f"{'concurrency':<12}{'base accept':>16}{'lora accept':>16}") + for concurrency in concurrencies: + base = table.get(("base_spec", concurrency), {}).get("accept_length") + lora = table.get(("lora_spec", concurrency), {}).get("accept_length") + fmt = lambda v: f"{v:.3f}" if v else "-" # noqa: E731 + log(f"{concurrency:<12}{fmt(base):>16}{fmt(lora):>16}") + log("") + log("ratios (>1 is better)") + log( + f"{'concurrency':<12}{'spec on base':>16}{'spec on lora':>16}" + f"{'lora vs base':>16}" + ) + for concurrency in concurrencies: + + def tps(label): + stats = table.get((label, concurrency)) + return stats["output_tps"] if stats else None + + base_nospec, base_spec = tps("base_nospec"), tps("base_spec") + lora_nospec, lora_spec = tps("lora_nospec"), tps("lora_spec") + ratio = lambda a, b: f"{a / b:.2f}x" if a and b else "-" # noqa: E731 + log( + f"{concurrency:<12}" + f"{ratio(base_spec, base_nospec):>16}" + f"{ratio(lora_spec, lora_nospec):>16}" + f"{ratio(lora_spec, base_spec):>16}" + ) + log("") + log(" spec on base = speculation speedup without adapters (the ceiling)") + log(" spec on lora = speculation speedup with adapters served (the ask)") + log(" lora vs base = cost of serving adapters, both speculating") + + +def resolve_adapters(config: dict, args): + """Materialize every adapter locally and return (name, local_path) pairs, + deriving a second distinct adapter when the base has only one public one. + + Config entries are (name, hf_path) or (name, hf_path, repo_type). + """ + adapters = [] + for entry in config["adapters"]: + name, hf_path = entry[0], entry[1] + repo_type = entry[2] if len(entry) > 2 else "model" + local_path = materialize_adapter(hf_path, repo_type) + + # A draft sharing the target's lm_head gets the unwrapped base layer, + # so these adapters cost accept rate, not correctness. Keep the + # weights as published so the run exercises that path. + lm_head_keys = adapter_lm_head_keys(local_path) + if lm_head_keys: + log( + f" {name}: carries {len(lm_head_keys)} lm_head tensor(s) " + f"({lm_head_keys[0]}) — exercising the shared-lm_head path" + ) + adapters.append((name, local_path)) + + log(f"prechecking {len(adapters)} adapter(s) against the lm_head rule") + for name, local_path in adapters: + check_adapter_supported(name, local_path) + + derive_name = config.get("derive_from") + if derive_name: + source = dict(adapters)[derive_name] + # Key the cache on the source path: several configs name their + # adapter "case", so a name-only key served one model's variant to + # another (mismatched layer counts). + source_key = hashlib.sha1(source.encode()).hexdigest()[:8] + adapters.append( + ( + derive_name + DERIVED_SUFFIX, + write_adapter_variant( + source, + os.path.join( + args.derived_dir, f"{derive_name}{DERIVED_SUFFIX}-{source_key}" + ), + args.keep_derived, + negate_b=True, + drop_lm_head=False, + ), + ) + ) + return adapters + + +def collect_shapes(base_url: str, routes, args): + """Outputs for every (shape, route, prompt) the checks compare. + + Three shapes: solo (one route per batch), mixed (every route interleaved + in one batch), wide (a batch larger than the cuda-graph capture, so the + eager verify path runs). + """ + out = {} + solo, accept = collect_solo(base_url, routes, args.max_new_tokens) + for route, texts in solo.items(): + for i, text in enumerate(texts): + out[("solo", route, i)] = text + + mixed_texts = [p for p in PROMPTS for _ in routes] + mixed_routes = [r for _ in PROMPTS for r in routes] + mixed, _ = generate(base_url, mixed_texts, mixed_routes, args.max_new_tokens) + for text, route, got in zip(mixed_texts, mixed_routes, mixed): + out[("mixed", route, PROMPTS.index(text))] = got + + wide = max(args.wide_batch, len(routes) * 2) + wide_routes = [routes[i % len(routes)] for i in range(wide)] + wide_texts = [PROMPTS[i % len(PROMPTS)] for i in range(wide)] + wide_out, _ = generate(base_url, wide_texts, wide_routes, args.max_new_tokens) + for text, route, got in zip(wide_texts, wide_routes, wide_out): + out[("wide", route, PROMPTS.index(text))] = got + + return out, accept, solo + + +def unstable_keys(ref_a: dict, ref_b: dict) -> set: + """(route, prompt) pairs the reference itself cannot reproduce. + + Greedy decoding is not bitwise reproducible here: batch composition + changes reduction order (MoE routing especially), and two runs of the + same server can diverge. Anything unstable *without* speculation is a + property of the model and backend, not of spec+LoRA, so the spec + comparison must exclude it or it reports noise as failure. + """ + unstable = set() + for (shape, route, index), text in ref_a.items(): + # run-to-run, identical shape + if ref_b.get((shape, route, index)) != text: + unstable.add((route, index)) + # shape-to-shape, same run + if ref_a.get(("solo", route, index)) != text: + unstable.add((route, index)) + return unstable + + +def is_reproducible(base_url: str, route, prompt: str, args, repeats: int = 3) -> bool: + """Whether one prompt yields the same greedy output across repeats.""" + outputs = { + generate(base_url, [prompt], [route], args.max_new_tokens)[0][0] + for _ in range(repeats) + } + return len(outputs) == 1 + + +def compare(ref_a, spec_a, spec_b, noisy, base_url: str, args) -> int: + """Count genuine spec-on vs spec-off differences. + + Skips pairs the reference could not reproduce, pairs where the spec side + disagrees with itself, and differences that land on an argmax tie (where + the token choice is arbitrary in the first place). + """ + log("comparing spec-on against spec-off on reproducible pairs only") + excluded = ties = failures = 0 + for key, want in ref_a.items(): + shape, route, index = key + # Both sides must be self-consistent before a difference between them + # means anything: the instability set is stochastic, so a single + # sample per side would report noise as a failure. + if (route, index) in noisy or spec_a.get(key) != spec_b.get(key): + excluded += 1 + continue + got = spec_a.get(key) + if got == want: + continue + + # A pair can look stable in two samples and still be stochastic, so + # re-sample this specific pair before calling it a defect. Verified + # necessary: a prompt that survived the blanket filter turned out to + # diverge spec-on vs spec-off with no adapters loaded at all, and to + # disagree with itself across repeats. + if not is_reproducible(base_url, route, PROMPTS[index], args): + ties += 1 + gap = min_logprob_gap(base_url, route, PROMPTS[index], args.max_new_tokens) + log( + f" UNSTABLE [{shape}] route={route} prompt#{index}: the spec " + f"server does not reproduce this prompt across repeats " + f"(min top-2 logprob gap {gap}); not attributable to spec" + ) + continue + + failures += 1 + log(f" MISMATCH [{shape}] route={route} prompt#{index}") + log(f" spec-off: {want!r}") + log(f" spec-on : {got!r}") + log( + f" ({excluded} skipped as baseline-nondeterministic, " + f"{ties} skipped as not reproducible on re-sampling)" + ) + return failures + + +def run(config: dict, args) -> int: + base_url = f"http://127.0.0.1:{args.port}" + adapters = resolve_adapters(config, args) + routes = [None] + [name for name, _ in adapters] + failures = 0 + process = None + + try: + # Phase 1 — reference: same weights and adapters, speculation off. + # Collected twice to measure how much this model/backend varies on + # its own, which sets the noise floor for the spec comparison. + process = launch(config, adapters, base_url, with_spec=False) + ref_a, _, solo_a = collect_shapes(base_url, routes, args) + ref_b, _, _ = collect_shapes(base_url, routes, args) + noisy = unstable_keys(ref_a, ref_b) + log( + f"baseline noise floor: {len(noisy)} of " + f"{len(routes) * len(PROMPTS)} (route, prompt) pairs are not " + "reproducible without speculation" + ) + + for route in routes[1:]: + if solo_a[route] == solo_a[None]: + failures += 1 + log( + f" FAIL adapter {route} output is identical to base on every prompt" + ) + kill_process_tree(process.pid) + process = None + + # Phase 2 — speculation on, same three shapes. + process = launch(config, adapters, base_url, with_spec=True) + spec_a, accept, _ = collect_shapes(base_url, routes, args) + spec_b, _, _ = collect_shapes(base_url, routes, args) + + failures += compare(ref_a, spec_a, spec_b, noisy, base_url, args) + finally: + if process is not None: + kill_process_tree(process.pid) + + log("=" * 68) + log("per-adapter accept length (tokens per verify step; 1.0 = no speedup)") + for route in routes: + value = accept.get(route) + label = "base" if route is None else route + log(f" {label:<24} {value if value is None else round(value, 3)}") + log("=" * 68) + log(f"RESULT: {'PASS' if failures == 0 else f'FAIL ({failures} mismatches)'}") + return 1 if failures else 0 + + +def main() -> int: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--config", help=f"one of: {', '.join(CONFIGS)}") + parser.add_argument("--list", action="store_true", help="list configs and exit") + parser.add_argument("--port", type=int, default=30000) + parser.add_argument("--max-new-tokens", type=int, default=32) + parser.add_argument( + "--wide-batch", + type=int, + default=16, + help="batch size for the eager-path check; keep it above --cuda-graph-max-bs", + ) + parser.add_argument( + "--lora-backend", help="override the config's LoRA kernel backend" + ) + parser.add_argument( + "--extra-arg", + action="append", + default=[], + dest="extra_args", + help="extra server flag, repeatable (e.g. --extra-arg=--enable-lora-overlap-loading)", + ) + parser.add_argument("--derived-dir", default="/tmp/sglang-derived-loras") + parser.add_argument("--keep-derived", action="store_true") + parser.add_argument( + "--mode", + choices=["correctness", "perf", "both"], + default="correctness", + help="correctness = the 4 output checks; perf = the 4-arm throughput " + "and accept-length comparison", + ) + parser.add_argument( + "--concurrency", + default="1,8,32", + help="comma-separated in-flight request counts to sweep (perf mode)", + ) + parser.add_argument("--num-prompts", type=int, default=64) + parser.add_argument("--input-len", type=int, default=512) + parser.add_argument("--output-len", type=int, default=256) + parser.add_argument("--seed", type=int, default=42) + args = parser.parse_args() + + if args.list or not args.config: + for name, config in CONFIGS.items(): + algo = next( + a.split("=")[1] for a in config["spec_args"] if "algorithm" in a + ) + n = len(config["adapters"]) + (1 if config.get("derive_from") else 0) + print( + f"{name:<24} tp={config['tp']} {algo:<7} adapters={n} {config['model']}" + ) + return 0 + + if args.config not in CONFIGS: + raise SystemExit(f"unknown config {args.config!r}; use --list") + started = time.time() + config = dict(CONFIGS[args.config]) + if args.lora_backend: + config["lora_backend"] = args.lora_backend + config["extra_args"] = args.extra_args + code = 0 + if args.mode in ("correctness", "both"): + code |= run(config, args) + if args.mode in ("perf", "both"): + code |= run_perf(config, args) + log(f"total wall time {time.time() - started:.0f}s") + return code + + +if __name__ == "__main__": + sys.exit(main()) diff --git a/test/registered/lora/test_lora_spec_decoding.py b/test/registered/lora/test_lora_spec_decoding.py new file mode 100644 index 000000000..050144fbd --- /dev/null +++ b/test/registered/lora/test_lora_spec_decoding.py @@ -0,0 +1,142 @@ +"""E2E test for multi-adapter LoRA + EAGLE-family speculative decoding. + +Adapters apply to the target model only; one shared draft runs unadapted. + +Deliberately asserts serving properties rather than exact output text. +Greedy decoding is not bitwise reproducible across batch shapes or server +restarts here -- reduction order changes flip a token and greedy amplifies +it -- so text equality between configurations is a flaky assertion, not a +correctness oracle. Losslessness is verified out of band by +test/manual/lora/run_spec_lora_matrix.py, which measures that noise floor +first. What this guards is what CI can hold stable: + +- the server starts at all with multiple adapters + speculation (it used to + crash loading the target's adapters into the draft model); +- adapters are actually applied during target-verify (adapter output differs + from base output); +- a mixed-adapter batch, and a batch wider than the cuda-graph capture, are + served without error (the eager verify path used to crash on + extend_seq_lens_cpu=None); +- speculation is really running (accept length above 1). +""" + +import unittest + +import requests + +from sglang.srt.utils import kill_process_tree +from sglang.test.ci.ci_register import register_cuda_ci +from sglang.test.test_utils import ( + DEFAULT_DRAFT_MODEL_EAGLE3, + DEFAULT_TARGET_MODEL_EAGLE3, + DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH, + DEFAULT_URL_FOR_TEST, + CustomTestCase, + popen_launch_server, +) + +register_cuda_ci(est_time=600, stage="base-b", runner_config="1-gpu-large") + +PROMPTS = [ + "What is the capital of France? Answer in one sentence.", + "List three primary colors.", + "Write a one-sentence story about a brave detective on Mars.", +] +# Ranks 8 and 64: a mixed-rank batch is what the per-request lora_ranks +# indexing has to get right, and uniform ranks would hide a mixup. +ADAPTERS = [ + ("fact", "algoprog/fact-generation-llama-3.1-8b-instruct-lora"), + ("guard", "nvidia/llama-3.1-nemoguard-8b-topic-control"), +] +SAMPLING = {"temperature": 0, "max_new_tokens": 32} + + +class TestEagle3MultiLoRA(CustomTestCase): + process = None + + @classmethod + def setUpClass(cls): + cls.base_url = DEFAULT_URL_FOR_TEST + cls.process = popen_launch_server( + DEFAULT_TARGET_MODEL_EAGLE3, + cls.base_url, + timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH, + other_args=[ + # Canonical EAGLE3 sglang config, as in + # test/registered/core/test_basic_sanity_eagle3.py: the draft + # checkpoint is fp16, and bf16 + flashinfer cutlass RMSNorm + # hits a dtype mismatch on the draft's input_layernorm. + "--dtype=float16", + "--attention-backend=triton", + "--speculative-algorithm=EAGLE3", + f"--speculative-draft-model-path={DEFAULT_DRAFT_MODEL_EAGLE3}", + "--speculative-num-steps=5", + "--speculative-eagle-topk=8", + "--speculative-num-draft-tokens=32", + "--enable-lora", + "--lora-backend=triton", + "--max-lora-rank=64", + # +1: the base model occupies a memory-pool slot too, and the + # batches below co-batch base requests with adapter requests. + f"--max-loras-per-batch={len(ADAPTERS) + 1}", + "--mem-fraction-static=0.7", + "--lora-paths", + ] + + [f"{name}={path}" for name, path in ADAPTERS], + ) + + @classmethod + def tearDownClass(cls): + if cls.process is not None: + kill_process_tree(cls.process.pid) + + def _generate(self, texts, lora_paths): + response = requests.post( + self.base_url + "/generate", + json={"text": texts, "lora_path": lora_paths, "sampling_params": SAMPLING}, + ) + self.assertEqual(response.status_code, 200, response.text) + results = response.json() + self.assertEqual(len(results), len(texts)) + for item in results: + self.assertTrue(item["text"].strip(), f"empty output in {results}") + return [item["text"] for item in results] + + def test_adapters_are_applied_under_speculation(self): + base = self._generate(PROMPTS, [None] * len(PROMPTS)) + for name, _ in ADAPTERS: + adapted = self._generate(PROMPTS, [name] * len(PROMPTS)) + self.assertNotEqual( + adapted, + base, + f"adapter {name} matched the base model on every prompt; LoRA " + "was likely not applied during target-verify", + ) + + def test_mixed_adapter_and_wide_batches_are_served(self): + routes = [None] + [name for name, _ in ADAPTERS] + self._generate( + [p for p in PROMPTS for _ in routes], + [r for _ in PROMPTS for r in routes], + ) + # Wider than the default cuda-graph capture, so target-verify falls + # back to the eager path. + wide = 24 + self._generate( + [PROMPTS[i % len(PROMPTS)] for i in range(wide)], + [routes[i % len(routes)] for i in range(wide)], + ) + + def test_speculation_is_active(self): + self._generate(PROMPTS, [ADAPTERS[0][0]] * len(PROMPTS)) + info = requests.get(self.base_url + "/get_server_info").json() + accept_length = info["internal_states"][0]["avg_spec_accept_length"] + self.assertGreater( + accept_length, + 1.0, + f"no drafts accepted with LoRA enabled: {accept_length}", + ) + + +if __name__ == "__main__": + unittest.main(verbosity=3) diff --git a/test/registered/unit/lora/test_lora_spec_verify_batch_info.py b/test/registered/unit/lora/test_lora_spec_verify_batch_info.py new file mode 100644 index 000000000..1aee79715 --- /dev/null +++ b/test/registered/unit/lora/test_lora_spec_verify_batch_info.py @@ -0,0 +1,109 @@ +"""Unit tests for LoRA batch-info preparation under TARGET_VERIFY. + +TARGET_VERIFY reports is_extend() True but ForwardBatch.init_new leaves +extend_seq_lens / extend_seq_lens_cpu as None (verify is routed through the +decode-style positions branch), while every request carries a uniform +spec_info.draft_token_num-token segment. Regression guards for the paths that +assumed extend fields exist on every is_extend() mode: + +- triton eager path: ``max(None)`` TypeError, and 1-token/req segments for a + draft_token_num-token layout (silent adapter mis-segmentation). +- MoE _add_moe_lora_info: ``sum(None)`` TypeError. +- static cuda-graph path: seg_lens are pre-filled at the captured width; a + verify batch of another width must fail loudly, not mis-segment. +""" + +import unittest +from types import SimpleNamespace +from unittest.mock import patch + +import torch + +from sglang.srt.lora.backend.triton_backend import TritonLoRABackend +from sglang.srt.model_executor.forward_batch_info import ForwardMode +from sglang.test.ci.ci_register import register_cuda_ci +from sglang.test.test_utils import CustomTestCase + +register_cuda_ci(est_time=30, stage="base-b", runner_config="1-gpu-small") + + +def _verify_batch(bs: int, draft_token_num: int) -> SimpleNamespace: + return SimpleNamespace( + forward_mode=ForwardMode.TARGET_VERIFY, + batch_size=bs, + spec_info=SimpleNamespace(draft_token_num=draft_token_num), + extend_seq_lens=None, + extend_seq_lens_cpu=None, + ) + + +class TestLoRASpecVerifyBatchInfo(CustomTestCase): + def _backend(self, max_loras_per_batch: int = 2) -> TritonLoRABackend: + return TritonLoRABackend( + max_loras_per_batch=max_loras_per_batch, device=torch.device("cuda") + ) + + def _prepare(self, backend, forward_batch, use_cuda_graph: bool): + backend.prepare_lora_batch( + forward_batch, + weight_indices=[0, 1], + lora_ranks=[8, 8], + scalings=[1.0, 1.0], + use_cuda_graph=use_cuda_graph, + ) + return backend.batch_info + + def test_eager_target_verify_builds_uniform_draft_width_segments(self): + backend = self._backend() + batch_info = self._prepare( + backend, _verify_batch(bs=2, draft_token_num=4), use_cuda_graph=False + ) + self.assertEqual(batch_info.max_len, 4) + self.assertEqual(batch_info.seg_lens[:2].tolist(), [4, 4]) + self.assertEqual(batch_info.seg_indptr[:3].tolist(), [0, 4, 8]) + self.assertEqual(batch_info.num_segments, 2) + + def test_graph_path_serves_the_captured_width_and_rejects_others(self): + """prepare_lora_batch predicts graph use before can_run_graph decides, + so a mismatched width used to silently apply the captured segment + layout to a differently-shaped batch (wrong adapter on wrong rows). + The matching case also pins that bs is rebound per batch, which is + what makes a mixed-adapter batch index the right slots.""" + backend = self._backend() + backend.init_cuda_graph_batch_info(max_bs_in_cuda_graph=4, num_tokens_per_req=4) + + batch_info = self._prepare( + backend, _verify_batch(bs=2, draft_token_num=4), use_cuda_graph=True + ) + self.assertIs(batch_info, backend.cuda_graph_batch_info) + self.assertEqual(batch_info.seg_lens.tolist(), [4, 4, 4, 4]) + self.assertEqual(batch_info.bs, 2) + + with self.assertRaisesRegex(AssertionError, "width"): + self._prepare( + backend, _verify_batch(bs=2, draft_token_num=8), use_cuda_graph=True + ) + + def test_moe_lora_info_uses_draft_width_token_counts_for_verify(self): + backend = self._backend() + backend.is_moe_lora = True + captured = {} + + def _capture(num_tokens, seg_indptr, lora_ranks, req_to_lora, *args, **kwargs): + captured["num_tokens"] = num_tokens + captured["max_len"] = kwargs["max_len"] + return None, None + + with patch( + "sglang.srt.lora.backend.base_backend._compute_moe_lora_info", + side_effect=_capture, + ): + self._prepare( + backend, _verify_batch(bs=2, draft_token_num=4), use_cuda_graph=False + ) + self.assertEqual(captured["num_tokens"], 8) + self.assertEqual(captured["max_len"], 4) + + +if __name__ == "__main__": + unittest.main() diff --git a/test/registered/unit/model_executor/test_draft_runner_skips_lora.py b/test/registered/unit/model_executor/test_draft_runner_skips_lora.py new file mode 100644 index 000000000..59f7f52a2 --- /dev/null +++ b/test/registered/unit/model_executor/test_draft_runner_skips_lora.py @@ -0,0 +1,51 @@ +"""The draft runner must not build a LoRA manager. + +Adapters apply to the target model only. Every worker is handed the same +published ServerArgs, so `enable_lora` is True for the draft too -- the +decision is the runner's own, keyed on is_draft_worker. Without it the draft +tries to load the target's adapters into the draft model, whose layer count +differs, and startup fails inside LoRAAdapter weight loading. + +The LoRA paths downstream then key on `lora_manager is not None` rather than +the config, so a draft runner skips them by construction. +""" + +import unittest +from unittest.mock import patch + +from sglang.srt.model_executor.model_runner import ModelRunner +from sglang.test.ci.ci_register import register_cpu_ci +from sglang.test.test_utils import CustomTestCase + +register_cpu_ci(est_time=5, suite="base-a-test-cpu") + + +class TestDraftRunnerSkipsLoRA(CustomTestCase): + def _init_lora_called(self, *, is_draft_worker: bool, enable_lora: bool) -> bool: + runner = ModelRunner.__new__(ModelRunner) + runner.is_draft_worker = is_draft_worker + runner.lora_manager = None + with patch.object(ModelRunner, "init_lora_manager") as init_lora: + with patch("sglang.srt.model_executor.model_runner.get_lora") as get_lora: + get_lora.return_value.enable_lora = enable_lora + runner.maybe_init_lora_manager() + return init_lora.called + + def test_only_the_target_runner_builds_a_lora_manager(self): + cases = [ + (False, True, True), # target + LoRA -> builds one + (True, True, False), # draft + LoRA -> must not + (False, False, False), # LoRA off -> nobody builds one + ] + for is_draft_worker, enable_lora, expected in cases: + with self.subTest(draft=is_draft_worker, enable_lora=enable_lora): + self.assertEqual( + self._init_lora_called( + is_draft_worker=is_draft_worker, enable_lora=enable_lora + ), + expected, + ) + + +if __name__ == "__main__": + unittest.main() diff --git a/test/registered/unit/model_executor/test_prefill_cuda_graph_runner.py b/test/registered/unit/model_executor/test_prefill_cuda_graph_runner.py index a00bb5c9b..526643a5e 100644 --- a/test/registered/unit/model_executor/test_prefill_cuda_graph_runner.py +++ b/test/registered/unit/model_executor/test_prefill_cuda_graph_runner.py @@ -69,9 +69,11 @@ class TestPrefillCudaGraphRunnerChunkedPrefix(CustomTestCase): device="cuda", gpu_id=0, is_draft_worker=False, + # A real ModelRunner always has this attribute; the prefill gate + # reads it rather than the process-wide LoRA config. + lora_manager=None, spec_algorithm=SimpleNamespace(is_eagle=lambda: False), server_args=SimpleNamespace( - enable_lora=False, cuda_graph_config=SimpleNamespace( prefill=SimpleNamespace(bs=[1], backend=Backend.BREAKABLE) ),