From 9121f2265644af2fd9803fc491206684956c8168 Mon Sep 17 00:00:00 2001 From: Alison Shao <54658187+alisonshao@users.noreply.github.com> Date: Sat, 24 Jan 2026 19:18:15 -0800 Subject: [PATCH] Add PyTorch .bin file validation to CI weight validation (#17533) --- .../srt/model_loader/ci_weight_validation.py | 73 ++++++++++++++++++- 1 file changed, 69 insertions(+), 4 deletions(-) diff --git a/python/sglang/srt/model_loader/ci_weight_validation.py b/python/sglang/srt/model_loader/ci_weight_validation.py index 96054d5d3..55978e939 100644 --- a/python/sglang/srt/model_loader/ci_weight_validation.py +++ b/python/sglang/srt/model_loader/ci_weight_validation.py @@ -1219,6 +1219,38 @@ def _validate_safetensors_file(file_path: str) -> bool: return False +def _validate_pytorch_bin_file(file_path: str) -> bool: + """ + Validate that a PyTorch .bin file is readable and not corrupted. + + This catches corruption issues like truncated downloads or invalid archives + that would cause errors like: + "RuntimeError: PytorchStreamReader failed reading file data/X: invalid header + or archive is corrupted" + + Args: + file_path: Path to the .bin file + + Returns: + True if the file is valid, False if corrupted + """ + try: + import torch + + # Use weights_only=True for security and to avoid executing arbitrary code + # mmap=False to fully read the file and catch all corruption + torch.load(file_path, map_location="cpu", weights_only=True, mmap=False) + return True + except Exception as e: + logger.warning( + "Corrupted PyTorch bin file detected: %s - %s: %s", + file_path, + type(e).__name__, + str(e), + ) + return False + + def _check_index_files_exist(snapshot_dir: str) -> Tuple[bool, Optional[str]]: """ Check if all files listed in safetensors index files actually exist on disk. @@ -1366,11 +1398,15 @@ def _validate_sharded_model( [], ) - # Validate safetensors files for corruption + # Validate weight files for corruption if group_info["suffix"] == "safetensors": for f in group_info["files"]: if not _validate_safetensors_file(f): corrupted_files.append(f) + elif group_info["suffix"] == "bin": + for f in group_info["files"]: + if not _validate_pytorch_bin_file(f): + corrupted_files.append(f) # Check for required index file for safetensors shards if group_info["suffix"] == "safetensors": @@ -1561,7 +1597,7 @@ def ci_validate_and_cleanup_local_snapshot( ) return False - # Also validate single (non-sharded) safetensors files + # Also validate single (non-sharded) weight files for f in local_weight_files: base_name = os.path.basename(f) # Check if this is a single model file (not sharded) @@ -1580,6 +1616,21 @@ def ci_validate_and_cleanup_local_snapshot( # Selective cleanup for single file _cleanup_corrupted_files_selective(model_name_or_path, [f]) return False + # Also validate single PyTorch .bin files + elif base_name in [ + "pytorch_model.bin", + "model.bin", + "adapter_model.bin", + ]: + if not _validate_pytorch_bin_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 False return True @@ -1613,12 +1664,15 @@ def _validate_weights_after_download( if not weight_files: return True # No weight files to validate - # Validate safetensors files + # Validate weight files (safetensors and .bin) corrupted_files = [] for f in weight_files: if f.endswith(".safetensors") and os.path.exists(f): if not _validate_safetensors_file(f): corrupted_files.append(os.path.basename(f)) + elif f.endswith(".bin") and os.path.exists(f): + if not _validate_pytorch_bin_file(f): + corrupted_files.append(os.path.basename(f)) if corrupted_files: # Clean up corrupted files so next attempt re-downloads them @@ -1777,9 +1831,20 @@ def ci_validate_and_clean_hf_cache(model_path: str) -> None: if not _validate_safetensors_file(sf_file): corrupted_files.append(sf_file) + # Also find and validate PyTorch .bin files + bin_files = glob_module.glob(os.path.join(snapshot_dir, "*.bin")) + + for bin_file in bin_files: + # Skip broken symlinks (os.path.exists returns False for them) + if not os.path.exists(bin_file): + continue + + if not _validate_pytorch_bin_file(bin_file): + corrupted_files.append(bin_file) + if corrupted_files: logger.warning( - "HFRunner: Found %d corrupted safetensors file(s) for %s. " + "HFRunner: Found %d corrupted weight file(s) for %s. " "Removing to force re-download.", len(corrupted_files), model_path,