[diffusion] chore: minor cleanups (#19123)

This commit is contained in:
Mick
2026-02-22 19:07:25 +08:00
committed by GitHub
parent 1a1c768d44
commit 87823722b3
6 changed files with 128 additions and 113 deletions
@@ -244,7 +244,7 @@ class MinimalA2AAttnOp(DistributedAttention):
SparseLinearAttentionBackend, SparseLinearAttentionBackend,
SageSparseLinearAttentionBackend, SageSparseLinearAttentionBackend,
): ):
logger.warning( logger.warning_once(
"TurboWan now only supports `sla_attn` or `sage_sla_attn` and has been automatically set to attention_type. Please set --attention-backend to `sla_attn` or `sage_sla_attn`." "TurboWan now only supports `sla_attn` or `sage_sla_attn` and has been automatically set to attention_type. Please set --attention-backend to `sla_attn` or `sage_sla_attn`."
) )
if attention_type == "sagesla": if attention_type == "sagesla":
@@ -1,19 +1,17 @@
import json
import os
from sglang.multimodal_gen.configs.models import ModelConfig from sglang.multimodal_gen.configs.models import ModelConfig
from sglang.multimodal_gen.runtime.loader.component_loaders.text_encoder_loader import ( from sglang.multimodal_gen.runtime.loader.component_loaders.text_encoder_loader import (
TextEncoderLoader, TextEncoderLoader,
) )
from sglang.multimodal_gen.runtime.loader.utils import _clean_hf_config_inplace
from sglang.multimodal_gen.runtime.server_args import ServerArgs from sglang.multimodal_gen.runtime.server_args import ServerArgs
from sglang.multimodal_gen.runtime.utils.hf_diffusers_utils import (
get_diffusers_component_config,
)
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
logger = init_logger(__name__) logger = init_logger(__name__)
class ImageEncoderLoader(TextEncoderLoader): class ImageEncoderLoader(TextEncoderLoader):
component_names = ["image_encoder"] component_names = ["image_encoder"]
expected_library = "transformers" expected_library = "transformers"
@@ -41,10 +39,9 @@ class ImageEncoderLoader(TextEncoderLoader):
# revision=server_args.revision, # revision=server_args.revision,
# model_override_args=None, # model_override_args=None,
# ) # )
with open(os.path.join(component_model_path, "config.json")) as f: model_config = get_diffusers_component_config(
model_config = json.load(f) component_path=component_model_path
_clean_hf_config_inplace(model_config) )
logger.debug("HF model config: %s", model_config)
encoder_config = server_args.pipeline_config.image_encoder_config encoder_config = server_args.pipeline_config.image_encoder_config
encoder_config.update_model_arch(model_config) encoder_config.update_model_arch(model_config)
@@ -21,7 +21,6 @@ from sglang.multimodal_gen.runtime.loader.component_loaders.component_loader imp
) )
from sglang.multimodal_gen.runtime.loader.fsdp_load import shard_model from sglang.multimodal_gen.runtime.loader.fsdp_load import shard_model
from sglang.multimodal_gen.runtime.loader.utils import ( from sglang.multimodal_gen.runtime.loader.utils import (
_clean_hf_config_inplace,
set_default_torch_dtype, set_default_torch_dtype,
skip_init_modules, skip_init_modules,
) )
@@ -173,20 +172,12 @@ class TextEncoderLoader(ComponentLoader):
self, component_model_path: str, server_args: ServerArgs, component_name: str self, component_model_path: str, server_args: ServerArgs, component_name: str
): ):
"""Load the text encoders based on the model path, and inference args.""" """Load the text encoders based on the model path, and inference args."""
# model_config: PretrainedConfig = get_hf_config(
# model=model_path,
# trust_remote_code=server_args.trust_remote_code,
# revision=server_args.revision,
# model_override_args=None,
# )
diffusers_pretrained_config = get_config( diffusers_pretrained_config = get_config(
component_model_path, trust_remote_code=True component_model_path, trust_remote_code=True
) )
model_config = get_diffusers_component_config( model_config = get_diffusers_component_config(
component_path=component_model_path component_path=component_model_path
) )
_clean_hf_config_inplace(model_config)
logger.debug("HF model config: %s", model_config)
def is_not_first_encoder(module_name): def is_not_first_encoder(module_name):
return "2" in module_name return "2" in module_name
@@ -135,22 +135,21 @@ class GPUWorker:
# otherwise empty offloaded weights could fail lora converting # otherwise empty offloaded weights could fail lora converting
if self.server_args.dit_layerwise_offload: if self.server_args.dit_layerwise_offload:
# enable layerwise offload if possible # enable layerwise offload if possible
for dit in filter( for module_name in [
None, "transformer",
[ "transformer_2",
self.pipeline.get_module("transformer"), "video_dit",
self.pipeline.get_module("transformer_2"), "video_dit_2",
self.pipeline.get_module("video_dit"), "audio_dit",
self.pipeline.get_module("video_dit_2"), ]:
self.pipeline.get_module("audio_dit"), dit = self.pipeline.get_module(module_name)
], if dit:
): if isinstance(dit, OffloadableDiTMixin):
if isinstance(dit, OffloadableDiTMixin): dit.configure_layerwise_offload(self.server_args)
dit.configure_layerwise_offload(self.server_args) else:
else: logger.info(
logger.info( f"Module {type(dit).__name__} does not support layerwise offload. Skipping."
f"Module {type(dit).__name__} does not support layerwise offload. Skipping." )
)
logger.info( logger.info(
f"Worker {self.rank}: Initialized device, model, and distributed environment." f"Worker {self.rank}: Initialized device, model, and distributed environment."
@@ -199,7 +198,7 @@ class GPUWorker:
logger.info( logger.info(
f"Peak GPU memory: {peak_reserved_gb:.2f} GB, " f"Peak GPU memory: {peak_reserved_gb:.2f} GB, "
f"Peak allocated: {peak_allocated_gb:.2f} GB, " f"Peak allocated: {peak_allocated_gb:.2f} GB, "
f"Memory pool overhead: {pool_overhead_gb:.2f} GB ({pool_overhead_gb/peak_reserved_gb*100:.1f}%), " f"Memory pool overhead: {pool_overhead_gb:.2f} GB ({pool_overhead_gb / peak_reserved_gb * 100:.1f}%), "
f"Remaining GPU memory at peak: {remaining_gpu_mem_gb:.2f} GB. " f"Remaining GPU memory at peak: {remaining_gpu_mem_gb:.2f} GB. "
f"Components that could stay resident (based on the last request workload): {can_stay_resident}. " f"Components that could stay resident (based on the last request workload): {can_stay_resident}. "
f"Related offload server args to disable: {suggested_args_str}" f"Related offload server args to disable: {suggested_args_str}"
@@ -40,12 +40,12 @@ from requests.exceptions import ConnectionError as RequestsConnectionError
from requests.exceptions import RequestException from requests.exceptions import RequestException
from safetensors import safe_open from safetensors import safe_open
from transformers import AutoConfig, PretrainedConfig from transformers import AutoConfig, PretrainedConfig
from transformers.models.auto.modeling_auto import MODEL_FOR_CAUSAL_LM_MAPPING_NAMES
from sglang.multimodal_gen.runtime.layers.quantization import ( from sglang.multimodal_gen.runtime.layers.quantization import (
QuantizationConfig, QuantizationConfig,
get_quantization_config, get_quantization_config,
) )
from sglang.multimodal_gen.runtime.loader.utils import _clean_hf_config_inplace
from sglang.multimodal_gen.runtime.loader.weight_utils import get_lock from sglang.multimodal_gen.runtime.loader.weight_utils import get_lock
from sglang.multimodal_gen.runtime.platforms import current_platform from sglang.multimodal_gen.runtime.platforms import current_platform
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
@@ -209,6 +209,34 @@ def _ci_validate_diffusers_model(model_path: str) -> tuple[bool, bool]:
return True, False return True, False
def _verify_diffusers_model_complete(path: str) -> bool:
"""Check if a diffusers model directory has all required component subdirectories."""
config_path = os.path.join(path, "model_index.json")
if not os.path.exists(config_path):
return False
try:
with open(config_path) as config_file:
model_index = json.load(config_file)
except Exception as exc:
logger.warning("Failed to read model_index.json at %s: %s", config_path, exc)
return False
component_keys = [
key
for key, value in model_index.items()
if isinstance(value, (list, tuple))
and len(value) == 2
and all(isinstance(item, str) for item in value)
]
if component_keys:
return all(os.path.exists(os.path.join(path, key)) for key in component_keys)
return os.path.exists(os.path.join(path, "transformer")) and os.path.exists(
os.path.join(path, "vae")
)
_CONFIG_REGISTRY: dict[str, type[PretrainedConfig]] = { _CONFIG_REGISTRY: dict[str, type[PretrainedConfig]] = {
# ChatGLMConfig.model_type: ChatGLMConfig, # ChatGLMConfig.model_type: ChatGLMConfig,
# DbrxConfig.model_type: DbrxConfig, # DbrxConfig.model_type: DbrxConfig,
@@ -235,8 +263,7 @@ def get_hf_config(
model_override_args: dict | None = None, model_override_args: dict | None = None,
**kwargs, **kwargs,
) -> PretrainedConfig: ) -> PretrainedConfig:
is_gguf = check_gguf_file(component_model_path) if check_gguf_file(component_model_path):
if is_gguf:
raise NotImplementedError("GGUF models are not supported.") raise NotImplementedError("GGUF models are not supported.")
config = AutoConfig.from_pretrained( config = AutoConfig.from_pretrained(
@@ -253,13 +280,6 @@ def get_hf_config(
if model_override_args: if model_override_args:
config.update(model_override_args) config.update(model_override_args)
# Special architecture mapping check for GGUF models
if is_gguf:
if config.model_type not in MODEL_FOR_CAUSAL_LM_MAPPING_NAMES:
raise RuntimeError(f"Can't get gguf config for {config.model_type}.")
model_type = MODEL_FOR_CAUSAL_LM_MAPPING_NAMES[config.model_type]
config.update({"architectures": [model_type]})
return config return config
@@ -270,14 +290,9 @@ def get_config(
model_override_args: Optional[dict] = None, model_override_args: Optional[dict] = None,
**kwargs, **kwargs,
): ):
try: return AutoConfig.from_pretrained(
config = AutoConfig.from_pretrained( model, trust_remote_code=trust_remote_code, revision=revision, **kwargs
model, trust_remote_code=trust_remote_code, revision=revision, **kwargs )
)
except ValueError as e:
raise e
return config
def load_dict(file_path): def load_dict(file_path):
@@ -305,7 +320,6 @@ def get_diffusers_component_config(
if not os.path.exists(component_path): if not os.path.exists(component_path):
component_path = maybe_download_model(component_path) component_path = maybe_download_model(component_path)
# tokenizer
config_names = ["generation_config.json"] config_names = ["generation_config.json"]
# By default, we load config.json, but scheduler_config.json for scheduler # By default, we load config.json, but scheduler_config.json for scheduler
if "scheduler" in component_path: if "scheduler" in component_path:
@@ -321,6 +335,10 @@ def get_diffusers_component_config(
lambda acc, path: acc | load_dict(path), config_file_paths, {} lambda acc, path: acc | load_dict(path), config_file_paths, {}
) )
_clean_hf_config_inplace(combined_config)
logger.debug("HF model config: %s", combined_config)
return combined_config return combined_config
@@ -414,9 +432,9 @@ CONTEXT_LENGTH_KEYS = [
def attach_additional_stop_token_ids(tokenizer): def attach_additional_stop_token_ids(tokenizer):
# Special handling for stop token <|eom_id|> generated by llama 3 tool use. # Special handling for stop token <|eom_id|> generated by llama 3 tool use.
if "<|eom_id|>" in tokenizer.get_added_vocab(): if "<|eom_id|>" in tokenizer.get_added_vocab():
tokenizer.additional_stop_token_ids = set( tokenizer.additional_stop_token_ids = {
[tokenizer.get_added_vocab()["<|eom_id|>"]] tokenizer.get_added_vocab()["<|eom_id|>"]
) }
else: else:
tokenizer.additional_stop_token_ids = None tokenizer.additional_stop_token_ids = None
@@ -635,45 +653,11 @@ def maybe_download_model(
Local path to the model Local path to the model
""" """
def _verify_diffusers_model_complete(path: str) -> bool:
"""Check if model directory (of a diffusers model, not a component) has required subdirectories."""
config_path = os.path.join(path, "model_index.json")
if not os.path.exists(config_path):
return False
try:
with open(config_path) as config_file:
model_index = json.load(config_file)
except Exception as exc:
logger.warning(
"Failed to read model_index.json at %s: %s", config_path, exc
)
return False
component_keys = [
key
for key, value in model_index.items()
if isinstance(value, (list, tuple))
and len(value) == 2
and all(isinstance(item, str) for item in value)
]
if component_keys:
return all(
os.path.exists(os.path.join(path, component_key))
for component_key in component_keys
)
transformer_dir = os.path.join(path, "transformer")
vae_dir = os.path.join(path, "vae")
return os.path.exists(transformer_dir) and os.path.exists(vae_dir)
# 1. Local path check: if path exists locally, verify it's complete (skip for LoRA) # 1. Local path check: if path exists locally, verify it's complete (skip for LoRA)
if os.path.exists(model_name_or_path): if os.path.exists(model_name_or_path):
# TODO: lots of duplication here
if not force_diffusers_model: if not force_diffusers_model:
return model_name_or_path return model_name_or_path
elif is_lora or _verify_diffusers_model_complete(model_name_or_path): if is_lora or _verify_diffusers_model_complete(model_name_or_path):
# CI validation: check all subdirectories for missing shards
if not is_lora: if not is_lora:
is_valid, cleanup_performed = _ci_validate_diffusers_model( is_valid, cleanup_performed = _ci_validate_diffusers_model(
model_name_or_path model_name_or_path
@@ -687,8 +671,6 @@ def maybe_download_model(
) )
# Fall through to download # Fall through to download
else: else:
# Local path is not in HF cache structure, can't clean up
# Raise error since we can't fix this automatically
raise ValueError( raise ValueError(
f"CI validation failed for local model at {model_name_or_path}. " f"CI validation failed for local model at {model_name_or_path}. "
"Some safetensors shards are missing. " "Some safetensors shards are missing. "
@@ -722,25 +704,21 @@ def maybe_download_model(
) )
if not force_diffusers_model: if not force_diffusers_model:
return str(local_path) return str(local_path)
elif is_lora or _verify_diffusers_model_complete(local_path): if is_lora or _verify_diffusers_model_complete(local_path):
# CI validation: check all subdirectories for missing shards
if not is_lora: if not is_lora:
is_valid, cleanup_performed = _ci_validate_diffusers_model(local_path) is_valid, cleanup_performed = _ci_validate_diffusers_model(local_path)
if not is_valid: if not is_valid:
if cleanup_performed: logger.warning(
logger.warning( "CI validation failed for cached model at %s, "
"CI validation failed for cached model at %s, " "%s, will re-download",
"cache has been cleaned up, will re-download", local_path,
local_path, (
) "cache has been cleaned up"
# Fall through to download if cleanup_performed
else: else "cleanup was not performed"
# This shouldn't happen for HF cache paths, but handle it ),
logger.warning( )
"CI validation failed for cached model at %s, " # Fall through to download
"but cleanup was not performed, will attempt re-download",
local_path,
)
else: else:
logger.info("Found complete model in cache at %s", local_path) logger.info("Found complete model in cache at %s", local_path)
return str(local_path) return str(local_path)
@@ -748,7 +726,6 @@ def maybe_download_model(
logger.info("Found complete model in cache at %s", local_path) logger.info("Found complete model in cache at %s", local_path)
return str(local_path) return str(local_path)
else: else:
# Model found in cache but incomplete
if not download: if not download:
raise ValueError( raise ValueError(
f"Model {model_name_or_path} found in cache but is incomplete and download=False." f"Model {model_name_or_path} found in cache but is incomplete and download=False."
@@ -930,3 +907,54 @@ def get_metadata_from_safetensors_file(file_path: str):
return metadata return metadata
except Exception as e: except Exception as e:
logger.warning(e) logger.warning(e)
def get_quant_config_from_safetensors_metadata(
file_path: str,
) -> Optional[QuantizationConfig]:
"""Extract quantization config from a safetensors file's metadata header.
Safetensors files can embed a flat string→string metadata dict in their header.
We expect a ``quantization_config`` key containing a JSON-encoded dict with at
least a ``quant_method`` field (e.g. ``"fp8"``), matching the format written by
``convert_hf_to_fp8.py`` when embedded into a config.json.
Returns None if no recognizable quantization metadata is found.
"""
metadata = get_metadata_from_safetensors_file(file_path)
if not metadata:
return None
quant_config_str = metadata.get("quantization_config")
if not quant_config_str:
return None
try:
quant_config_dict = json.loads(quant_config_str)
except Exception as e:
logger.warning(
"failed to parse quantization_config from safetensors metadata: %s", e
)
return None
quant_method = quant_config_dict.get("quant_method")
if not quant_method:
logger.warning(
"quantization_config in safetensors metadata is missing 'quant_method'"
)
return None
try:
quant_cls = get_quantization_config(quant_method)
config = quant_cls.from_config(quant_config_dict)
logger.info(
"loaded quantization config (%s) from safetensors metadata: %s",
quant_method,
file_path,
)
return config
except Exception as e:
logger.warning(
"failed to build QuantizationConfig from safetensors metadata: %s", e
)
return None
@@ -115,7 +115,7 @@ class ConversionResult:
with self.lock: with self.lock:
for k, v in q_weights.items(): for k, v in q_weights.items():
self.weight_map[k] = filename self.weight_map[k] = filename
self.param_count += len(v) self.param_count += v.numel()
self.modules_to_not_convert.extend(module_names) self.modules_to_not_convert.extend(module_names)