From 01835998e1c33656f1567f38012bf0f97de95631 Mon Sep 17 00:00:00 2001 From: Alison Shao <54658187+alisonshao@users.noreply.github.com> Date: Tue, 9 Dec 2025 20:36:54 -0800 Subject: [PATCH] fix: race condition between validation and download locks (#14761) --- .../sglang/srt/model_loader/weight_utils.py | 376 ++++++++++-------- 1 file changed, 216 insertions(+), 160 deletions(-) diff --git a/python/sglang/srt/model_loader/weight_utils.py b/python/sglang/srt/model_loader/weight_utils.py index a1359d79c..e4c1255ac 100644 --- a/python/sglang/srt/model_loader/weight_utils.py +++ b/python/sglang/srt/model_loader/weight_utils.py @@ -263,13 +263,19 @@ def get_quant_config( return quant_cls.from_config(config) -def find_local_hf_snapshot_dir( +def _find_local_hf_snapshot_dir_unlocked( model_name_or_path: str, cache_dir: Optional[str], allow_patterns: List[str], revision: Optional[str] = None, ) -> Optional[str]: - """If the weights are already local, skip downloading and returns the path.""" + """Find local HF snapshot directory without locking. + + IMPORTANT: Caller MUST hold the model lock before calling this function + to prevent race conditions during validation and cleanup. + + If the weights are already local, skip downloading and returns the path. + """ if os.path.isdir(model_name_or_path): return None @@ -315,131 +321,144 @@ def find_local_hf_snapshot_dir( if found_local_snapshot_dir is None: return None + # Check if snapshot dir exists (might have been cleaned by another process + # before we acquired the lock) + 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.", + ) + _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: + for pattern in allow_patterns: + matched_files = glob.glob(os.path.join(found_local_snapshot_dir, pattern)) + for f in matched_files: + # os.path.exists returns False for broken symlinks. + if not os.path.exists(f): + continue + local_weight_files.append(f) + except Exception as e: + logger.warning( + "Failed to scan local snapshot %s with patterns %s: %s", + found_local_snapshot_dir, + allow_patterns, + e, + ) + 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): + 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( + logger, + f"Found local HF snapshot for {model_name_or_path} at " + f"{found_local_snapshot_dir}; skipping download.", + ) + return found_local_snapshot_dir + else: + log_info_on_rank0( + logger, + f"Local HF snapshot at {found_local_snapshot_dir} has no files matching " + f"{allow_patterns}; will attempt download.", + ) + return None + + +def find_local_hf_snapshot_dir( + model_name_or_path: str, + cache_dir: Optional[str], + allow_patterns: List[str], + revision: Optional[str] = None, +) -> Optional[str]: + """If the weights are already local, skip downloading and returns the path. + + This function acquires a lock to prevent race conditions during validation + and cleanup. For use within download_weights_from_hf, use + _find_local_hf_snapshot_dir_unlocked instead with an external lock. + """ + # For local paths, no locking needed + if os.path.isdir(model_name_or_path): + return None + # Use file lock to prevent multiple processes (TP ranks) from # validating and cleaning up the same model cache simultaneously. - # This prevents race conditions where multiple ranks detect corruption - # and try to delete the same files at the same time. - with get_lock(model_name_or_path, cache_dir, suffix="-validation"): - # Re-check if snapshot dir still exists after acquiring lock - # (another process may have already cleaned it up) - 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, "..", "..") + with get_lock(model_name_or_path, cache_dir): + return _find_local_hf_snapshot_dir_unlocked( + model_name_or_path, cache_dir, allow_patterns, revision ) - 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: - for pattern in allow_patterns: - matched_files = glob.glob( - os.path.join(found_local_snapshot_dir, pattern) - ) - for f in matched_files: - # os.path.exists returns False for broken symlinks. - if not os.path.exists(f): - continue - local_weight_files.append(f) - except Exception as e: - logger.warning( - "Failed to scan local snapshot %s with patterns %s: %s", - found_local_snapshot_dir, - allow_patterns, - e, - ) - 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): - 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( - logger, - f"Found local HF snapshot for {model_name_or_path} at " - f"{found_local_snapshot_dir}; skipping download.", - ) - return found_local_snapshot_dir - else: - log_info_on_rank0( - logger, - f"Local HF snapshot at {found_local_snapshot_dir} has no files matching " - f"{allow_patterns}; will attempt download.", - ) - return None def _validate_weights_after_download( hf_folder: str, allow_patterns: List[str], model_name_or_path: str, -) -> None: +) -> bool: """Validate downloaded weight files to catch corruption early. This function validates safetensors files after download to catch @@ -451,8 +470,8 @@ def _validate_weights_after_download( allow_patterns: Patterns used to match weight files model_name_or_path: Model identifier for error messages - Raises: - RuntimeError: If any weight files are corrupted + Returns: + True if all files are valid, False if corrupted files were found and cleaned up """ import glob as glob_module @@ -462,7 +481,7 @@ def _validate_weights_after_download( weight_files.extend(glob_module.glob(os.path.join(hf_folder, pattern))) if not weight_files: - return # No weight files to validate + return True # No weight files to validate # Validate safetensors files corrupted_files = [] @@ -477,11 +496,15 @@ def _validate_weights_after_download( model_name_or_path, [os.path.join(hf_folder, f) for f in corrupted_files], ) - raise RuntimeError( + log_info_on_rank0( + logger, f"Downloaded model files are corrupted for {model_name_or_path}: " f"{corrupted_files}. The corrupted files have been removed. " - "Please retry to re-download the model." + "Will retry download.", ) + return False + + return True def download_weights_from_hf( @@ -490,6 +513,7 @@ def download_weights_from_hf( allow_patterns: List[str], revision: Optional[str] = None, ignore_patterns: Optional[Union[str, List[str]]] = None, + max_retries: int = 3, ) -> str: """Download model weights from Hugging Face Hub. @@ -504,52 +528,84 @@ def download_weights_from_hf( ignore_patterns (Optional[Union[str, List[str]]]): The patterns to filter out the weight files. Files matched by any of the patterns will be ignored. + max_retries (int): Maximum number of download retries if corruption + is detected. Defaults to 3. Returns: str: The path to the downloaded model weights. """ + # For local paths, no HF operations needed + if os.path.isdir(model_name_or_path): + return model_name_or_path - # Always check for valid local cache first. - # This validates cached files and cleans up corrupted ones. - path = find_local_hf_snapshot_dir( - model_name_or_path, cache_dir, allow_patterns, revision - ) - if path is not None: - # Valid local cache found, skip download - return path - - # In CI, skip HF API calls if we're in offline mode or want to avoid rate limits - # But we already checked for local cache above, so if we're here we need to download - if not huggingface_hub.constants.HF_HUB_OFFLINE: - # Before we download we look at what is available: - fs = HfFileSystem() - file_list = fs.ls(model_name_or_path, detail=False, revision=revision) - - # depending on what is available we download different things - for pattern in allow_patterns: - matching = fnmatch.filter(file_list, pattern) - if len(matching) > 0: - allow_patterns = [pattern] - break - - log_info_on_rank0(logger, f"Using model weights format {allow_patterns}") - # Use file lock to prevent multiple processes from - # downloading the same model weights at the same time. + # Use a SINGLE lock for the entire operation (validation + cleanup + download) + # to prevent race conditions where: + # 1. Process A validates, finds corruption, deletes corrupted file + # 2. Process B validates, sees missing file, deletes ENTIRE cache + # 3. Process A tries to download but cache is gone + # By using one lock, validation/cleanup and download are atomic. with get_lock(model_name_or_path, cache_dir): - 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, + # Check for valid local cache first (validates and cleans up if needed) + path = _find_local_hf_snapshot_dir_unlocked( + model_name_or_path, cache_dir, allow_patterns, revision ) + if path is not None: + # Valid local cache found, skip download + return path - # Validate downloaded files to catch corruption early - _validate_weights_after_download(hf_folder, allow_patterns, model_name_or_path) + # In CI, skip HF API calls if we're in offline mode or want to avoid rate limits + # But we already checked for local cache above, so if we're here we need to download + if not huggingface_hub.constants.HF_HUB_OFFLINE: + # Before we download we look at what is available: + fs = HfFileSystem() + file_list = fs.ls(model_name_or_path, detail=False, revision=revision) - return hf_folder + # depending on what is available we download different things + for pattern in allow_patterns: + matching = fnmatch.filter(file_list, pattern) + if len(matching) > 0: + allow_patterns = [pattern] + break + + log_info_on_rank0(logger, f"Using model weights format {allow_patterns}") + + # 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 def download_safetensors_index_file_from_hf(