[loader] enable private loader (#14620)
This commit is contained in:
@@ -29,6 +29,7 @@ class LoadFormat(str, enum.Enum):
|
||||
REMOTE_INSTANCE = "remote_instance"
|
||||
RDMA = "rdma"
|
||||
LOCAL_CACHED = "local_cached"
|
||||
PRIVATE = "private"
|
||||
|
||||
|
||||
@dataclass
|
||||
|
||||
@@ -953,7 +953,7 @@ class ModelRunner:
|
||||
return 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
|
||||
|
||||
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:
|
||||
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)
|
||||
|
||||
@@ -696,7 +696,7 @@ class Qwen2_5_VLForConditionalGeneration(nn.Module):
|
||||
if name in params_dict.keys():
|
||||
param = params_dict[name]
|
||||
else:
|
||||
continue
|
||||
raise ValueError(f"Weight {name} not found in params_dict")
|
||||
except KeyError:
|
||||
print(params_dict.keys())
|
||||
raise
|
||||
|
||||
@@ -82,6 +82,7 @@ LOAD_FORMAT_CHOICES = [
|
||||
"flash_rl",
|
||||
"remote",
|
||||
"remote_instance",
|
||||
"private",
|
||||
]
|
||||
|
||||
QUANTIZATION_CHOICES = [
|
||||
|
||||
Reference in New Issue
Block a user