Move LoRA cuda-graph buffers and logging into LoRAManager (#31151)
This commit is contained in:
@@ -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):
|
||||||
|
|||||||
Reference in New Issue
Block a user