fix: add ModelScope cache lookup and speculative path support (#20098)
This commit is contained in:
@@ -919,18 +919,7 @@ class ServerArgs:
|
|||||||
|
|
||||||
# Handle ModelScope model downloads
|
# Handle ModelScope model downloads
|
||||||
if get_bool_env_var("SGLANG_USE_MODELSCOPE"):
|
if get_bool_env_var("SGLANG_USE_MODELSCOPE"):
|
||||||
if not os.path.exists(self.model_path):
|
self._handle_modelscope_paths()
|
||||||
from modelscope import snapshot_download
|
|
||||||
|
|
||||||
self.model_path = snapshot_download(
|
|
||||||
self.model_path, cache_dir=self.download_dir, revision=self.revision
|
|
||||||
)
|
|
||||||
self.tokenizer_path = snapshot_download(
|
|
||||||
self.tokenizer_path,
|
|
||||||
cache_dir=self.download_dir,
|
|
||||||
revision=self.revision,
|
|
||||||
ignore_patterns=["*.bin", "*.safetensors"],
|
|
||||||
)
|
|
||||||
|
|
||||||
# Mamba scheduler strategy
|
# Mamba scheduler strategy
|
||||||
if self.mamba_scheduler_strategy == "auto":
|
if self.mamba_scheduler_strategy == "auto":
|
||||||
@@ -946,6 +935,70 @@ class ServerArgs:
|
|||||||
elif self.speculative_draft_model_quantization == "unquant":
|
elif self.speculative_draft_model_quantization == "unquant":
|
||||||
self.speculative_draft_model_quantization = None
|
self.speculative_draft_model_quantization = None
|
||||||
|
|
||||||
|
def _handle_modelscope_paths(self):
|
||||||
|
"""Resolve model / tokenizer / speculative-draft paths from the local
|
||||||
|
ModelScope cache when possible, falling back to ``snapshot_download``
|
||||||
|
for any path that is not already present on disk.
|
||||||
|
|
||||||
|
Note: ``speculative_token_map`` is intentionally NOT handled here
|
||||||
|
because its value uses ``repo_id/filename`` semantics rather than a
|
||||||
|
plain repo ID. That resolution lives in
|
||||||
|
:func:`sglang.srt.speculative.spec_utils.load_token_map`.
|
||||||
|
"""
|
||||||
|
|
||||||
|
ms_root = None
|
||||||
|
ms_snapshot_download = None
|
||||||
|
|
||||||
|
def _resolve_or_download(
|
||||||
|
path: Optional[str],
|
||||||
|
ignore_patterns: Optional[list] = None,
|
||||||
|
revision: Optional[str] = None,
|
||||||
|
) -> Optional[str]:
|
||||||
|
nonlocal ms_root, ms_snapshot_download
|
||||||
|
if path is None:
|
||||||
|
return None
|
||||||
|
if not path or os.path.exists(path):
|
||||||
|
return path
|
||||||
|
|
||||||
|
if ms_snapshot_download is None:
|
||||||
|
from modelscope.hub.snapshot_download import (
|
||||||
|
snapshot_download as _ms_snapshot_download,
|
||||||
|
)
|
||||||
|
from modelscope.utils.file_utils import get_model_cache_root
|
||||||
|
|
||||||
|
ms_snapshot_download = _ms_snapshot_download
|
||||||
|
ms_root = get_model_cache_root()
|
||||||
|
|
||||||
|
# Check ModelScope default cache
|
||||||
|
cached = os.path.join(ms_root, path)
|
||||||
|
if os.path.exists(cached):
|
||||||
|
return cached
|
||||||
|
# Check user-specified download dir
|
||||||
|
if self.download_dir:
|
||||||
|
alt = os.path.join(self.download_dir, path)
|
||||||
|
if os.path.exists(alt):
|
||||||
|
return alt
|
||||||
|
|
||||||
|
# Cache miss — download from ModelScope hub
|
||||||
|
return ms_snapshot_download(
|
||||||
|
path,
|
||||||
|
cache_dir=self.download_dir,
|
||||||
|
revision=revision,
|
||||||
|
**({"ignore_patterns": ignore_patterns} if ignore_patterns else {}),
|
||||||
|
)
|
||||||
|
|
||||||
|
self.model_path = _resolve_or_download(self.model_path, revision=self.revision)
|
||||||
|
self.tokenizer_path = _resolve_or_download(
|
||||||
|
self.tokenizer_path,
|
||||||
|
ignore_patterns=["*.bin", "*.safetensors"],
|
||||||
|
revision=self.revision,
|
||||||
|
)
|
||||||
|
if self.speculative_draft_model_path:
|
||||||
|
self.speculative_draft_model_path = _resolve_or_download(
|
||||||
|
self.speculative_draft_model_path,
|
||||||
|
revision=self.speculative_draft_model_revision or "main",
|
||||||
|
)
|
||||||
|
|
||||||
def _handle_hpu_backends(self):
|
def _handle_hpu_backends(self):
|
||||||
if self.device == "hpu":
|
if self.device == "hpu":
|
||||||
self.attention_backend = "torch_native"
|
self.attention_backend = "torch_native"
|
||||||
|
|||||||
@@ -688,11 +688,30 @@ def generate_token_bitmask(
|
|||||||
|
|
||||||
def load_token_map(token_map_path: str) -> List[int]:
|
def load_token_map(token_map_path: str) -> List[int]:
|
||||||
if not os.path.exists(token_map_path):
|
if not os.path.exists(token_map_path):
|
||||||
cache_dir = snapshot_download(
|
repo_id = os.path.dirname(token_map_path)
|
||||||
os.path.dirname(token_map_path),
|
file_name = os.path.basename(token_map_path)
|
||||||
|
|
||||||
|
cache_dir = None
|
||||||
|
if envs.SGLANG_USE_MODELSCOPE.get():
|
||||||
|
from modelscope.utils.file_utils import get_model_cache_root
|
||||||
|
|
||||||
|
cached_repo_path = os.path.join(get_model_cache_root(), repo_id)
|
||||||
|
if os.path.exists(cached_repo_path):
|
||||||
|
cache_dir = cached_repo_path
|
||||||
|
|
||||||
|
if cache_dir is None:
|
||||||
|
if envs.SGLANG_USE_MODELSCOPE.get():
|
||||||
|
from modelscope.hub.snapshot_download import (
|
||||||
|
snapshot_download as download_func,
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
download_func = snapshot_download
|
||||||
|
cache_dir = download_func(
|
||||||
|
repo_id,
|
||||||
ignore_patterns=["*.bin", "*.safetensors"],
|
ignore_patterns=["*.bin", "*.safetensors"],
|
||||||
)
|
)
|
||||||
token_map_path = os.path.join(cache_dir, os.path.basename(token_map_path))
|
|
||||||
|
token_map_path = os.path.join(cache_dir, file_name)
|
||||||
hot_token_id = torch.load(token_map_path, weights_only=True)
|
hot_token_id = torch.load(token_map_path, weights_only=True)
|
||||||
return torch.tensor(hot_token_id, dtype=torch.int64)
|
return torch.tensor(hot_token_id, dtype=torch.int64)
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user