Move LoRA cuda-graph buffers and logging into LoRAManager (#31151)

This commit is contained in:
fzyzcjy
2026-07-14 15:56:18 +08:00
committed by GitHub
parent e20c346541
commit 205a2f2de4
2 changed files with 84 additions and 70 deletions
+73 -2
View File
@@ -46,7 +46,7 @@ from sglang.srt.lora.utils import (
from sglang.srt.managers.io_struct import LoRAUpdateOutput from sglang.srt.managers.io_struct import LoRAUpdateOutput
from sglang.srt.model_executor.forward_batch_info import ForwardBatch from sglang.srt.model_executor.forward_batch_info import ForwardBatch
from sglang.srt.server_args import ServerArgs from sglang.srt.server_args import ServerArgs
from sglang.srt.utils import replace_submodule from sglang.srt.utils import get_available_gpu_memory, replace_submodule
from sglang.srt.utils.hf_transformers_utils import AutoConfig from sglang.srt.utils.hf_transformers_utils import AutoConfig
_SGLANG_EXPERIMENTAL_LORA_OPTI = envs.SGLANG_EXPERIMENTAL_LORA_OPTI.get() _SGLANG_EXPERIMENTAL_LORA_OPTI = envs.SGLANG_EXPERIMENTAL_LORA_OPTI.get()
@@ -164,6 +164,18 @@ class LoRAManager:
) )
def load_lora_adapter(self, lora_ref: LoRARef) -> LoRAUpdateOutput: def load_lora_adapter(self, lora_ref: LoRARef) -> LoRAUpdateOutput:
logger.info(
f"LoRA adapter loading starts: {lora_ref}. "
f"avail mem={get_available_gpu_memory(self.device.type, self.device.index):.2f} GB"
)
result = self._load_lora_adapter(lora_ref)
logger.info(
f"LoRA adapter loading completes: {lora_ref}. "
f"avail mem={get_available_gpu_memory(self.device.type, self.device.index):.2f} GB"
)
return result
def _load_lora_adapter(self, lora_ref: LoRARef) -> LoRAUpdateOutput:
""" """
Load a single LoRA adapter from the specified path. Load a single LoRA adapter from the specified path.
@@ -246,6 +258,18 @@ class LoRAManager:
) )
def unload_lora_adapter(self, lora_ref: LoRARef) -> LoRAUpdateOutput: def unload_lora_adapter(self, lora_ref: LoRARef) -> LoRAUpdateOutput:
logger.info(
f"LoRA adapter unloading starts: {lora_ref}. "
f"avail mem={get_available_gpu_memory(self.device.type, self.device.index):.2f} GB"
)
result = self._unload_lora_adapter(lora_ref)
logger.info(
f"LoRA adapter unloading completes: {lora_ref}. "
f"avail mem={get_available_gpu_memory(self.device.type, self.device.index):.2f} GB"
)
return result
def _unload_lora_adapter(self, lora_ref: LoRARef) -> LoRAUpdateOutput:
""" """
Unload LoRA adapters by their names. This will remove the adapters from the memory pool and Unload LoRA adapters by their names. This will remove the adapters from the memory pool and
delete the corresponding LoRA modules. delete the corresponding LoRA modules.
@@ -484,7 +508,7 @@ class LoRAManager:
if lora_paths: if lora_paths:
for lora_ref in lora_paths: for lora_ref in lora_paths:
result = self.load_lora_adapter(lora_ref) result = self._load_lora_adapter(lora_ref)
if not result.success: if not result.success:
raise RuntimeError( raise RuntimeError(
f"Failed to load LoRA adapter {lora_ref.lora_name}: {result.error_message}" f"Failed to load LoRA adapter {lora_ref.lora_name}: {result.error_message}"
@@ -686,6 +710,20 @@ class LoRAManager:
tensors: Dict[str, torch.Tensor], tensors: Dict[str, torch.Tensor],
config_dict: Dict, config_dict: Dict,
added_tokens_config: Optional[Dict] = None, added_tokens_config: Optional[Dict] = None,
) -> LoRAUpdateOutput:
logger.info(f"LoRA adapter loading from tensors starts: {lora_ref}.")
result = self._load_lora_adapter_from_tensors(
lora_ref, tensors, config_dict, added_tokens_config
)
logger.info(f"LoRA adapter loading from tensors completes: {lora_ref}.")
return result
def _load_lora_adapter_from_tensors(
self,
lora_ref: LoRARef,
tensors: Dict[str, torch.Tensor],
config_dict: Dict,
added_tokens_config: Optional[Dict] = None,
) -> LoRAUpdateOutput: ) -> LoRAUpdateOutput:
""" """
Load a single LoRA adapter from tensors and config dict. Load a single LoRA adapter from tensors and config dict.
@@ -856,3 +894,36 @@ class LoRAManager:
lora_module.experts_shared_outer_loras = self.experts_shared_outer_loras lora_module.experts_shared_outer_loras = self.experts_shared_outer_loras
lora_module.lora_use_virtual_experts = self.lora_use_virtual_experts lora_module.lora_use_virtual_experts = self.lora_use_virtual_experts
self.lora_modules[layer_id][module_name] = lora_module self.lora_modules[layer_id][module_name] = lora_module
def init_lora_cuda_graph_moe_buffers(
*,
server_args: ServerArgs,
model: torch.nn.Module,
lora_manager: LoRAManager,
dtype: torch.dtype,
):
"""Phase 1 of LoRA CUDA graph init: pre-allocate MoE intermediate buffers.
Must be called before init_memory_pool() so that memory profiling
sees the reduced available memory and sizes KV cache correctly.
All MoE LoRA layers share one set of buffers (managed by the
lora_backend) since they execute sequentially during forward.
Phase 2 (dense LoRA batch metadata) is handled later in
CudaGraphRunner.__init__() via lora_manager.init_cuda_graph_batch_info(),
because it needs capture-time parameters (max_bs, num_tokens_per_req)
that are only available at that stage.
"""
from sglang.srt.lora.layers import FusedMoEWithLoRA
max_bs = server_args.cuda_graph_config.decode.max_bs
max_loras = server_args.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)
logger.info(
f"Pre-allocated shared MoE LoRA CUDA graph buffers "
f"(max_bs={max_bs}, max_loras={max_loras})"
)
break
@@ -114,7 +114,7 @@ from sglang.srt.layers.moe.topk import TopK
from sglang.srt.layers.sampler import create_sampler from sglang.srt.layers.sampler import create_sampler
from sglang.srt.layers.torchao_utils import apply_torchao_config_to_model from sglang.srt.layers.torchao_utils import apply_torchao_config_to_model
from sglang.srt.layers.utils.cp_utils import is_mla_prefill_cp_enabled from sglang.srt.layers.utils.cp_utils import is_mla_prefill_cp_enabled
from sglang.srt.lora.lora_manager import LoRAManager from sglang.srt.lora.lora_manager import LoRAManager, init_lora_cuda_graph_moe_buffers
from sglang.srt.lora.lora_registry import LoRARef from sglang.srt.lora.lora_registry import LoRARef
from sglang.srt.managers.schedule_batch import sanity_check_mm_pad_shift_value from sglang.srt.managers.schedule_batch import sanity_check_mm_pad_shift_value
from sglang.srt.mem_cache import kv_cache_dtype from sglang.srt.mem_cache import kv_cache_dtype
@@ -734,13 +734,6 @@ class ModelRunner(ModelRunnerKVCacheMixin):
# Init lora # Init lora
if server_args.enable_lora: if server_args.enable_lora:
self.init_lora_manager() self.init_lora_manager()
if not cuda_graph_fully_disabled():
# Phase 1 of LoRA CUDA graph init: pre-allocate large MoE
# intermediate buffers before init_memory_pool() so memory
# profiling accounts for them. The buffers are reused by
# any captured graph (decode today; widen here so any
# future prefill capture path also picks them up).
self._init_lora_cuda_graph_moe_buffers()
# Enable batch invariant mode # Enable batch invariant mode
if server_args.enable_deterministic_inference: if server_args.enable_deterministic_inference:
@@ -1692,78 +1685,28 @@ class ModelRunner(ModelRunnerKVCacheMixin):
target_modules=self.server_args.lora_target_modules, target_modules=self.server_args.lora_target_modules,
lora_paths=self.server_args.lora_paths, lora_paths=self.server_args.lora_paths,
) )
if not cuda_graph_fully_disabled():
def _init_lora_cuda_graph_moe_buffers(self): init_lora_cuda_graph_moe_buffers(
"""Phase 1 of LoRA CUDA graph init: pre-allocate MoE intermediate buffers. server_args=self.server_args,
model=self.model,
Must be called before init_memory_pool() so that memory profiling lora_manager=self.lora_manager,
sees the reduced available memory and sizes KV cache correctly. dtype=self.dtype,
All MoE LoRA layers share one set of buffers (managed by the
lora_backend) since they execute sequentially during forward.
Phase 2 (dense LoRA batch metadata) is handled later in
CudaGraphRunner.__init__() via lora_manager.init_cuda_graph_batch_info(),
because it needs capture-time parameters (max_bs, num_tokens_per_req)
that are only available at that stage.
"""
from sglang.srt.lora.layers import FusedMoEWithLoRA
max_bs = self.server_args.cuda_graph_config.decode.max_bs
max_loras = self.server_args.max_loras_per_batch
for module in self.model.modules():
if isinstance(module, FusedMoEWithLoRA):
self.lora_manager.init_cuda_graph_moe_buffers(
max_bs, max_loras, self.dtype, module
) )
logger.info(
f"Pre-allocated shared MoE LoRA CUDA graph buffers "
f"(max_bs={max_bs}, max_loras={max_loras})"
)
break
def load_lora_adapter(self, lora_ref: LoRARef): def load_lora_adapter(self, lora_ref: LoRARef):
"""Load a new lora adapter from disk or huggingface.""" """Load a new lora adapter from disk or huggingface."""
return self.lora_manager.load_lora_adapter(lora_ref)
logger.info(
f"LoRA adapter loading starts: {lora_ref}. "
f"avail mem={get_available_gpu_memory(self.device, self.gpu_id):.2f} GB"
)
result = self.lora_manager.load_lora_adapter(lora_ref)
logger.info(
f"LoRA adapter loading completes: {lora_ref}. "
f"avail mem={get_available_gpu_memory(self.device, self.gpu_id):.2f} GB"
)
return result
def load_lora_adapter_from_tensors( def load_lora_adapter_from_tensors(
self, lora_ref: LoRARef, tensors, config_dict, added_tokens_config=None self, lora_ref: LoRARef, tensors, config_dict, added_tokens_config=None
): ):
logger.info(f"LoRA adapter loading from tensors starts: {lora_ref}.") return self.lora_manager.load_lora_adapter_from_tensors(
result = self.lora_manager.load_lora_adapter_from_tensors(
lora_ref, tensors, config_dict, added_tokens_config lora_ref, tensors, config_dict, added_tokens_config
) )
logger.info(f"LoRA adapter loading from tensors completes: {lora_ref}.")
return result
def unload_lora_adapter(self, lora_ref: LoRARef): def unload_lora_adapter(self, lora_ref: LoRARef):
"""Unload a lora adapter that was previously loaded during initialization or dynamic loading.""" """Unload a lora adapter that was previously loaded during initialization or dynamic loading."""
return self.lora_manager.unload_lora_adapter(lora_ref)
logger.info(
f"LoRA adapter unloading starts: {lora_ref}. "
f"avail mem={get_available_gpu_memory(self.device, self.gpu_id):.2f} GB"
)
result = self.lora_manager.unload_lora_adapter(lora_ref)
logger.info(
f"LoRA adapter unloading completes: {lora_ref}. "
f"avail mem={get_available_gpu_memory(self.device, self.gpu_id):.2f} GB"
)
return result
@property @property
def effective_max_total_num_tokens(self): def effective_max_total_num_tokens(self):