[diffusion] CI: add CI validation for diffusion model downloads (#16311)

This commit is contained in:
Alison Shao
2026-01-03 14:12:45 +08:00
committed by GitHub
parent d8ac5eecf7
commit 5b4f790200
@@ -19,8 +19,10 @@
"""Utilities for Huggingface Transformers.""" """Utilities for Huggingface Transformers."""
import contextlib import contextlib
import glob
import json import json
import os import os
import shutil
import time import time
from functools import reduce from functools import reduce
from pathlib import Path from pathlib import Path
@@ -43,8 +45,166 @@ from transformers.models.auto.modeling_auto import MODEL_FOR_CAUSAL_LM_MAPPING_N
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
from sglang.utils import is_in_ci
logger = init_logger(__name__) logger = init_logger(__name__)
def _check_index_files_for_missing_shards(
model_path: str,
) -> tuple[bool, list[str], list[str]]:
"""
Check all subdirectories for missing shards based on index files.
This catches cases where a model download was interrupted, leaving
some safetensors shards missing while the index file exists.
Args:
model_path: Path to the model directory
Returns:
Tuple of (all_valid, missing_files, checked_subdirs)
"""
missing_files = []
checked_subdirs = []
# Check the root directory and all subdirectories that might contain model weights
dirs_to_check = [model_path]
# Add common subdirectories for diffusers models
try:
subdirs = os.listdir(model_path)
except OSError as e:
logger.warning("Failed to list model directory %s: %s", model_path, e)
return True, [], [] # Assume valid if we can't check
for subdir in subdirs:
subdir_path = os.path.join(model_path, subdir)
if os.path.isdir(subdir_path):
dirs_to_check.append(subdir_path)
for dir_path in dirs_to_check:
# Find all safetensors index files
index_files = glob.glob(os.path.join(dir_path, "*.safetensors.index.json"))
for index_file in index_files:
checked_subdirs.append(os.path.basename(dir_path))
try:
with open(index_file) as f:
index_data = json.load(f)
weight_map = index_data.get("weight_map", {})
if not weight_map:
continue
# Get unique files referenced in weight_map
required_files = set(weight_map.values())
for file_name in required_files:
file_path = os.path.join(dir_path, file_name)
if not os.path.exists(file_path):
relative_path = os.path.relpath(file_path, model_path)
missing_files.append(relative_path)
except Exception as e:
logger.warning("Failed to read index file %s: %s", index_file, e)
continue
return len(missing_files) == 0, missing_files, checked_subdirs
def _cleanup_model_cache(model_path: str, reason: str) -> bool:
"""
Remove the model cache directory to force a clean re-download.
Args:
model_path: Path to the model directory (snapshot path)
reason: Reason for cleanup (for logging)
Returns:
True if cleanup was performed, False otherwise
"""
# Navigate up to the model root directory: snapshots/hash -> snapshots -> model_root
# HF cache structure: models--org--name/snapshots/hash/
try:
snapshot_dir = os.path.abspath(model_path)
snapshots_dir = os.path.dirname(snapshot_dir)
repo_folder = os.path.dirname(snapshots_dir)
# Verify this looks like an HF cache structure
if os.path.basename(snapshots_dir) != "snapshots":
logger.warning(
"Model path %s doesn't appear to be in HF cache structure, skipping cleanup",
model_path,
)
return False
logger.warning(
"Removing model cache at %s. Reason: %s",
repo_folder,
reason,
)
shutil.rmtree(repo_folder)
logger.info("Successfully removed corrupted cache directory")
return True
except Exception as e:
logger.error(
"Failed to remove corrupted cache directory %s: %s. "
"Manual cleanup may be required.",
model_path,
e,
)
return False
def _ci_validate_diffusers_model(model_path: str) -> tuple[bool, bool]:
"""
CI-specific validation for diffusers models.
Checks all subdirectories (transformer, transformer_2, vae, etc.) for
missing shards based on their index files. If issues are found in CI,
cleans up the cache to force re-download.
Args:
model_path: Path to the model directory
Returns:
Tuple of (is_valid, cleanup_performed)
- is_valid: True if the model is valid
- cleanup_performed: True if cleanup was performed (only relevant when is_valid=False)
"""
if not is_in_ci():
return True, False
is_valid, missing_files, checked_subdirs = _check_index_files_for_missing_shards(
model_path
)
if not is_valid:
logger.error(
"CI validation failed for %s. Missing %d file(s): %s. "
"Checked subdirectories: %s",
model_path,
len(missing_files),
missing_files[:5] if len(missing_files) > 5 else missing_files,
checked_subdirs,
)
cleanup_performed = _cleanup_model_cache(
model_path,
f"Missing {len(missing_files)} shard file(s): {missing_files[:3]}",
)
return False, cleanup_performed
if checked_subdirs:
logger.info(
"CI validation passed for %s. Checked subdirectories: %s",
model_path,
checked_subdirs,
)
return True, False
_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,
@@ -402,8 +562,33 @@ def maybe_download_model(
# 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):
if is_lora or _verify_model_complete(model_name_or_path): if is_lora or _verify_model_complete(model_name_or_path):
logger.info("Model already exists locally and is complete") # CI validation: check all subdirectories for missing shards
return model_name_or_path if not is_lora:
is_valid, cleanup_performed = _ci_validate_diffusers_model(
model_name_or_path
)
if not is_valid:
if cleanup_performed:
logger.warning(
"CI validation failed for local model at %s, "
"cache has been cleaned up, will re-download",
model_name_or_path,
)
# Fall through to download
else:
# Local path is not in HF cache structure, can't clean up
# Raise error since we can't fix this automatically
raise ValueError(
f"CI validation failed for local model at {model_name_or_path}. "
"Some safetensors shards are missing. "
"Please manually delete the model directory and retry."
)
else:
logger.info("Model already exists locally and is complete")
return model_name_or_path
else:
logger.info("Model already exists locally and is complete")
return model_name_or_path
else: else:
logger.warning( logger.warning(
"Local model at %s appears incomplete (missing transformer/ or vae/), " "Local model at %s appears incomplete (missing transformer/ or vae/), "
@@ -427,8 +612,30 @@ def maybe_download_model(
etag_timeout=60, etag_timeout=60,
) )
if is_lora or _verify_model_complete(local_path): if is_lora or _verify_model_complete(local_path):
logger.info("Found complete model in cache at %s", local_path) # CI validation: check all subdirectories for missing shards
return str(local_path) if not is_lora:
is_valid, cleanup_performed = _ci_validate_diffusers_model(local_path)
if not is_valid:
if cleanup_performed:
logger.warning(
"CI validation failed for cached model at %s, "
"cache has been cleaned up, will re-download",
local_path,
)
# Fall through to download
else:
# This shouldn't happen for HF cache paths, but handle it
logger.warning(
"CI validation failed for cached model at %s, "
"but cleanup was not performed, will attempt re-download",
local_path,
)
else:
logger.info("Found complete model in cache at %s", local_path)
return str(local_path)
else:
logger.info("Found complete model in cache at %s", local_path)
return str(local_path)
else: else:
# Model found in cache but incomplete # Model found in cache but incomplete
if not download: if not download:
@@ -498,6 +705,17 @@ def maybe_download_model(
"The model repository may be missing required components (model_index.json, transformer/, or vae/)." "The model repository may be missing required components (model_index.json, transformer/, or vae/)."
) )
# CI validation: check all subdirectories for missing shards after download
if not is_lora:
is_valid, cleanup_performed = _ci_validate_diffusers_model(local_path)
if not is_valid:
# In CI, if validation fails after download, we have a serious issue
# If cleanup was performed, the next retry should get a fresh download
raise ValueError(
f"CI validation failed for downloaded model at {local_path}. "
f"Some safetensors shards are missing. Cleanup performed: {cleanup_performed}."
)
logger.info("Downloaded model to %s", local_path) logger.info("Downloaded model to %s", local_path)
return str(local_path) return str(local_path)