diff --git a/python/sglang/multimodal_gen/runtime/loader/component_loaders/transformer_loader.py b/python/sglang/multimodal_gen/runtime/loader/component_loaders/transformer_loader.py index 4b4050c7d..723cf2899 100644 --- a/python/sglang/multimodal_gen/runtime/loader/component_loaders/transformer_loader.py +++ b/python/sglang/multimodal_gen/runtime/loader/component_loaders/transformer_loader.py @@ -25,9 +25,9 @@ from sglang.multimodal_gen.runtime.loader.minimax_h3_weights import ( ) from sglang.multimodal_gen.runtime.loader.transformer_load_utils import ( TransformerQuantLoadSpec, + resolve_transformer_checkpoint_files, resolve_transformer_gguf_to_load, resolve_transformer_quant_load_spec, - resolve_transformer_safetensors_to_load, ) from sglang.multimodal_gen.runtime.loader.utils import _normalize_component_type from sglang.multimodal_gen.runtime.loader.weight_load_plan import WeightLoadPlan @@ -214,10 +214,13 @@ class TransformerLoader(ComponentLoader): # A GGUF file holds the whole transformer; the remaining components # still load from the base model path. safetensors_list = [] + transformer_override_config_path = None else: - safetensors_list = resolve_transformer_safetensors_to_load( + checkpoint_files = resolve_transformer_checkpoint_files( component_server_args, component_model_path ) + safetensors_list = list(checkpoint_files.safetensors) + transformer_override_config_path = checkpoint_files.config_path # 2. dit config # Config from Diffusers supersedes sgl_diffusion's model config @@ -286,6 +289,7 @@ class TransformerLoader(ComponentLoader): component_name=component_name, gguf_file=gguf_file, checkpoint_quant_config=checkpoint_quant_config, + transformer_override_config_path=transformer_override_config_path, ) if quant_spec.gguf_file is not None and is_minimax_h3: assert quant_spec.quant_config is not None diff --git a/python/sglang/multimodal_gen/runtime/loader/transformer_load_utils.py b/python/sglang/multimodal_gen/runtime/loader/transformer_load_utils.py index 3216ce3bc..8694ad8a4 100644 --- a/python/sglang/multimodal_gen/runtime/loader/transformer_load_utils.py +++ b/python/sglang/multimodal_gen/runtime/loader/transformer_load_utils.py @@ -45,11 +45,6 @@ from sglang.multimodal_gen.runtime.managers.memory_managers.component_residency ) from sglang.multimodal_gen.runtime.platforms import current_platform from sglang.multimodal_gen.runtime.server_args import ServerArgs -from sglang.multimodal_gen.runtime.utils.hf_diffusers_utils import ( - hf_hub_download, - maybe_download_model, - snapshot_download, -) from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger from sglang.multimodal_gen.runtime.utils.precision import resolve_precision from sglang.multimodal_gen.runtime.utils.quantization_utils import ( @@ -58,6 +53,11 @@ from sglang.multimodal_gen.runtime.utils.quantization_utils import ( get_quant_config, get_quant_config_from_safetensors_metadata, ) +from sglang.multimodal_gen.runtime.weights.source import ( + materialize_weight_set, + materialize_weight_set_config, + resolve_safetensors_weight_set, +) from sglang.srt.utils.hf_transformers import ( check_gguf_file, resolve_hf_gguf_reference, @@ -71,11 +71,6 @@ _PRECISION_VARIANT_SUFFIX_RE = re.compile( r"^(?P.+?)(?P\.(?:fp16|bf16|fp32))(?P-\d+-of-\d+)?(?P\.safetensors)$" ) _MIXED_SAFETENSORS_RE = re.compile(r".*-mixed(?:-\d+-of-\d+)?\.safetensors$") -_HF_SAFETENSORS_URL_RE = re.compile( - r"https?://huggingface\.co/(?P[^/]+/[^/]+)/" - r"(?:blob|resolve)/(?P[^/]+)/(?P.+\.safetensors)$", - re.IGNORECASE, -) def _get_quant_config_name(config: Optional[QuantizationConfig]) -> Optional[str]: @@ -213,6 +208,14 @@ class TransformerQuantLoadSpec: ) +@dataclass(frozen=True) +class TransformerCheckpointFiles: + """Files from one checkpoint revision needed during transformer loading.""" + + safetensors: tuple[str, ...] + config_path: str | None + + class _TransformerQuantAdapter: def prepare(self) -> None: """initialize""" @@ -612,86 +615,55 @@ def resolve_transformer_gguf_to_load( return resolved -def resolve_transformer_safetensors_to_load( +def resolve_transformer_checkpoint_files( server_args: ServerArgs, component_model_path: str -) -> list[str]: +) -> TransformerCheckpointFiles: """Resolve transformer weights from the base component path or an override.""" quantized_path = server_args.transformer_weights_path if quantized_path: - original_quantized_path = quantized_path - direct_url = _HF_SAFETENSORS_URL_RE.fullmatch(original_quantized_path) - if direct_url is not None: - quantized_path = hf_hub_download( - repo_id=direct_url.group("repo"), - filename=direct_url.group("filename"), - revision=direct_url.group("revision"), - ) - else: - parts = original_quantized_path.strip("/").split("/") - is_hub_file = ( - not os.path.exists(original_quantized_path) - and not os.path.isabs(original_quantized_path) - and not original_quantized_path.startswith((".", "~")) - and len(parts) > 2 - and original_quantized_path.endswith(".safetensors") - ) - quantized_path = ( - hf_hub_download( - repo_id="/".join(parts[:2]), - filename="/".join(parts[2:]), - revision=server_args.revision, - ) - if is_hub_file - else maybe_download_model(original_quantized_path) - ) - logger.info("using quantized transformer weights from: %s", quantized_path) - if os.path.isfile(quantized_path) and quantized_path.endswith(".safetensors"): - safetensors_list = [quantized_path] - else: - safetensors_list = _list_safetensors_files(quantized_path) - if not safetensors_list and not os.path.exists(original_quantized_path): - logger.warning( - "No safetensors files found in cached transformer weights path " - "%s; refreshing snapshot for %s", - quantized_path, - original_quantized_path, - ) - quantized_path = snapshot_download( - repo_id=original_quantized_path, - ignore_patterns=["*.onnx", "*.msgpack"], - allow_patterns=[ - "*.json", - "*.safetensors", - "*.safetensors.index.json", - ], - max_workers=8, - ) - safetensors_list = _list_safetensors_files(quantized_path) - else: - safetensors_list = _list_safetensors_files(component_model_path) + resolved_set = resolve_safetensors_weight_set( + quantized_path, + revision=server_args.revision, + select_unindexed_weight=_select_single_mixed_safetensors_file, + ) + safetensors_list = materialize_weight_set(resolved_set) + logger.info( + "using transformer weight set from %s: %s", + quantized_path, + safetensors_list, + ) + return TransformerCheckpointFiles( + safetensors=safetensors_list, + config_path=materialize_weight_set_config(resolved_set), + ) + safetensors_list = _list_safetensors_files(component_model_path) if safetensors_list: - # Diffusers repos occasionally ship more than one shard split for the - # same checkpoint (e.g. a 4-way and an 8-way split side by side). The - # index file is the authoritative source for which files belong to - # the checkpoint that was actually exported; anything else is a - # leftover sibling variant. + # Preserve legacy cleanup for the base component. Explicit overrides + # are resolved above, where an index is already the final authority. safetensors_list = filter_duplicate_safetensors_files( safetensors_list, os.path.dirname(safetensors_list[0]), SAFE_WEIGHTS_INDEX_NAME, ) - - safetensors_list = _prefer_mixed_safetensors_files(safetensors_list) - safetensors_list = _filter_duplicate_precision_variant_safetensors(safetensors_list) - - if not safetensors_list: - raise ValueError( - f"no safetensors files found in {quantized_path or component_model_path}" + safetensors_list = _prefer_mixed_safetensors_files(safetensors_list) + safetensors_list = _filter_duplicate_precision_variant_safetensors( + safetensors_list ) - return safetensors_list + if not safetensors_list: + raise ValueError(f"no safetensors files found in {component_model_path}") + + return TransformerCheckpointFiles(tuple(safetensors_list), None) + + +def _select_single_mixed_safetensors_file( + candidates: tuple[str, ...], +) -> str | None: + """Preserve the transformer's established mixed-export preference.""" + mixed = tuple(path for path in candidates if _MIXED_SAFETENSORS_RE.fullmatch(path)) + return mixed[0] if len(mixed) == 1 else None def _prefer_mixed_safetensors_files(safetensors_list: list[str]) -> list[str]: @@ -772,6 +744,7 @@ def resolve_transformer_quant_load_spec( component_name: str | None = None, gguf_file: str | None = None, checkpoint_quant_config: QuantizationConfig | None = None, + transformer_override_config_path: str | None = None, ) -> TransformerQuantLoadSpec: if gguf_file is not None: if checkpoint_quant_config is not None: @@ -803,6 +776,7 @@ def resolve_transformer_quant_load_spec( server_args=server_args, safetensors_list=safetensors_list, component_model_path=component_model_path, + transformer_override_config_path=transformer_override_config_path, ) if quant_config is not None: @@ -945,37 +919,15 @@ def _build_transformer_quant_adapters( def _resolve_quant_config_from_transformer_override( - transformer_weights_path: str, + override_config_path: str, ) -> Optional[QuantizationConfig]: """Resolve quant config from an override transformer repo or directory.""" - expanded_path = os.path.expanduser(transformer_weights_path) - if os.path.isfile(expanded_path): - return None - - # A single local safetensors file does not carry a directory-level config.json. - # Let downstream metadata probing handle it instead of misrouting it through HF. - if expanded_path.endswith(".safetensors") and ( - os.path.isabs(expanded_path) - or expanded_path.startswith(".") - or os.sep in expanded_path - or (os.path.altsep and os.path.altsep in expanded_path) - ): - return None - - override_quantized_path = maybe_download_model(transformer_weights_path) - if not os.path.isdir(override_quantized_path): - return None - - override_config_path = os.path.join(override_quantized_path, "config.json") - if not os.path.isfile(override_config_path): - return None - with open(override_config_path, encoding="utf-8") as f: override_hf_config = json.load(f) return get_quant_config( override_hf_config, - override_quantized_path, + os.path.dirname(override_config_path), ) @@ -985,6 +937,7 @@ def _resolve_quant_config( server_args: ServerArgs, safetensors_list: list[str], component_model_path: str, + transformer_override_config_path: str | None = None, ) -> Optional[QuantizationConfig]: """ resolve quant config from checkpoints' metadata @@ -1055,11 +1008,11 @@ def _resolve_quant_config( fallback_group_size, ) quant_config = _merge_modelopt_fp4_configs(quant_config, inferred_nvfp4_config) - if quant_config is not None or not server_args.transformer_weights_path: + if quant_config is not None or transformer_override_config_path is None: return quant_config quant_config = _resolve_quant_config_from_transformer_override( - server_args.transformer_weights_path + transformer_override_config_path, ) quant_config = _merge_modelopt_fp4_configs(quant_config, inferred_nvfp4_config) if quant_config is not None: diff --git a/python/sglang/multimodal_gen/runtime/weights/source.py b/python/sglang/multimodal_gen/runtime/weights/source.py index d2a7ae3c3..d5412158c 100644 --- a/python/sglang/multimodal_gen/runtime/weights/source.py +++ b/python/sglang/multimodal_gen/runtime/weights/source.py @@ -2,8 +2,10 @@ from __future__ import annotations +import json import os -from dataclasses import dataclass +from collections.abc import Callable, Mapping +from dataclasses import dataclass, replace from pathlib import Path, PurePosixPath from typing import Literal from urllib.parse import unquote, urlparse @@ -14,6 +16,8 @@ from huggingface_hub.utils import validate_repo_id WeightSourceKind = Literal["local", "huggingface"] _WEIGHT_SUFFIXES = (".safetensors", ".gguf", ".bin", ".pt", ".pth", ".ckpt") +_SAFETENSORS_INDEX_SUFFIX = ".safetensors.index.json" +_WEIGHT_REFERENCE_SUFFIXES = _WEIGHT_SUFFIXES + (_SAFETENSORS_INDEX_SUFFIX,) @dataclass(frozen=True) @@ -40,6 +44,17 @@ class ResolvedWeight: selected_file: str +@dataclass(frozen=True) +class ResolvedWeightSet: + inventory: WeightInventory + selected_files: tuple[str, ...] + index_file: str | None = None + + +class NoSafetensorsWeightsError(FileNotFoundError): + """The source has no safetensors payload to resolve as a weight set.""" + + def is_explicit_weight_file_reference(source: str) -> bool: """Whether a component override names one weight file, not a component root.""" expanded = os.path.expanduser(source) @@ -65,6 +80,23 @@ def _merge_revision(url_revision: str | None, revision: str | None) -> str | Non return url_revision or revision +def _local_index_inventory(index_path: Path) -> tuple[str, ...]: + """List only an exact local index and the shards it declares.""" + with index_path.open(encoding="utf-8") as index_stream: + index = json.load(index_stream) + weight_map = index.get("weight_map") if isinstance(index, Mapping) else None + shard_names = weight_map.values() if isinstance(weight_map, Mapping) else () + files = {index_path.name} + for shard_name in shard_names: + if not isinstance(shard_name, str): + continue + shard_name = _validate_relative_hub_path(shard_name, "index shard") + shard_path = index_path.parent / shard_name + if shard_path.is_file(): + files.add(shard_path.relative_to(index_path.parent).as_posix()) + return tuple(sorted(files)) + + def _parse_huggingface_url(source: str, revision: str | None) -> WeightSource: parsed = urlparse(source) if parsed.netloc.lower() not in ("huggingface.co", "www.huggingface.co"): @@ -149,7 +181,7 @@ def parse_weight_source( tail = "/".join(parts[2:]) or None filename = ( _validate_relative_hub_path(tail, "filename") - if tail is not None and tail.lower().endswith(_WEIGHT_SUFFIXES) + if tail is not None and tail.lower().endswith(_WEIGHT_REFERENCE_SUFFIXES) else None ) subfolder = tail if filename is None else None @@ -169,12 +201,11 @@ def _filter_inventory_files( files: tuple[str, ...], source: WeightSource ) -> tuple[str, ...]: if source.filename is not None: - selected = tuple(path for path in files if path == source.filename) - if not selected: + if source.filename not in files: raise FileNotFoundError( f"Weight file {source.filename!r} was not found in {source.repo_id}" ) - return selected + return (source.filename,) if source.subfolder is None: return files prefix = source.subfolder.rstrip("/") + "/" @@ -245,7 +276,7 @@ def select_weight_file( path for path in inventory.files if path.lower().endswith(_WEIGHT_SUFFIXES) ) if inventory.source.filename is not None: - return inventory.files[0] + return _select_named_file(candidates, inventory.source.filename) if weight_name is not None: return _select_named_file(candidates, weight_name) if len(candidates) == 1: @@ -274,18 +305,210 @@ def resolve_weight( ) -def materialize_weight(resolved: ResolvedWeight) -> str: - """Return the selected local file, downloading one pinned Hub file if needed.""" - source = resolved.inventory.source +def _materialize_inventory_file(inventory: WeightInventory, filename: str) -> str: + source = inventory.source if source.kind == "local": assert source.local_path is not None if os.path.isfile(source.local_path): - return source.local_path - return os.path.join(source.local_path, resolved.selected_file) + local_path = Path(source.local_path) + if filename == local_path.name: + return source.local_path + return str(local_path.parent / filename) + return os.path.join(source.local_path, filename) assert source.repo_id is not None return hf_hub_download( repo_id=source.repo_id, - filename=resolved.selected_file, - revision=resolved.inventory.resolved_revision, + filename=filename, + revision=inventory.resolved_revision, ) + + +def _resolve_index_shard( + inventory: WeightInventory, index_file: str, shard_name: str +) -> str: + shard_name = _validate_relative_hub_path(shard_name, "index shard") + if not shard_name.lower().endswith(".safetensors"): + raise ValueError( + f"Safetensors index {index_file!r} references a non-safetensors " + f"shard: {shard_name!r}" + ) + index_parent = PurePosixPath(index_file).parent + candidates = (shard_name, (index_parent / shard_name).as_posix()) + matches = tuple( + dict.fromkeys(path for path in candidates if path in inventory.files) + ) + if len(matches) == 1: + return matches[0] + if not matches: + raise FileNotFoundError( + f"Safetensors index {index_file!r} references missing shard {shard_name!r}" + ) + raise ValueError( + f"Safetensors index shard {shard_name!r} is ambiguous: {list(matches)}" + ) + + +def _read_safetensors_index( + inventory: WeightInventory, index_file: str +) -> tuple[str, ...]: + with open( + _materialize_inventory_file(inventory, index_file), encoding="utf-8" + ) as index_stream: + index = json.load(index_stream) + weight_map = index.get("weight_map") if isinstance(index, Mapping) else None + if not isinstance(weight_map, Mapping) or not weight_map: + raise ValueError( + f"Safetensors index {index_file!r} must contain a non-empty weight_map" + ) + shard_names = tuple(weight_map.values()) + if not all(isinstance(name, str) for name in shard_names): + raise ValueError( + f"Safetensors index {index_file!r} contains a non-string shard name" + ) + return tuple( + sorted( + { + _resolve_index_shard(inventory, index_file, shard_name) + for shard_name in shard_names + } + ) + ) + + +def resolve_safetensors_weight_set( + source: str, + *, + revision: str | None = None, + weight_name: str | None = None, + select_unindexed_weight: Callable[[tuple[str, ...]], str | None] | None = None, +) -> ResolvedWeightSet: + """Resolve one safetensors checkpoint, using its index as shard authority.""" + parsed_source = parse_weight_source(source, revision=revision) + if ( + parsed_source.kind == "local" + and parsed_source.local_path is not None + and os.path.isfile(parsed_source.local_path) + and parsed_source.local_path.lower().endswith(_SAFETENSORS_INDEX_SUFFIX) + ): + inventory = WeightInventory( + parsed_source, + None, + _local_index_inventory(Path(parsed_source.local_path)), + ) + elif parsed_source.kind == "huggingface" and parsed_source.filename is not None: + parent = PurePosixPath(parsed_source.filename).parent + scoped_inventory = resolve_weight_inventory( + replace( + parsed_source, + filename=None, + subfolder=None if parent == PurePosixPath(".") else parent.as_posix(), + ) + ) + inventory = WeightInventory( + parsed_source, + scoped_inventory.resolved_revision, + scoped_inventory.files, + ) + else: + inventory = resolve_weight_inventory(parsed_source) + weights = tuple( + path for path in inventory.files if path.lower().endswith(".safetensors") + ) + indexes = tuple( + path + for path in inventory.files + if path.lower().endswith(_SAFETENSORS_INDEX_SUFFIX) + ) + selected_name = inventory.source.filename or weight_name + if ( + selected_name is None + and inventory.source.kind == "local" + and inventory.source.local_path is not None + and os.path.isfile(inventory.source.local_path) + ): + selected_name = Path(inventory.source.local_path).name + if selected_name is not None: + if not selected_name.lower().endswith( + (".safetensors", _SAFETENSORS_INDEX_SUFFIX) + ): + raise NoSafetensorsWeightsError( + f"Selected file is not safetensors: {selected_name!r}" + ) + selected = _select_named_file(weights + indexes, selected_name) + if selected in weights: + return ResolvedWeightSet(inventory, (selected,)) + return ResolvedWeightSet( + inventory, + _read_safetensors_index(inventory, selected), + index_file=selected, + ) + if len(indexes) == 1: + return ResolvedWeightSet( + inventory, + _read_safetensors_index(inventory, indexes[0]), + index_file=indexes[0], + ) + if len(indexes) > 1: + raise ValueError( + "Source contains multiple safetensors indexes; select one with an " + f"exact index name. Candidates: {list(indexes)}" + ) + if len(weights) == 1: + return ResolvedWeightSet(inventory, weights) + if not weights: + raise NoSafetensorsWeightsError("Source contains no safetensors weights") + if select_unindexed_weight is not None: + selected = select_unindexed_weight(weights) + if selected is not None: + if selected not in weights: + raise ValueError( + "Unindexed weight selector returned a file outside the source: " + f"{selected!r}" + ) + return ResolvedWeightSet(inventory, (selected,)) + raise ValueError( + "Source contains multiple safetensors files without an index; they may " + "be independent variants. Select one exact file or provide a standard " + f"safetensors index. Candidates: {list(weights)}" + ) + + +def materialize_weight_set(resolved: ResolvedWeightSet) -> tuple[str, ...]: + """Materialize every selected file from one pinned checkpoint revision.""" + return tuple( + _materialize_inventory_file(resolved.inventory, filename) + for filename in resolved.selected_files + ) + + +def materialize_weight_set_config(resolved: ResolvedWeightSet) -> str | None: + """Materialize adjacent runtime configuration from the pinned revision.""" + source = resolved.inventory.source + if source.kind == "local" and source.local_path is not None: + local_path = Path(source.local_path) + if local_path.is_file(): + config_path = local_path.with_name("config.json") + return str(config_path) if config_path.is_file() else None + + anchor = resolved.index_file or resolved.selected_files[0] + config_file = PurePosixPath(anchor).with_name("config.json").as_posix() + if config_file not in resolved.inventory.files: + return None + config_parent = PurePosixPath(config_file).parent + metadata_files = tuple( + path + for path in resolved.inventory.files + if PurePosixPath(path).parent == config_parent + and PurePosixPath(path).name.startswith("quant_model_description") + and path.lower().endswith(".json") + ) + config_path = _materialize_inventory_file(resolved.inventory, config_file) + for metadata_file in metadata_files: + _materialize_inventory_file(resolved.inventory, metadata_file) + return config_path + + +def materialize_weight(resolved: ResolvedWeight) -> str: + """Return the selected local file, downloading one pinned Hub file if needed.""" + return _materialize_inventory_file(resolved.inventory, resolved.selected_file) diff --git a/python/sglang/multimodal_gen/test/unit/test_transformer_quant.py b/python/sglang/multimodal_gen/test/unit/test_transformer_quant.py index 9212a64eb..18c75ba27 100644 --- a/python/sglang/multimodal_gen/test/unit/test_transformer_quant.py +++ b/python/sglang/multimodal_gen/test/unit/test_transformer_quant.py @@ -3,7 +3,6 @@ This unittest is introduced in #22360, preventing duplicate transformer safetens """ import json -import os import sys import tempfile import types @@ -100,8 +99,8 @@ from sglang.multimodal_gen.runtime.loader.transformer_load_utils import ( _Flux2Nvfp4FallbackAdapter, _needs_device_weight_postprocess, _resolve_quant_config, + resolve_transformer_checkpoint_files, resolve_transformer_quant_load_spec, - resolve_transformer_safetensors_to_load, ) from sglang.multimodal_gen.runtime.loader.weight_load_plan import WeightLoadPlan from sglang.multimodal_gen.runtime.models.dits.flux import FluxSingleTransformerBlock @@ -262,50 +261,79 @@ class TestTransformerQuantHelpers(unittest.TestCase): self.assertIsNone(backend) - def test_resolve_transformer_safetensors_to_load_uses_single_override_file(self): + def test_resolve_transformer_checkpoint_files_uses_single_override_file(self): with tempfile.NamedTemporaryFile(suffix=".safetensors") as f: server_args = self._make_server_args(transformer_weights_path=f.name) - resolved = resolve_transformer_safetensors_to_load( - server_args, "/unused/component/path" - ) + with ( + patch( + "sglang.multimodal_gen.runtime.weights.source.HfApi.model_info" + ) as model_info, + patch( + "sglang.multimodal_gen.runtime.weights.source.hf_hub_download" + ) as download, + ): + resolved = resolve_transformer_checkpoint_files( + server_args, "/unused/component/path" + ) - self.assertEqual(resolved, [f.name]) + self.assertEqual(resolved.safetensors, (f.name,)) + self.assertIsNone(resolved.config_path) + model_info.assert_not_called() + download.assert_not_called() - @patch( - "sglang.multimodal_gen.runtime.loader.transformer_load_utils.hf_hub_download", - return_value="/cache/model.safetensors", - ) - def test_resolve_transformer_safetensors_to_load_uses_hf_file_reference( - self, mock_download - ): - filename = "diffusion_models/minimax_h3_fl2va_pruned_bf16.safetensors" + def test_resolve_transformer_checkpoint_files_uses_one_hf_revision(self): + filename = "weights/model.safetensors" references = ( ( - f"https://huggingface.co/Comfy-Org/MiniMax-H3/resolve/main/{filename}", + f"https://huggingface.co/owner/repo/resolve/main/{filename}", "main", ), - (f"Comfy-Org/MiniMax-H3/{filename}", "test-revision"), + (f"owner/repo/{filename}", "test-revision"), ) for reference, revision in references: with self.subTest(reference=reference): - server_args = self._make_server_args(transformer_weights_path=reference) - with patch( - "os.path.isfile", - side_effect=lambda path: path == "/cache/model.safetensors", - ): - self.assertEqual( - resolve_transformer_safetensors_to_load( - server_args, "/unused/component/path" - ), - ["/cache/model.safetensors"], - ) - mock_download.assert_called_once_with( - repo_id="Comfy-Org/MiniMax-H3", - filename=filename, + server_args = self._make_server_args( + transformer_weights_path=reference, revision=revision, ) - mock_download.reset_mock() + model_info = SimpleNamespace( + sha="immutable-sha", + siblings=[ + SimpleNamespace(rfilename=filename), + SimpleNamespace(rfilename="weights/config.json"), + ], + ) + + def download(*, filename, **_kwargs): + return f"/cache/{filename.rsplit('/', 1)[-1]}" + + with ( + patch( + "sglang.multimodal_gen.runtime.weights.source.HfApi.model_info", + return_value=model_info, + ) as model_info_call, + patch( + "sglang.multimodal_gen.runtime.weights.source.hf_hub_download", + side_effect=download, + ) as download, + ): + resolved = resolve_transformer_checkpoint_files( + server_args, "/unused/component/path" + ) + self.assertEqual(resolved.safetensors, ("/cache/model.safetensors",)) + self.assertEqual(resolved.config_path, "/cache/config.json") + model_info_call.assert_called_once_with("owner/repo", revision=revision) + self.assertEqual( + {call.kwargs["filename"] for call in download.call_args_list}, + {filename, "weights/config.json"}, + ) + self.assertTrue( + all( + call.kwargs["revision"] == "immutable-sha" + for call in download.call_args_list + ) + ) def test_inspect_minimax_h3_safetensors_detects_curve_and_comfy_format(self): marker = json.dumps( @@ -739,13 +767,7 @@ class TestTransformerQuantHelpers(unittest.TestCase): checkpoint_quant_config=ComfyFp8Config({}), ) - @patch( - "sglang.multimodal_gen.runtime.loader.transformer_load_utils.maybe_download_model", - side_effect=lambda path, **kw: path, - ) - def test_resolve_transformer_safetensors_to_load_prefers_mixed_export( - self, _mock_download - ): + def test_resolve_transformer_override_prefers_single_mixed_export(self): with tempfile.TemporaryDirectory() as tmpdir: mixed = f"{tmpdir}/flux2-dev-nvfp4-mixed.safetensors" full = f"{tmpdir}/flux2-dev-nvfp4.safetensors" @@ -753,39 +775,11 @@ class TestTransformerQuantHelpers(unittest.TestCase): open(full, "a").close() server_args = self._make_server_args(transformer_weights_path=tmpdir) - resolved = resolve_transformer_safetensors_to_load( + resolved = resolve_transformer_checkpoint_files( server_args, "/unused/component/path" ) - self.assertEqual(resolved, [mixed]) - - @patch( - "sglang.multimodal_gen.runtime.loader.transformer_load_utils.snapshot_download", - ) - @patch( - "sglang.multimodal_gen.runtime.loader.transformer_load_utils.maybe_download_model", - ) - def test_resolve_transformer_safetensors_to_load_refreshes_empty_cached_repo( - self, mock_download_model, mock_snapshot_download - ): - with tempfile.TemporaryDirectory() as cached_dir: - repo_id = "black-forest-labs/FLUX.2-dev-NVFP4" - mixed = os.path.join(cached_dir, "flux2-dev-nvfp4-mixed.safetensors") - mock_download_model.return_value = cached_dir - - def _snapshot_download(**_kwargs): - open(mixed, "a").close() - return cached_dir - - mock_snapshot_download.side_effect = _snapshot_download - - server_args = self._make_server_args(transformer_weights_path=repo_id) - resolved = resolve_transformer_safetensors_to_load( - server_args, "/unused/component/path" - ) - - self.assertEqual(resolved, [mixed]) - mock_snapshot_download.assert_called_once() + self.assertEqual(resolved.safetensors, (mixed,)) def test_filter_transformer_precision_variants_prefers_canonical_file(self): files = [ @@ -957,9 +951,6 @@ class TestTransformerQuantHelpers(unittest.TestCase): "sglang.multimodal_gen.runtime.loader.transformer_load_utils.build_nvfp4_config_from_safetensors_list", return_value=None, ) - @patch( - "sglang.multimodal_gen.runtime.loader.transformer_load_utils.maybe_download_model" - ) @patch( "sglang.multimodal_gen.runtime.loader.transformer_load_utils.get_quant_config_from_safetensors_metadata", return_value=None, @@ -967,21 +958,12 @@ class TestTransformerQuantHelpers(unittest.TestCase): @patch( "sglang.multimodal_gen.runtime.loader.transformer_load_utils.get_metadata_from_safetensors_file" ) - @patch( - "sglang.multimodal_gen.runtime.loader.transformer_load_utils.maybe_download_model", - side_effect=lambda path, **kw: path, - ) def test_resolve_transformer_quant_load_spec_keeps_nunchaku_hook( self, - _mock_download, mock_metadata, _mock_quant_metadata, - mock_maybe_download, _mock_nvfp4, ): - mock_maybe_download.side_effect = AssertionError( - "local safetensors path should not trigger maybe_download_model" - ) mock_metadata.return_value = { "config": json.dumps({"_class_name": _FakeFluxTransformer.__name__}) } @@ -1006,7 +988,6 @@ class TestTransformerQuantHelpers(unittest.TestCase): self.assertIsNone(spec.param_dtype) self.assertEqual(len(spec.post_load_hooks), 1) self.assertIs(nunchaku_config.model_cls, _FakeFluxTransformer) - mock_maybe_download.assert_not_called() def test_flux2_mixed_nvfp4_fallback_disables_conflicting_offloads(self): server_args = self._make_server_args( diff --git a/python/sglang/multimodal_gen/test/unit/test_weight_source.py b/python/sglang/multimodal_gen/test/unit/test_weight_source.py index 8332d6938..77ddc2c86 100644 --- a/python/sglang/multimodal_gen/test/unit/test_weight_source.py +++ b/python/sglang/multimodal_gen/test/unit/test_weight_source.py @@ -4,9 +4,13 @@ from unittest.mock import patch import pytest from sglang.multimodal_gen.runtime.weights.source import ( + NoSafetensorsWeightsError, is_explicit_weight_file_reference, materialize_weight, + materialize_weight_set, + materialize_weight_set_config, parse_weight_source, + resolve_safetensors_weight_set, resolve_weight, resolve_weight_inventory, ) @@ -18,6 +22,9 @@ def test_explicit_weight_file_reference_does_not_claim_directories(tmp_path): assert not is_explicit_weight_file_reference(str(component)) assert is_explicit_weight_file_reference("owner/repo/model.safetensors") + assert not is_explicit_weight_file_reference( + "owner/repo/model.safetensors.index.json" + ) assert is_explicit_weight_file_reference( "https://huggingface.co/owner/repo/resolve/main/model.gguf?download=true" ) @@ -95,6 +102,137 @@ def test_weight_source_rejects_ambiguous_files(tmp_path): resolve_weight(str(tmp_path)) +def test_safetensors_index_selects_only_declared_shards(tmp_path): + (tmp_path / "model-00001-of-00002.safetensors").write_bytes(b"one") + (tmp_path / "model-00002-of-00002.safetensors").write_bytes(b"two") + (tmp_path / "alternate.safetensors").write_bytes(b"other variant") + (tmp_path / "model.safetensors.index.json").write_text( + '{"weight_map":{"a":"model-00001-of-00002.safetensors",' + '"b":"model-00002-of-00002.safetensors"}}' + ) + + resolved = resolve_safetensors_weight_set(str(tmp_path)) + + assert resolved.index_file == "model.safetensors.index.json" + assert resolved.selected_files == ( + "model-00001-of-00002.safetensors", + "model-00002-of-00002.safetensors", + ) + assert materialize_weight_set(resolved) == tuple( + str(tmp_path / filename) for filename in resolved.selected_files + ) + + +def test_exact_local_safetensors_index_resolves_adjacent_shards(tmp_path): + shard = tmp_path / "model-00001-of-00001.safetensors" + shard.write_bytes(b"one") + index = tmp_path / "model.safetensors.index.json" + index.write_text('{"weight_map":{"a":"model-00001-of-00001.safetensors"}}') + + resolved = resolve_safetensors_weight_set(str(index)) + + assert resolved.index_file == index.name + assert materialize_weight_set(resolved) == (str(shard),) + + +def test_safetensors_weight_set_rejects_unindexed_variants(tmp_path): + (tmp_path / "base.safetensors").write_bytes(b"base") + (tmp_path / "distilled.safetensors").write_bytes(b"distilled") + + with pytest.raises(ValueError, match="without an index"): + resolve_safetensors_weight_set(str(tmp_path)) + + +def test_safetensors_index_rejects_non_weight_shard(tmp_path): + (tmp_path / "config.json").write_text("{}") + (tmp_path / "model.safetensors.index.json").write_text( + '{"weight_map":{"a":"config.json"}}' + ) + + with pytest.raises(ValueError, match="non-safetensors shard"): + resolve_safetensors_weight_set(str(tmp_path)) + + +def test_safetensors_index_missing_shard_is_not_no_weights_fallback(tmp_path): + (tmp_path / "model.safetensors.index.json").write_text( + '{"weight_map":{"a":"missing.safetensors"}}' + ) + + with pytest.raises(FileNotFoundError, match="missing shard") as error: + resolve_safetensors_weight_set(str(tmp_path)) + + assert not isinstance(error.value, NoSafetensorsWeightsError) + + +def test_source_without_safetensors_has_distinct_error(tmp_path): + (tmp_path / "pytorch_model.bin").write_bytes(b"bin") + + with pytest.raises(NoSafetensorsWeightsError): + resolve_safetensors_weight_set(str(tmp_path)) + + +def test_exact_non_safetensors_file_has_distinct_error(tmp_path): + weights = tmp_path / "pytorch_model.bin" + weights.write_bytes(b"bin") + + with pytest.raises(NoSafetensorsWeightsError): + resolve_safetensors_weight_set(str(weights)) + + +def test_remote_safetensors_shards_use_one_pinned_revision(tmp_path): + index_path = tmp_path / "model.safetensors.index.json" + index_path.write_text( + '{"weight_map":{"a":"model-00001-of-00002.safetensors",' + '"b":"model-00002-of-00002.safetensors"}}' + ) + config_path = tmp_path / "config.json" + config_path.write_text("{}") + model_info = SimpleNamespace( + sha="immutable-sha", + siblings=[ + SimpleNamespace(rfilename=f"transformer/{filename}") + for filename in ( + "model.safetensors.index.json", + "model-00001-of-00002.safetensors", + "model-00002-of-00002.safetensors", + "config.json", + "quant_model_description_w8a8.json", + ) + ], + ) + + def download(*, filename, **_kwargs): + if filename.endswith("index.json"): + return str(index_path) + if filename.endswith("config.json"): + return str(config_path) + if "quant_model_description" in filename: + return str(tmp_path / "quant_model_description_w8a8.json") + return f"/{filename}" + + with ( + patch( + "sglang.multimodal_gen.runtime.weights.source.HfApi.model_info", + return_value=model_info, + ), + patch( + "sglang.multimodal_gen.runtime.weights.source.hf_hub_download", + side_effect=download, + ) as hub_download, + ): + resolved = resolve_safetensors_weight_set("owner/repo/transformer") + materialize_weight_set(resolved) + assert materialize_weight_set_config(resolved) == str(config_path) + + assert all( + call.kwargs["revision"] == "immutable-sha" + for call in hub_download.call_args_list + ) + assert "transformer/quant_model_description_w8a8.json" in { + call.kwargs["filename"] for call in hub_download.call_args_list + } + + def test_materialize_local_weight_returns_selected_file(tmp_path): checkpoint = tmp_path / "model.safetensors" checkpoint.write_bytes(b"fixture")