Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -67,12 +67,13 @@
remap_gguf_tensor_meta,
)
from sglang.multimodal_gen.runtime.loader.utils import (
_list_safetensors_files,
checkpoint_bytes,
get_param_names_mapping,
set_default_torch_dtype,
skip_init_modules,
)
from sglang.multimodal_gen.runtime.loader.weight_utils import (
filter_duplicate_safetensors_files,
filter_files_not_needed_for_inference,
pt_weights_iterator,
safetensors_weights_iterator,
Expand Down Expand Up @@ -462,31 +463,16 @@ def _require_quantized_encoder_layers(
)


def _checkpoint_bytes(model_path: str) -> int:
"""On-disk size of a checkpoint, readable before any weight of it is."""
if os.path.isfile(model_path):
return os.path.getsize(model_path)
total = 0
for path in glob.glob(
os.path.join(str(model_path), "**", "*.safetensors"), recursive=True
):
try:
total += os.path.getsize(path)
except OSError:
continue
return total


def _keep_this_checkpoint_mapped(model_path: str) -> bool:
"""Whether this encoder's weights should stay on their file mapping."""
checkpoint_bytes = _checkpoint_bytes(model_path)
if not host_copies_would_not_fit(checkpoint_bytes):
weight_bytes = checkpoint_bytes(model_path)
if not host_copies_would_not_fit(weight_bytes):
return False
logger.info(
"Text encoder checkpoint is %.2f GiB against %.2f GiB of host memory, "
"so its compatible weights stay on the checkpoint mapping instead of "
"being copied in.",
checkpoint_bytes / 1024**3,
weight_bytes / 1024**3,
host_memory_available_bytes() / 1024**3,
)
return True
Expand Down Expand Up @@ -598,20 +584,20 @@ def _prepare_weights(

hf_weights_files: list[str] = []
for pattern in allow_patterns:
hf_weights_files += glob.glob(os.path.join(hf_folder, pattern))
if pattern == "*.safetensors":
hf_weights_files = _list_safetensors_files(
hf_folder,
index_file=index_file,
key_filter=key_filter,
)
else:
hf_weights_files = glob.glob(os.path.join(hf_folder, pattern))
if len(hf_weights_files) > 0:
if pattern == "*.safetensors":
use_safetensors = True
break

if use_safetensors:
hf_weights_files = filter_duplicate_safetensors_files(
hf_weights_files,
hf_folder,
index_file,
key_filter=key_filter,
)
else:
if not use_safetensors:
hf_weights_files = filter_files_not_needed_for_inference(hf_weights_files)

if len(hf_weights_files) == 0:
Expand Down
Original file line number Diff line number Diff line change
@@ -1,4 +1,3 @@
import glob
import json
import os
import re
Expand All @@ -10,6 +9,7 @@
from sglang.multimodal_gen.runtime.loader.component_loaders.component_loader import (
PlainStateDictComponentLoader,
)
from sglang.multimodal_gen.runtime.loader.utils import _list_safetensors_files
from sglang.multimodal_gen.runtime.models.upsampler.latent_upsampler import (
LatentUpsampler,
)
Expand Down Expand Up @@ -59,7 +59,7 @@ def _find_safetensors_file(path: str) -> str:
return path

if os.path.isdir(path):
files = sorted(glob.glob(os.path.join(path, "*.safetensors")))
files = _list_safetensors_files(path)
if len(files) == 1:
return files[0]
elif len(files) > 1:
Expand All @@ -75,7 +75,7 @@ def _find_safetensors_file(path: str) -> str:
try:
maybe_downloaded = maybe_download_model(path)
if os.path.isdir(maybe_downloaded):
files = sorted(glob.glob(os.path.join(maybe_downloaded, "*.safetensors")))
files = _list_safetensors_files(maybe_downloaded)
if len(files) == 1:
return files[0]
elif len(files) > 1:
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -42,6 +42,9 @@
resolve_component_precision,
resolve_decode_precision,
)
from sglang.multimodal_gen.runtime.weights.source import (
filter_duplicate_precision_variant_safetensors,
)
from sglang.multimodal_gen.utils import PRECISION_TO_TYPE
from sglang.srt.model_loader.checkpoint_quantization import (
resolve_checkpoint_quant_spec,
Expand Down Expand Up @@ -568,14 +571,21 @@ def load_customized(
)
safetensors_list = [component_weights_path]
else:
safetensors_list = _list_safetensors_files(component_weights_path)
# VAE configs may explicitly choose a precision variant, so their
# selector must run before the canonical fallback.
safetensors_list = _list_safetensors_files(
component_weights_path, raw_candidates=True
)
safetensors_list = self.select_weight_files(
safetensors_list,
component_weights_path,
server_args,
component_name,
vae_precision,
)
safetensors_list = filter_duplicate_precision_variant_safetensors(
safetensors_list
)

assert len(safetensors_list) >= 1, (
f"Found no safetensors files in {component_weights_path}"
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -14,7 +14,6 @@
from typing import Callable, Optional

import torch
from diffusers.utils import SAFE_WEIGHTS_INDEX_NAME
from safetensors import safe_open
from torch import nn

Expand All @@ -38,9 +37,6 @@
read_gguf_tensor_meta,
)
from sglang.multimodal_gen.runtime.loader.utils import _list_safetensors_files
from sglang.multimodal_gen.runtime.loader.weight_utils import (
filter_duplicate_safetensors_files,
)
from sglang.multimodal_gen.runtime.managers.memory_managers.component_residency import (
COMPONENT_OFFLOAD,
ComponentResidencyError,
Expand Down Expand Up @@ -72,9 +68,6 @@

PostLoadHook = Callable[[nn.Module], None]

_PRECISION_VARIANT_SUFFIX_RE = re.compile(
r"^(?P<stem>.+?)(?P<precision>\.(?:fp16|bf16|fp32))(?P<shard>-\d+-of-\d+)?(?P<ext>\.safetensors)$"
)
_MIXED_SAFETENSORS_RE = re.compile(r".*-mixed(?:-\d+-of-\d+)?\.safetensors$")


Expand Down Expand Up @@ -645,17 +638,7 @@ def resolve_transformer_checkpoint_files(

safetensors_list = _list_safetensors_files(component_model_path)
if safetensors_list:
# 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 {component_model_path}")
Expand Down Expand Up @@ -696,48 +679,6 @@ def _prefer_mixed_safetensors_files(safetensors_list: list[str]) -> list[str]:
return mixed_files


def _filter_duplicate_precision_variant_safetensors(
safetensors_list: list[str],
) -> list[str]:
"""Drop precision-specific duplicates when a canonical file is present.

Diffusers checkpoints sometimes ship both `foo.safetensors` and
`foo.fp16.safetensors` (and their sharded variants) in the same directory.
Loading both is unsafe because duplicate parameter names race and whichever
tensor arrives last wins, leading to non-deterministic behavior

If a canonical unsuffixed (non bf16|fp32) file exists, prefer it and drop the precision
variant from the same family. Precision-only families are left untouched.
"""
canonical_paths = set(safetensors_list)
filtered: list[str] = []
removed: list[str] = []

for path in safetensors_list:
match = _PRECISION_VARIANT_SUFFIX_RE.match(path)
if match is None:
filtered.append(path)
continue

canonical_path = (
f"{match.group('stem')}{match.group('shard') or ''}{match.group('ext')}"
)
if canonical_path in canonical_paths:
removed.append(path)
continue

filtered.append(path)

if removed:
logger.info(
"Filtered %d duplicate transformer precision variant file(s): %s",
len(removed),
removed,
)

return filtered


def resolve_transformer_quant_load_spec(
*,
hf_config: dict,
Expand Down
84 changes: 57 additions & 27 deletions python/sglang/multimodal_gen/runtime/loader/utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -18,9 +18,14 @@
from torch import nn

from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
from sglang.multimodal_gen.runtime.weights.source import (
filter_duplicate_precision_variant_safetensors,
)

logger = init_logger(__name__)

_DEFAULT_SAFETENSORS_INDEX = "diffusion_pytorch_model.safetensors.index.json"

_QUANTIZED_DTYPES = {
torch.uint8,
torch.float8_e4m3fn,
Expand Down Expand Up @@ -250,11 +255,15 @@ def _try_redownload_missing_shards(model_path: str, missing: list[str]) -> bool:


def checkpoint_bytes(model_path: str) -> int:
"""On-disk size of every safetensors under a path, readable before any is."""
"""On-disk size of the selected safetensors checkpoint files."""
if os.path.isfile(model_path):
return os.path.getsize(model_path)

paths = sorted(
glob.glob(os.path.join(str(model_path), "**", "*.safetensors"), recursive=True)
)
total = 0
for path in glob.glob(
os.path.join(str(model_path), "**", "*.safetensors"), recursive=True
):
for path in filter_duplicate_precision_variant_safetensors(paths):
try:
total += os.path.getsize(path)
except OSError:
Expand Down Expand Up @@ -289,26 +298,41 @@ def keep_checkpoint_mapped(*, weight_bytes: int, component: str) -> bool:
return True


def _list_safetensors_files(model_path: str) -> list[str]:
"""List all .safetensors files under a directory.
def _select_safetensors_index_file(model_path: str, preferred_name: str) -> str | None:
preferred_path = os.path.join(str(model_path), preferred_name)
if os.path.exists(preferred_path):
return preferred_path

candidates = filter_duplicate_precision_variant_safetensors(
sorted(glob.glob(os.path.join(str(model_path), "*.safetensors.index.json")))
)
return candidates[0] if len(candidates) == 1 else None


def _list_safetensors_files(
model_path: str,
*,
index_file: str = _DEFAULT_SAFETENSORS_INDEX,
key_filter: Callable[[str], bool] | None = None,
raw_candidates: bool = False,
) -> list[str]:
"""Resolve the safetensors files to load from a local component path.

If a safetensors index file is present, verifies that every shard listed
in the index actually exists on disk. Missing shards are first repaired
automatically via HuggingFace Hub (if the path is an HF cache entry);
if repair fails a clear RuntimeError is raised.
An index is authoritative when present. Otherwise canonical files are
preferred over precision-suffixed copies. ``raw_candidates`` is reserved
for model-specific selectors that must choose a precision variant first.
"""
if os.path.isfile(model_path):
return [str(model_path)] if str(model_path).endswith(".safetensors") else []

found = sorted(glob.glob(os.path.join(str(model_path), "*.safetensors")))

index_path = os.path.join(
str(model_path), "diffusion_pytorch_model.safetensors.index.json"
)
if os.path.exists(index_path):
index_path = _select_safetensors_index_file(model_path, index_file)
if index_path is not None:
with open(index_path) as f:
index = json.load(f)
expected_shards = sorted(set(index.get("weight_map", {}).values()))
weight_map = index.get("weight_map", {})
expected_shards = sorted(set(weight_map.values()))
found_basenames = {os.path.basename(p) for p in found}
missing = [s for s in expected_shards if s not in found_basenames]
if missing:
Expand All @@ -325,24 +349,30 @@ def _list_safetensors_files(model_path: str) -> list[str]:
f"`huggingface-cli download {os.path.basename(model_path)}`)."
)

return found
if not raw_candidates:
selected_shards = {
shard
for weight_name, shard in weight_map.items()
if key_filter is None or key_filter(weight_name)
}
return [
os.path.join(str(model_path), shard)
for shard in sorted(selected_shards)
]

if raw_candidates:
return found
return filter_duplicate_precision_variant_safetensors(found)


def load_safetensors_state_dict(model_path: str) -> dict[str, torch.Tensor]:
"""Load one safetensors checkpoint, including an indexed sharded set."""
index_path = os.path.join(
str(model_path), "diffusion_pytorch_model.safetensors.index.json"
)
index_path = _select_safetensors_index_file(model_path, _DEFAULT_SAFETENSORS_INDEX)
safetensors_files = _list_safetensors_files(model_path)
if os.path.exists(index_path):
with open(index_path) as f:
index = json.load(f)
shard_names = sorted(set(index.get("weight_map", {}).values()))
if index_path is not None:
state_dict: dict[str, torch.Tensor] = {}
for shard_name in shard_names:
state_dict.update(
safetensors_load_file(os.path.join(str(model_path), shard_name))
)
for path in safetensors_files:
state_dict.update(safetensors_load_file(path))
return state_dict

if not safetensors_files:
Expand Down
Loading
Loading