[diffusion] feat: support single-file component weight overrides (#35979)
This commit is contained in:
@@ -150,6 +150,9 @@ for the current component matrix.
|
|||||||
|
|
||||||
Component overrides accept a local component directory, a standalone Hub
|
Component overrides accept a local component directory, a standalone Hub
|
||||||
repository, or a Hub component subfolder written as `owner/repo/subfolder`.
|
repository, or a Hub component subfolder written as `owner/repo/subfolder`.
|
||||||
|
For transformer and native encoder loaders, an explicit weight filename keeps
|
||||||
|
the base component configuration and replaces only its weights; the selected
|
||||||
|
file format must still be supported by that loader.
|
||||||
|
|
||||||
For supported realtime causal video models, `--kv-cache-quant {off|int4|int2}`
|
For supported realtime causal video models, `--kv-cache-quant {off|int4|int2}`
|
||||||
compresses completed KV-cache chunks independently of transformer weight
|
compresses completed KV-cache chunks independently of transformer weight
|
||||||
|
|||||||
@@ -16,7 +16,7 @@ precision are resolved independently.
|
|||||||
Use these paths:
|
Use these paths:
|
||||||
|
|
||||||
- `--model-path`: the base or original model
|
- `--model-path`: the base or original model
|
||||||
- `--component-paths.<component>` / `--<component>-path`: replace a component from `model_index.json` or a native registered module with an independent repo or local directory
|
- `--component-paths.<component>` / `--<component>-path`: replace a component from `model_index.json` or a native registered module with an independent repo/local directory; for transformer and native encoder loaders, an explicit weight filename retains the base component config
|
||||||
- `--transformer-path`: a quantized transformers-style transformer component directory that already contains its own `config.json`
|
- `--transformer-path`: a quantized transformers-style transformer component directory that already contains its own `config.json`
|
||||||
- `--transformer-weights-path`: replacement transformer weights in safetensors
|
- `--transformer-weights-path`: replacement transformer weights in safetensors
|
||||||
format (file, directory, or Hub repository/file) or a supported GGUF file
|
format (file, directory, or Hub repository/file) or a supported GGUF file
|
||||||
|
|||||||
+7
-1
@@ -23,6 +23,11 @@ class ImageEncoderLoader(TextEncoderLoader):
|
|||||||
component_name: str = "image_encoder",
|
component_name: str = "image_encoder",
|
||||||
):
|
):
|
||||||
"""Load the text encoders based on the model path, and inference args."""
|
"""Load the text encoders based on the model path, and inference args."""
|
||||||
|
component_weights_path = self.resolve_model_weights_path(
|
||||||
|
component_model_path,
|
||||||
|
server_args,
|
||||||
|
component_name,
|
||||||
|
)
|
||||||
# model_config: PretrainedConfig = get_hf_config(
|
# model_config: PretrainedConfig = get_hf_config(
|
||||||
# model=model_path,
|
# model=model_path,
|
||||||
# trust_remote_code=server_args.trust_remote_code,
|
# trust_remote_code=server_args.trust_remote_code,
|
||||||
@@ -39,6 +44,7 @@ class ImageEncoderLoader(TextEncoderLoader):
|
|||||||
encoder_config,
|
encoder_config,
|
||||||
model_config,
|
model_config,
|
||||||
component_model_path,
|
component_model_path,
|
||||||
|
component_weights_path,
|
||||||
component_name,
|
component_name,
|
||||||
)
|
)
|
||||||
# real dims are populated now; resolve fold vs replicate
|
# real dims are populated now; resolve fold vs replicate
|
||||||
@@ -50,7 +56,7 @@ class ImageEncoderLoader(TextEncoderLoader):
|
|||||||
# Always start with local device; load_model will adjust for offload if needed
|
# Always start with local device; load_model will adjust for offload if needed
|
||||||
# TODO(will): add support for other dtypes
|
# TODO(will): add support for other dtypes
|
||||||
return self.load_model(
|
return self.load_model(
|
||||||
component_model_path,
|
component_weights_path,
|
||||||
encoder_config,
|
encoder_config,
|
||||||
server_args,
|
server_args,
|
||||||
server_args.pipeline_config.image_encoder_precision,
|
server_args.pipeline_config.image_encoder_precision,
|
||||||
|
|||||||
+74
-6
@@ -62,7 +62,14 @@ from sglang.multimodal_gen.runtime.utils.hf_diffusers_utils import (
|
|||||||
load_dict,
|
load_dict,
|
||||||
)
|
)
|
||||||
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
||||||
from sglang.multimodal_gen.runtime.utils.quantization_utils import get_quant_config
|
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,
|
||||||
|
resolve_weight,
|
||||||
|
)
|
||||||
from sglang.multimodal_gen.utils import PRECISION_TO_TYPE
|
from sglang.multimodal_gen.utils import PRECISION_TO_TYPE
|
||||||
from sglang.srt.environ import envs
|
from sglang.srt.environ import envs
|
||||||
|
|
||||||
@@ -93,11 +100,29 @@ def _delegate_standard_bnb4_to_transformers(
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _get_encoder_quant_config(
|
||||||
|
component_config: dict,
|
||||||
|
component_model_path: str,
|
||||||
|
component_weights_path: str,
|
||||||
|
):
|
||||||
|
quant_config = get_quant_config(component_config, component_model_path)
|
||||||
|
if (
|
||||||
|
quant_config is None
|
||||||
|
and component_weights_path != component_model_path
|
||||||
|
and component_weights_path.endswith(".safetensors")
|
||||||
|
):
|
||||||
|
quant_config = get_quant_config_from_safetensors_metadata(
|
||||||
|
component_weights_path
|
||||||
|
)
|
||||||
|
return quant_config
|
||||||
|
|
||||||
|
|
||||||
def _configure_encoder_quantization(
|
def _configure_encoder_quantization(
|
||||||
model_config: EncoderConfig,
|
model_config: EncoderConfig,
|
||||||
model_cls: type[nn.Module],
|
model_cls: type[nn.Module],
|
||||||
component_config: dict,
|
component_config: dict,
|
||||||
component_model_path: str,
|
component_model_path: str,
|
||||||
|
component_weights_path: str,
|
||||||
component_name: str,
|
component_name: str,
|
||||||
) -> None:
|
) -> None:
|
||||||
if getattr(model_cls, "manages_checkpoint_quantization", False):
|
if getattr(model_cls, "manages_checkpoint_quantization", False):
|
||||||
@@ -111,9 +136,10 @@ def _configure_encoder_quantization(
|
|||||||
component_name,
|
component_name,
|
||||||
)
|
)
|
||||||
try:
|
try:
|
||||||
quant_config = get_quant_config(
|
quant_config = _get_encoder_quant_config(
|
||||||
component_config,
|
component_config,
|
||||||
component_model_path,
|
component_model_path,
|
||||||
|
component_weights_path,
|
||||||
)
|
)
|
||||||
except (KeyError, TypeError, ValueError) as error:
|
except (KeyError, TypeError, ValueError) as error:
|
||||||
raise ComponentCheckpointUnsupportedError(
|
raise ComponentCheckpointUnsupportedError(
|
||||||
@@ -134,6 +160,7 @@ def _resolve_and_configure_encoder_quantization(
|
|||||||
model_config: EncoderConfig,
|
model_config: EncoderConfig,
|
||||||
component_config: dict,
|
component_config: dict,
|
||||||
component_model_path: str,
|
component_model_path: str,
|
||||||
|
component_weights_path: str,
|
||||||
component_name: str,
|
component_name: str,
|
||||||
) -> type[nn.Module]:
|
) -> type[nn.Module]:
|
||||||
architectures = getattr(model_config, "architectures", [])
|
architectures = getattr(model_config, "architectures", [])
|
||||||
@@ -145,7 +172,11 @@ def _resolve_and_configure_encoder_quantization(
|
|||||||
component_name,
|
component_name,
|
||||||
)
|
)
|
||||||
try:
|
try:
|
||||||
quant_config = get_quant_config(component_config, component_model_path)
|
quant_config = _get_encoder_quant_config(
|
||||||
|
component_config,
|
||||||
|
component_model_path,
|
||||||
|
component_weights_path,
|
||||||
|
)
|
||||||
except Exception as quantization_error:
|
except Exception as quantization_error:
|
||||||
raise ComponentCheckpointUnsupportedError(
|
raise ComponentCheckpointUnsupportedError(
|
||||||
f"Cannot parse checkpoint quantization for {component_name!r}: "
|
f"Cannot parse checkpoint quantization for {component_name!r}: "
|
||||||
@@ -163,6 +194,7 @@ def _resolve_and_configure_encoder_quantization(
|
|||||||
model_cls,
|
model_cls,
|
||||||
component_config,
|
component_config,
|
||||||
component_model_path,
|
component_model_path,
|
||||||
|
component_weights_path,
|
||||||
component_name,
|
component_name,
|
||||||
)
|
)
|
||||||
return model_cls
|
return model_cls
|
||||||
@@ -242,6 +274,8 @@ def _require_quantized_encoder_layers(
|
|||||||
|
|
||||||
def _checkpoint_bytes(model_path: str) -> int:
|
def _checkpoint_bytes(model_path: str) -> int:
|
||||||
"""On-disk size of a checkpoint, readable before any weight of it is."""
|
"""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
|
total = 0
|
||||||
for path in glob.glob(
|
for path in glob.glob(
|
||||||
os.path.join(str(model_path), "**", "*.safetensors"), recursive=True
|
os.path.join(str(model_path), "**", "*.safetensors"), recursive=True
|
||||||
@@ -274,6 +308,23 @@ class TextEncoderLoader(ComponentLoader):
|
|||||||
component_names = ["text_encoder"]
|
component_names = ["text_encoder"]
|
||||||
expected_library = "transformers"
|
expected_library = "transformers"
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def resolve_model_weights_path(
|
||||||
|
component_model_path: str,
|
||||||
|
server_args: ServerArgs,
|
||||||
|
component_name: str,
|
||||||
|
) -> str:
|
||||||
|
weights_override = server_args.component_weights_paths.get(component_name)
|
||||||
|
if weights_override is None:
|
||||||
|
return component_model_path
|
||||||
|
model_weights_path = materialize_weight(resolve_weight(weights_override))
|
||||||
|
logger.info(
|
||||||
|
"Using weight-file override for %s: %s",
|
||||||
|
component_name,
|
||||||
|
model_weights_path,
|
||||||
|
)
|
||||||
|
return model_weights_path
|
||||||
|
|
||||||
@dataclasses.dataclass
|
@dataclasses.dataclass
|
||||||
class Source:
|
class Source:
|
||||||
"""A source for weights."""
|
"""A source for weights."""
|
||||||
@@ -320,8 +371,19 @@ class TextEncoderLoader(ComponentLoader):
|
|||||||
# model_name_or_path = (self._maybe_download_from_modelscope(
|
# model_name_or_path = (self._maybe_download_from_modelscope(
|
||||||
# model_name_or_path, revision) or model_name_or_path)
|
# model_name_or_path, revision) or model_name_or_path)
|
||||||
|
|
||||||
is_local = os.path.isdir(model_name_or_path)
|
if os.path.isfile(model_name_or_path):
|
||||||
assert is_local, "Model path must be a local directory"
|
if model_name_or_path.endswith(".safetensors"):
|
||||||
|
return os.path.dirname(model_name_or_path), [model_name_or_path], True
|
||||||
|
if fall_back_to_pt and model_name_or_path.endswith((".bin", ".pt")):
|
||||||
|
return os.path.dirname(model_name_or_path), [model_name_or_path], False
|
||||||
|
raise ValueError(
|
||||||
|
"Native encoder weight overrides currently support one "
|
||||||
|
f"safetensors, bin, or pt file, got {model_name_or_path!r}"
|
||||||
|
)
|
||||||
|
if not os.path.isdir(model_name_or_path):
|
||||||
|
raise ValueError(
|
||||||
|
f"Model path must be a local file or directory: {model_name_or_path!r}"
|
||||||
|
)
|
||||||
|
|
||||||
use_safetensors = False
|
use_safetensors = False
|
||||||
index_file = SAFE_WEIGHTS_INDEX_NAME
|
index_file = SAFE_WEIGHTS_INDEX_NAME
|
||||||
@@ -447,6 +509,11 @@ class TextEncoderLoader(ComponentLoader):
|
|||||||
component_starts_on_cpu: bool | None = None,
|
component_starts_on_cpu: bool | None = None,
|
||||||
):
|
):
|
||||||
"""Load the text encoders based on the model path, and inference args."""
|
"""Load the text encoders based on the model path, and inference args."""
|
||||||
|
component_weights_path = self.resolve_model_weights_path(
|
||||||
|
component_model_path,
|
||||||
|
server_args,
|
||||||
|
component_name,
|
||||||
|
)
|
||||||
diffusers_pretrained_config = get_config(
|
diffusers_pretrained_config = get_config(
|
||||||
component_model_path, trust_remote_code=True
|
component_model_path, trust_remote_code=True
|
||||||
)
|
)
|
||||||
@@ -478,6 +545,7 @@ class TextEncoderLoader(ComponentLoader):
|
|||||||
encoder_config,
|
encoder_config,
|
||||||
model_config,
|
model_config,
|
||||||
component_model_path,
|
component_model_path,
|
||||||
|
component_weights_path,
|
||||||
component_name,
|
component_name,
|
||||||
)
|
)
|
||||||
encoder_dp_group = get_encoder_data_parallel_group()
|
encoder_dp_group = get_encoder_data_parallel_group()
|
||||||
@@ -499,7 +567,7 @@ class TextEncoderLoader(ComponentLoader):
|
|||||||
]
|
]
|
||||||
# TODO(will): add support for other dtypes
|
# TODO(will): add support for other dtypes
|
||||||
return self.load_model(
|
return self.load_model(
|
||||||
component_model_path,
|
component_weights_path,
|
||||||
encoder_config,
|
encoder_config,
|
||||||
server_args,
|
server_args,
|
||||||
encoder_dtype,
|
encoder_dtype,
|
||||||
|
|||||||
@@ -99,15 +99,7 @@ def _server_args_for_transformer_component(
|
|||||||
server_args: ServerArgs, component_name: str
|
server_args: ServerArgs, component_name: str
|
||||||
) -> ServerArgs:
|
) -> ServerArgs:
|
||||||
"""Mask global quantized override flags for secondary transformer components."""
|
"""Mask global quantized override flags for secondary transformer components."""
|
||||||
if component_name not in ("transformer_2", "unconditional_transformer"):
|
component_weights_path = server_args.component_weights_paths.get(component_name)
|
||||||
return server_args
|
|
||||||
|
|
||||||
# Some pipelines have secondary DiT components with their own quantized
|
|
||||||
# weight file. Keep the mapping model-owned and the loader generic.
|
|
||||||
component_weights_paths = getattr(
|
|
||||||
server_args, "component_transformer_weights_paths", {}
|
|
||||||
)
|
|
||||||
component_weights_path = component_weights_paths.get(component_name)
|
|
||||||
if component_weights_path is not None:
|
if component_weights_path is not None:
|
||||||
component_server_args = copy.copy(server_args)
|
component_server_args = copy.copy(server_args)
|
||||||
component_server_args.transformer_weights_path = component_weights_path
|
component_server_args.transformer_weights_path = component_weights_path
|
||||||
@@ -119,6 +111,9 @@ def _server_args_for_transformer_component(
|
|||||||
)
|
)
|
||||||
return component_server_args
|
return component_server_args
|
||||||
|
|
||||||
|
if component_name not in ("transformer_2", "unconditional_transformer"):
|
||||||
|
return server_args
|
||||||
|
|
||||||
if (
|
if (
|
||||||
server_args.transformer_weights_path is None
|
server_args.transformer_weights_path is None
|
||||||
and server_args.nunchaku_config is None
|
and server_args.nunchaku_config is None
|
||||||
|
|||||||
@@ -244,25 +244,19 @@ class Ideogram4Nvfp4Pipeline(Ideogram4Pipeline):
|
|||||||
# The loader treats transformer_weights_path as the base DiT override.
|
# The loader treats transformer_weights_path as the base DiT override.
|
||||||
# Route the sibling unconditional DiT weights through the generic
|
# Route the sibling unconditional DiT weights through the generic
|
||||||
# per-component override map instead of hard-coding Ideogram there.
|
# per-component override map instead of hard-coding Ideogram there.
|
||||||
component_transformer_weights_paths = dict(
|
component_weights_paths = dict(server_args.component_weights_paths)
|
||||||
getattr(server_args, "component_transformer_weights_paths", {})
|
component_weights_paths.setdefault(
|
||||||
)
|
|
||||||
component_transformer_weights_paths.setdefault(
|
|
||||||
"unconditional_transformer",
|
"unconditional_transformer",
|
||||||
model_resolution.unconditional_transformer_weights_path,
|
model_resolution.unconditional_transformer_weights_path,
|
||||||
)
|
)
|
||||||
server_args.component_transformer_weights_paths = (
|
server_args.component_weights_paths = component_weights_paths
|
||||||
component_transformer_weights_paths
|
|
||||||
)
|
|
||||||
logger.info(
|
logger.info(
|
||||||
"NVFP4 transformer weights: %s",
|
"NVFP4 transformer weights: %s",
|
||||||
model_resolution.transformer_weights_path,
|
model_resolution.transformer_weights_path,
|
||||||
)
|
)
|
||||||
logger.info(
|
logger.info(
|
||||||
"NVFP4 unconditional transformer weights: %s",
|
"NVFP4 unconditional transformer weights: %s",
|
||||||
server_args.component_transformer_weights_paths.get(
|
server_args.component_weights_paths.get("unconditional_transformer"),
|
||||||
"unconditional_transformer"
|
|
||||||
),
|
|
||||||
)
|
)
|
||||||
return super().load_modules(server_args, loaded_modules)
|
return super().load_modules(server_args, loaded_modules)
|
||||||
|
|
||||||
|
|||||||
@@ -77,6 +77,9 @@ from sglang.multimodal_gen.runtime.utils.logging_utils import (
|
|||||||
configure_logger,
|
configure_logger,
|
||||||
init_logger,
|
init_logger,
|
||||||
)
|
)
|
||||||
|
from sglang.multimodal_gen.runtime.weights.source import (
|
||||||
|
is_explicit_weight_file_reference,
|
||||||
|
)
|
||||||
from sglang.multimodal_gen.utils import (
|
from sglang.multimodal_gen.utils import (
|
||||||
FlexibleArgumentParser,
|
FlexibleArgumentParser,
|
||||||
StoreBoolean,
|
StoreBoolean,
|
||||||
@@ -301,6 +304,8 @@ class ServerArgs(DisaggServerArgsMixin):
|
|||||||
|
|
||||||
# Component path overrides (key = model_index.json component name, value = path)
|
# Component path overrides (key = model_index.json component name, value = path)
|
||||||
component_paths: dict[str, str] = field(default_factory=dict)
|
component_paths: dict[str, str] = field(default_factory=dict)
|
||||||
|
# Exact weight-file overrides retain the base component configuration.
|
||||||
|
component_weights_paths: dict[str, str] = field(default_factory=dict)
|
||||||
# Optional LTX-2.5 decoder is large enough to load only when requested.
|
# Optional LTX-2.5 decoder is large enough to load only when requested.
|
||||||
load_diffusion_decoder: bool = False
|
load_diffusion_decoder: bool = False
|
||||||
|
|
||||||
@@ -313,12 +318,6 @@ class ServerArgs(DisaggServerArgsMixin):
|
|||||||
# Widest timestep plan the rebuild slab is sized for; see
|
# Widest timestep plan the rebuild slab is sized for; see
|
||||||
# MINIMAX_H3_ADALN_MAX_PLAN_WIDTH.
|
# MINIMAX_H3_ADALN_MAX_PLAN_WIDTH.
|
||||||
minimax_h3_adaln_plan_width: int = 4
|
minimax_h3_adaln_plan_width: int = 4
|
||||||
# Per-component transformer weight overrides (key = model_index.json component name).
|
|
||||||
# Pipelines use this when a checkpoint ships separate quantized weights for
|
|
||||||
# secondary DiT components; the generic loader consumes it without model-specific
|
|
||||||
# filename logic.
|
|
||||||
component_transformer_weights_paths: dict[str, str] = field(default_factory=dict)
|
|
||||||
|
|
||||||
# Explicit quantization method override (e.g. "mxfp8", "fp8", "modelslim").
|
# Explicit quantization method override (e.g. "mxfp8", "fp8", "modelslim").
|
||||||
# When set, the transformer loader uses it instead of auto-detection.
|
# When set, the transformer loader uses it instead of auto-detection.
|
||||||
quantization: str | None = None
|
quantization: str | None = None
|
||||||
@@ -1699,6 +1698,30 @@ class ServerArgs(DisaggServerArgsMixin):
|
|||||||
# configure logger before use
|
# configure logger before use
|
||||||
configure_logger(server_args=self)
|
configure_logger(server_args=self)
|
||||||
|
|
||||||
|
component_paths: dict[str, str] = {}
|
||||||
|
component_weights_paths = dict(self.component_weights_paths)
|
||||||
|
for component, path in self.component_paths.items():
|
||||||
|
supports_weight_file_override = (
|
||||||
|
is_dit_component_name(component)
|
||||||
|
or is_text_encoder_component_name(component)
|
||||||
|
or is_image_encoder_component_name(component)
|
||||||
|
)
|
||||||
|
if (
|
||||||
|
not supports_weight_file_override
|
||||||
|
or not is_explicit_weight_file_reference(path)
|
||||||
|
):
|
||||||
|
component_paths[component] = path
|
||||||
|
continue
|
||||||
|
existing = component_weights_paths.get(component)
|
||||||
|
if existing is not None and existing != path:
|
||||||
|
raise ValueError(
|
||||||
|
f"Conflicting weight overrides for component {component!r}: "
|
||||||
|
f"{existing!r} and {path!r}"
|
||||||
|
)
|
||||||
|
component_weights_paths[component] = path
|
||||||
|
self.component_paths = component_paths
|
||||||
|
self.component_weights_paths = component_weights_paths
|
||||||
|
|
||||||
# Convert string disagg_role to enum (from CLI/config)
|
# Convert string disagg_role to enum (from CLI/config)
|
||||||
if isinstance(self.disagg_role, str):
|
if isinstance(self.disagg_role, str):
|
||||||
self.disagg_role = RoleType.from_string(self.disagg_role)
|
self.disagg_role = RoleType.from_string(self.disagg_role)
|
||||||
|
|||||||
@@ -8,7 +8,7 @@ from pathlib import Path, PurePosixPath
|
|||||||
from typing import Literal
|
from typing import Literal
|
||||||
from urllib.parse import unquote, urlparse
|
from urllib.parse import unquote, urlparse
|
||||||
|
|
||||||
from huggingface_hub import HfApi
|
from huggingface_hub import HfApi, hf_hub_download
|
||||||
from huggingface_hub.utils import validate_repo_id
|
from huggingface_hub.utils import validate_repo_id
|
||||||
|
|
||||||
WeightSourceKind = Literal["local", "huggingface"]
|
WeightSourceKind = Literal["local", "huggingface"]
|
||||||
@@ -40,6 +40,14 @@ class ResolvedWeight:
|
|||||||
selected_file: str
|
selected_file: str
|
||||||
|
|
||||||
|
|
||||||
|
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)
|
||||||
|
if os.path.isdir(expanded):
|
||||||
|
return False
|
||||||
|
return urlparse(source).path.lower().endswith(_WEIGHT_SUFFIXES)
|
||||||
|
|
||||||
|
|
||||||
def _validate_relative_hub_path(path: str, field_name: str) -> str:
|
def _validate_relative_hub_path(path: str, field_name: str) -> str:
|
||||||
normalized = str(PurePosixPath(path))
|
normalized = str(PurePosixPath(path))
|
||||||
pure_path = PurePosixPath(normalized)
|
pure_path = PurePosixPath(normalized)
|
||||||
@@ -264,3 +272,20 @@ def resolve_weight(
|
|||||||
inventory=inventory,
|
inventory=inventory,
|
||||||
selected_file=selected_file,
|
selected_file=selected_file,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def materialize_weight(resolved: ResolvedWeight) -> str:
|
||||||
|
"""Return the selected local file, downloading one pinned Hub file if needed."""
|
||||||
|
source = resolved.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)
|
||||||
|
|
||||||
|
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,
|
||||||
|
)
|
||||||
|
|||||||
@@ -509,7 +509,7 @@ class TestIdeogram4(unittest.TestCase):
|
|||||||
server_args = SimpleNamespace(
|
server_args = SimpleNamespace(
|
||||||
transformer_weights_path="/unused/override.safetensors",
|
transformer_weights_path="/unused/override.safetensors",
|
||||||
nunchaku_config={"enabled": True},
|
nunchaku_config={"enabled": True},
|
||||||
component_transformer_weights_paths={},
|
component_weights_paths={},
|
||||||
)
|
)
|
||||||
component_args = _server_args_for_transformer_component(
|
component_args = _server_args_for_transformer_component(
|
||||||
server_args, "unconditional_transformer"
|
server_args, "unconditional_transformer"
|
||||||
@@ -524,7 +524,7 @@ class TestIdeogram4(unittest.TestCase):
|
|||||||
"/ckpt/diffusion_models/ideogram4_nvfp4_mixed.safetensors"
|
"/ckpt/diffusion_models/ideogram4_nvfp4_mixed.safetensors"
|
||||||
),
|
),
|
||||||
nunchaku_config={"enabled": True},
|
nunchaku_config={"enabled": True},
|
||||||
component_transformer_weights_paths={
|
component_weights_paths={
|
||||||
"unconditional_transformer": (
|
"unconditional_transformer": (
|
||||||
"/ckpt/diffusion_models/"
|
"/ckpt/diffusion_models/"
|
||||||
"ideogram4_unconditional_nvfp4_mixed.safetensors"
|
"ideogram4_unconditional_nvfp4_mixed.safetensors"
|
||||||
|
|||||||
@@ -27,6 +27,7 @@ class TestImageEncoderQuantizationAdmission(unittest.TestCase):
|
|||||||
image_encoder_precision="bf16",
|
image_encoder_precision="bf16",
|
||||||
native_only_components=(),
|
native_only_components=(),
|
||||||
),
|
),
|
||||||
|
component_weights_paths={},
|
||||||
encoder_parallel="replicate",
|
encoder_parallel="replicate",
|
||||||
resolve_component_attention_backend=lambda _name: (None, None),
|
resolve_component_attention_backend=lambda _name: (None, None),
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -163,6 +163,23 @@ class TestServerArgsPathExpansion(unittest.TestCase):
|
|||||||
args.component_paths["vae"], os.path.expanduser("~/fake/local/vae")
|
args.component_paths["vae"], os.path.expanduser("~/fake/local/vae")
|
||||||
)
|
)
|
||||||
|
|
||||||
|
def test_component_weight_file_keeps_base_component_config(self):
|
||||||
|
args = self._from_dict_without_model_resolution(
|
||||||
|
{
|
||||||
|
"model_path": "/data/my-model",
|
||||||
|
"component_paths": {
|
||||||
|
"text_encoder": "owner/repo/text_encoder/model.safetensors",
|
||||||
|
"vae": "owner/repo/vae",
|
||||||
|
},
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(args.component_paths, {"vae": "owner/repo/vae"})
|
||||||
|
self.assertEqual(
|
||||||
|
args.component_weights_paths,
|
||||||
|
{"text_encoder": "owner/repo/text_encoder/model.safetensors"},
|
||||||
|
)
|
||||||
|
|
||||||
def test_component_attention_backends_are_normalized(self):
|
def test_component_attention_backends_are_normalized(self):
|
||||||
args = self._from_dict_without_model_resolution(
|
args = self._from_dict_without_model_resolution(
|
||||||
{
|
{
|
||||||
|
|||||||
@@ -187,10 +187,31 @@ class TestTextEncoderQuantization(unittest.TestCase):
|
|||||||
TextEncoder,
|
TextEncoder,
|
||||||
{},
|
{},
|
||||||
"/model/text_encoder",
|
"/model/text_encoder",
|
||||||
|
"/model/text_encoder",
|
||||||
"text_encoder",
|
"text_encoder",
|
||||||
)
|
)
|
||||||
self.assertIs(model_config.quant_config, self.serialized)
|
self.assertIs(model_config.quant_config, self.serialized)
|
||||||
|
|
||||||
|
def test_weight_file_metadata_configures_native_encoder(self):
|
||||||
|
model_config = SimpleNamespace(quant_config=None)
|
||||||
|
self.get_quant_config.return_value = None
|
||||||
|
with mock.patch(
|
||||||
|
"sglang.multimodal_gen.runtime.loader.component_loaders."
|
||||||
|
"text_encoder_loader.get_quant_config_from_safetensors_metadata",
|
||||||
|
return_value=self.serialized,
|
||||||
|
) as get_file_quant_config:
|
||||||
|
_configure_encoder_quantization(
|
||||||
|
model_config,
|
||||||
|
TextEncoder,
|
||||||
|
{},
|
||||||
|
"/model/text_encoder",
|
||||||
|
"/weights/encoder.safetensors",
|
||||||
|
"text_encoder",
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertIs(model_config.quant_config, self.serialized)
|
||||||
|
get_file_quant_config.assert_called_once_with("/weights/encoder.safetensors")
|
||||||
|
|
||||||
def test_encoder_must_use_native_loader(self):
|
def test_encoder_must_use_native_loader(self):
|
||||||
model_config = SimpleNamespace(quant_config=None)
|
model_config = SimpleNamespace(quant_config=None)
|
||||||
with self.assertRaisesRegex(
|
with self.assertRaisesRegex(
|
||||||
@@ -201,6 +222,7 @@ class TestTextEncoderQuantization(unittest.TestCase):
|
|||||||
nn.Module,
|
nn.Module,
|
||||||
{},
|
{},
|
||||||
"/model/text_encoder",
|
"/model/text_encoder",
|
||||||
|
"/model/text_encoder",
|
||||||
"text_encoder",
|
"text_encoder",
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -225,6 +247,7 @@ class TestTextEncoderQuantization(unittest.TestCase):
|
|||||||
SimpleNamespace(architectures=[architecture], quant_config=None),
|
SimpleNamespace(architectures=[architecture], quant_config=None),
|
||||||
component_config,
|
component_config,
|
||||||
"/model/text_encoder",
|
"/model/text_encoder",
|
||||||
|
"/model/text_encoder",
|
||||||
"text_encoder",
|
"text_encoder",
|
||||||
)
|
)
|
||||||
self.get_quant_config.assert_not_called()
|
self.get_quant_config.assert_not_called()
|
||||||
@@ -244,6 +267,7 @@ class TestTextEncoderQuantization(unittest.TestCase):
|
|||||||
}
|
}
|
||||||
},
|
},
|
||||||
"/model/text_encoder",
|
"/model/text_encoder",
|
||||||
|
"/model/text_encoder",
|
||||||
"text_encoder",
|
"text_encoder",
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -264,6 +288,7 @@ class TestTextEncoderQuantization(unittest.TestCase):
|
|||||||
}
|
}
|
||||||
},
|
},
|
||||||
"/model/text_encoder",
|
"/model/text_encoder",
|
||||||
|
"/model/text_encoder",
|
||||||
"text_encoder",
|
"text_encoder",
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -284,6 +309,7 @@ class TestTextEncoderQuantization(unittest.TestCase):
|
|||||||
}
|
}
|
||||||
},
|
},
|
||||||
"/model/text_encoder",
|
"/model/text_encoder",
|
||||||
|
"/model/text_encoder",
|
||||||
"text_encoder",
|
"text_encoder",
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|||||||
@@ -4,12 +4,25 @@ from unittest.mock import patch
|
|||||||
import pytest
|
import pytest
|
||||||
|
|
||||||
from sglang.multimodal_gen.runtime.weights.source import (
|
from sglang.multimodal_gen.runtime.weights.source import (
|
||||||
|
is_explicit_weight_file_reference,
|
||||||
|
materialize_weight,
|
||||||
parse_weight_source,
|
parse_weight_source,
|
||||||
resolve_weight,
|
resolve_weight,
|
||||||
resolve_weight_inventory,
|
resolve_weight_inventory,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def test_explicit_weight_file_reference_does_not_claim_directories(tmp_path):
|
||||||
|
component = tmp_path / "component.safetensors"
|
||||||
|
component.mkdir()
|
||||||
|
|
||||||
|
assert not is_explicit_weight_file_reference(str(component))
|
||||||
|
assert is_explicit_weight_file_reference("owner/repo/model.safetensors")
|
||||||
|
assert is_explicit_weight_file_reference(
|
||||||
|
"https://huggingface.co/owner/repo/resolve/main/model.gguf?download=true"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
def test_parse_weight_source_accepts_repo_subfolder_and_exact_url():
|
def test_parse_weight_source_accepts_repo_subfolder_and_exact_url():
|
||||||
subfolder = parse_weight_source("owner/repo/text_encoder", revision="v1")
|
subfolder = parse_weight_source("owner/repo/text_encoder", revision="v1")
|
||||||
repo_file = parse_weight_source("owner/repo/adapter.safetensors")
|
repo_file = parse_weight_source("owner/repo/adapter.safetensors")
|
||||||
@@ -80,3 +93,10 @@ def test_weight_source_rejects_ambiguous_files(tmp_path):
|
|||||||
|
|
||||||
with pytest.raises(ValueError, match="multiple independent weight files"):
|
with pytest.raises(ValueError, match="multiple independent weight files"):
|
||||||
resolve_weight(str(tmp_path))
|
resolve_weight(str(tmp_path))
|
||||||
|
|
||||||
|
|
||||||
|
def test_materialize_local_weight_returns_selected_file(tmp_path):
|
||||||
|
checkpoint = tmp_path / "model.safetensors"
|
||||||
|
checkpoint.write_bytes(b"fixture")
|
||||||
|
|
||||||
|
assert materialize_weight(resolve_weight(str(checkpoint))) == str(checkpoint)
|
||||||
|
|||||||
Reference in New Issue
Block a user