Small cleanups related to LoRA weight loading (#13474)

This commit is contained in:
Glen Liu
2025-11-18 09:14:34 -08:00
committed by GitHub
parent 6e9b154981
commit d79e12941c
3 changed files with 4 additions and 20 deletions
+3 -8
View File
@@ -19,13 +19,13 @@
# https://github.com/vllm-project/vllm/blob/4abf6336ec65c270343eb895e7b18786e9274176/vllm/lora/layers.py # https://github.com/vllm-project/vllm/blob/4abf6336ec65c270343eb895e7b18786e9274176/vllm/lora/layers.py
import logging import logging
import re
from typing import Dict, List from typing import Dict, List
import torch import torch
from torch import nn from torch import nn
from sglang.srt.configs.load_config import LoadConfig from sglang.srt.configs.load_config import LoadConfig
from sglang.srt.layers.utils import get_layer_id
from sglang.srt.lora.backend.base_backend import BaseLoRABackend from sglang.srt.lora.backend.base_backend import BaseLoRABackend
from sglang.srt.lora.backend.lora_registry import LORA_SUPPORTED_BACKENDS from sglang.srt.lora.backend.lora_registry import LORA_SUPPORTED_BACKENDS
from sglang.srt.lora.lora_config import LoRAConfig from sglang.srt.lora.lora_config import LoRAConfig
@@ -71,8 +71,6 @@ class LoRAAdapter(nn.Module):
] ]
) )
self.weights: Dict[str, torch.Tensor] = {}
# initialize the LoRA weights to cpu # initialize the LoRA weights to cpu
def initialize_weights(self): def initialize_weights(self):
model_path = self.config.path model_path = self.config.path
@@ -83,12 +81,9 @@ class LoRAAdapter(nn.Module):
model_path, revision=revision, fall_back_to_pt=True model_path, revision=revision, fall_back_to_pt=True
) )
): ):
match = re.search(r"layers\.(\d+)\.", name) layer_id = get_layer_id(name)
if match is not None: if layer_id is not None:
layer_id = int(match.group(1))
self.layers[layer_id].weights[name] = loaded_weight.cpu() self.layers[layer_id].weights[name] = loaded_weight.cpu()
else:
self.weights[name] = loaded_weight.cpu()
# normalize kv_proj and gate_up_proj # normalize kv_proj and gate_up_proj
for layer in self.layers: for layer in self.layers:
+1 -1
View File
@@ -21,6 +21,7 @@ from typing import Dict, Iterable, List, Optional
import torch import torch
from sglang.srt.configs.load_config import LoadConfig from sglang.srt.configs.load_config import LoadConfig
from sglang.srt.layers.utils import get_layer_id
from sglang.srt.lora.backend.base_backend import BaseLoRABackend from sglang.srt.lora.backend.base_backend import BaseLoRABackend
from sglang.srt.lora.backend.lora_registry import get_backend_from_name from sglang.srt.lora.backend.lora_registry import get_backend_from_name
from sglang.srt.lora.layers import BaseLayerWithLoRA, get_lora_layer from sglang.srt.lora.layers import BaseLayerWithLoRA, get_lora_layer
@@ -30,7 +31,6 @@ from sglang.srt.lora.lora_registry import LoRARef
from sglang.srt.lora.mem_pool import LoRAMemoryPool from sglang.srt.lora.mem_pool import LoRAMemoryPool
from sglang.srt.lora.utils import ( from sglang.srt.lora.utils import (
LoRAType, LoRAType,
get_layer_id,
get_normalized_target_modules, get_normalized_target_modules,
get_target_module_name, get_target_module_name,
) )
-11
View File
@@ -1,4 +1,3 @@
import re
from dataclasses import dataclass from dataclasses import dataclass
from enum import Enum from enum import Enum
from typing import Iterable, Optional, Set, Tuple from typing import Iterable, Optional, Set, Tuple
@@ -46,16 +45,6 @@ class LoRAType(Enum):
LORA_B = 1 LORA_B = 1
def get_layer_id(name: str) -> int:
"""
Extract integer id of layer from its name in string.
"""
match = re.search(r"layers\.(\d+)\.", name)
if match is None:
return None
return int(match.group(1))
def get_hidden_dim( def get_hidden_dim(
module_name: str, config: AutoConfig, base_model: torch.nn.Module, layer_idx: int module_name: str, config: AutoConfig, base_model: torch.nn.Module, layer_idx: int
) -> Tuple[int]: ) -> Tuple[int]: