[loader] enable private loader (#14620)
This commit is contained in:
@@ -29,6 +29,7 @@ class LoadFormat(str, enum.Enum):
|
|||||||
REMOTE_INSTANCE = "remote_instance"
|
REMOTE_INSTANCE = "remote_instance"
|
||||||
RDMA = "rdma"
|
RDMA = "rdma"
|
||||||
LOCAL_CACHED = "local_cached"
|
LOCAL_CACHED = "local_cached"
|
||||||
|
PRIVATE = "private"
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
|
|||||||
@@ -953,7 +953,7 @@ class ModelRunner:
|
|||||||
return iter
|
return iter
|
||||||
|
|
||||||
def model_load_weights(model, iter):
|
def model_load_weights(model, iter):
|
||||||
DefaultModelLoader.load_weights_and_postprocess(model, iter, target_device)
|
loader.load_weights_and_postprocess(model, iter, target_device)
|
||||||
return model
|
return model
|
||||||
|
|
||||||
with set_default_torch_dtype(self.model_config.dtype):
|
with set_default_torch_dtype(self.model_config.dtype):
|
||||||
|
|||||||
@@ -2585,4 +2585,13 @@ def get_model_loader(
|
|||||||
if load_config.load_format == LoadFormat.REMOTE_INSTANCE:
|
if load_config.load_format == LoadFormat.REMOTE_INSTANCE:
|
||||||
return RemoteInstanceModelLoader(load_config)
|
return RemoteInstanceModelLoader(load_config)
|
||||||
|
|
||||||
|
if load_config.load_format == LoadFormat.PRIVATE:
|
||||||
|
import importlib
|
||||||
|
|
||||||
|
try:
|
||||||
|
module = importlib.import_module("sglang.private.private_model_loader")
|
||||||
|
return module.PrivateModelLoader(load_config)
|
||||||
|
except ImportError:
|
||||||
|
raise ValueError("Failed to import sglang.private.private_model_loader")
|
||||||
|
|
||||||
return DefaultModelLoader(load_config)
|
return DefaultModelLoader(load_config)
|
||||||
|
|||||||
@@ -696,7 +696,7 @@ class Qwen2_5_VLForConditionalGeneration(nn.Module):
|
|||||||
if name in params_dict.keys():
|
if name in params_dict.keys():
|
||||||
param = params_dict[name]
|
param = params_dict[name]
|
||||||
else:
|
else:
|
||||||
continue
|
raise ValueError(f"Weight {name} not found in params_dict")
|
||||||
except KeyError:
|
except KeyError:
|
||||||
print(params_dict.keys())
|
print(params_dict.keys())
|
||||||
raise
|
raise
|
||||||
|
|||||||
@@ -82,6 +82,7 @@ LOAD_FORMAT_CHOICES = [
|
|||||||
"flash_rl",
|
"flash_rl",
|
||||||
"remote",
|
"remote",
|
||||||
"remote_instance",
|
"remote_instance",
|
||||||
|
"private",
|
||||||
]
|
]
|
||||||
|
|
||||||
QUANTIZATION_CHOICES = [
|
QUANTIZATION_CHOICES = [
|
||||||
|
|||||||
Reference in New Issue
Block a user