From b6523a4f729393369c70217fa8687f9daf83c62b Mon Sep 17 00:00:00 2001 From: Alison Shao <54658187+alisonshao@users.noreply.github.com> Date: Wed, 10 Dec 2025 16:03:53 -0800 Subject: [PATCH] fix: restrict cache validation behaviors to CI only (#14849) --- .../sglang/srt/model_loader/weight_utils.py | 206 ++++++++++-------- 1 file changed, 115 insertions(+), 91 deletions(-) diff --git a/python/sglang/srt/model_loader/weight_utils.py b/python/sglang/srt/model_loader/weight_utils.py index e4c1255ac..82e69d88c 100644 --- a/python/sglang/srt/model_loader/weight_utils.py +++ b/python/sglang/srt/model_loader/weight_utils.py @@ -47,6 +47,7 @@ from sglang.srt.model_loader.weight_validation import ( _validate_sharded_model, ) from sglang.srt.utils import find_local_repo_dir, log_info_on_rank0, print_warning_once +from sglang.utils import is_in_ci logger = logging.getLogger(__name__) @@ -326,27 +327,32 @@ def _find_local_hf_snapshot_dir_unlocked( if not os.path.isdir(found_local_snapshot_dir): return None - # Check for incomplete files and clean up if found - repo_folder = os.path.abspath(os.path.join(found_local_snapshot_dir, "..", "..")) - blobs_dir = os.path.join(repo_folder, "blobs") - - # Check for incomplete download markers - incomplete_files = [] - if os.path.isdir(blobs_dir): - incomplete_files = glob.glob(os.path.join(blobs_dir, "*.incomplete")) - - if incomplete_files: - log_info_on_rank0( - logger, - f"Found {len(incomplete_files)} .incomplete files in {blobs_dir} for " - f"{model_name_or_path}. Will clean up and re-download.", + # Only perform cache validation and cleanup in CI to avoid + # unnecessary overhead for regular users + if is_in_ci(): + # Check for incomplete files and clean up if found + repo_folder = os.path.abspath( + os.path.join(found_local_snapshot_dir, "..", "..") ) - _cleanup_corrupted_model_cache( - model_name_or_path, - found_local_snapshot_dir, - f"Incomplete download detected ({len(incomplete_files)} incomplete files)", - ) - return None + blobs_dir = os.path.join(repo_folder, "blobs") + + # Check for incomplete download markers + incomplete_files = [] + if os.path.isdir(blobs_dir): + incomplete_files = glob.glob(os.path.join(blobs_dir, "*.incomplete")) + + if incomplete_files: + log_info_on_rank0( + logger, + f"Found {len(incomplete_files)} .incomplete files in {blobs_dir} for " + f"{model_name_or_path}. Will clean up and re-download.", + ) + _cleanup_corrupted_model_cache( + model_name_or_path, + found_local_snapshot_dir, + f"Incomplete download detected ({len(incomplete_files)} incomplete files)", + ) + return None local_weight_files: List[str] = [] try: @@ -366,53 +372,57 @@ def _find_local_hf_snapshot_dir_unlocked( ) local_weight_files = [] - # Validate sharded models and check for corruption - if local_weight_files: - is_valid, error_msg, corrupted_files = _validate_sharded_model( - found_local_snapshot_dir, local_weight_files - ) - if not is_valid: - if corrupted_files: - # Selective cleanup: only remove corrupted files - log_info_on_rank0( - logger, - f"Found {len(corrupted_files)} corrupted file(s) for " - f"{model_name_or_path}: {error_msg}. " - "Will selectively clean and re-download only these files.", - ) - _cleanup_corrupted_files_selective(model_name_or_path, corrupted_files) - return None - else: - # Cannot selectively clean (e.g., missing shards) - remove entire cache - log_info_on_rank0( - logger, - f"Validation failed for {model_name_or_path}: {error_msg}. " - "Will remove entire cache and re-download.", - ) - _cleanup_corrupted_model_cache( - model_name_or_path, found_local_snapshot_dir, error_msg - ) - return None - - # Also validate single (non-sharded) safetensors files - for f in local_weight_files: - base_name = os.path.basename(f) - # Check if this is a single model file (not sharded) - # Include adapter_model.safetensors for LoRA adapters - if base_name in [ - "model.safetensors", - "pytorch_model.safetensors", - "adapter_model.safetensors", - ]: - if not _validate_safetensors_file(f): + # Only perform cache validation and cleanup in CI + if is_in_ci(): + # Validate sharded models and check for corruption + if local_weight_files: + is_valid, error_msg, corrupted_files = _validate_sharded_model( + found_local_snapshot_dir, local_weight_files + ) + if not is_valid: + if corrupted_files: + # Selective cleanup: only remove corrupted files log_info_on_rank0( logger, - f"Corrupted model file {base_name} for {model_name_or_path}. " - "Will selectively clean and re-download this file.", + f"Found {len(corrupted_files)} corrupted file(s) for " + f"{model_name_or_path}: {error_msg}. " + "Will selectively clean and re-download only these files.", + ) + _cleanup_corrupted_files_selective( + model_name_or_path, corrupted_files ) - # Selective cleanup for single file - _cleanup_corrupted_files_selective(model_name_or_path, [f]) return None + else: + # Cannot selectively clean (e.g., missing shards) - remove entire cache + log_info_on_rank0( + logger, + f"Validation failed for {model_name_or_path}: {error_msg}. " + "Will remove entire cache and re-download.", + ) + _cleanup_corrupted_model_cache( + model_name_or_path, found_local_snapshot_dir, error_msg + ) + return None + + # Also validate single (non-sharded) safetensors files + for f in local_weight_files: + base_name = os.path.basename(f) + # Check if this is a single model file (not sharded) + # Include adapter_model.safetensors for LoRA adapters + if base_name in [ + "model.safetensors", + "pytorch_model.safetensors", + "adapter_model.safetensors", + ]: + if not _validate_safetensors_file(f): + log_info_on_rank0( + logger, + f"Corrupted model file {base_name} for {model_name_or_path}. " + "Will selectively clean and re-download this file.", + ) + # Selective cleanup for single file + _cleanup_corrupted_files_selective(model_name_or_path, [f]) + return None if len(local_weight_files) > 0: log_info_on_rank0( @@ -569,8 +579,47 @@ def download_weights_from_hf( log_info_on_rank0(logger, f"Using model weights format {allow_patterns}") - # Retry loop for handling corrupted downloads - for attempt in range(max_retries): + # Only perform validation and retry in CI to avoid overhead for regular users + if is_in_ci(): + # Retry loop for handling corrupted downloads + for attempt in range(max_retries): + hf_folder = snapshot_download( + model_name_or_path, + allow_patterns=allow_patterns, + ignore_patterns=ignore_patterns, + cache_dir=cache_dir, + tqdm_class=DisabledTqdm, + revision=revision, + local_files_only=huggingface_hub.constants.HF_HUB_OFFLINE, + ) + + # Validate downloaded files to catch corruption early + is_valid = _validate_weights_after_download( + hf_folder, allow_patterns, model_name_or_path + ) + + if is_valid: + return hf_folder + + # Validation failed, corrupted files were cleaned up + if attempt < max_retries - 1: + log_info_on_rank0( + logger, + f"Retrying download for {model_name_or_path} " + f"(attempt {attempt + 2}/{max_retries})...", + ) + else: + raise RuntimeError( + f"Downloaded model files are still corrupted for " + f"{model_name_or_path} after {max_retries} attempts. " + "This may indicate a persistent issue with the model files " + "on Hugging Face Hub or network problems." + ) + + # This should never be reached, but just in case + return hf_folder + else: + # Simple download without validation for non-CI environments hf_folder = snapshot_download( model_name_or_path, allow_patterns=allow_patterns, @@ -580,32 +629,7 @@ def download_weights_from_hf( revision=revision, local_files_only=huggingface_hub.constants.HF_HUB_OFFLINE, ) - - # Validate downloaded files to catch corruption early - is_valid = _validate_weights_after_download( - hf_folder, allow_patterns, model_name_or_path - ) - - if is_valid: - return hf_folder - - # Validation failed, corrupted files were cleaned up - if attempt < max_retries - 1: - log_info_on_rank0( - logger, - f"Retrying download for {model_name_or_path} " - f"(attempt {attempt + 2}/{max_retries})...", - ) - else: - raise RuntimeError( - f"Downloaded model files are still corrupted for " - f"{model_name_or_path} after {max_retries} attempts. " - "This may indicate a persistent issue with the model files " - "on Hugging Face Hub or network problems." - ) - - # This should never be reached, but just in case - return hf_folder + return hf_folder def download_safetensors_index_file_from_hf(