[misc] Add a comment style rule to .claude/rules (#35597)
This commit is contained in:
@@ -1,17 +1,8 @@
|
||||
"""
|
||||
CI-specific weight validation and cache cleanup utilities.
|
||||
"""CI-only weight validation and cache cleanup.
|
||||
|
||||
This module contains validation and cleanup logic that is ONLY used in CI environments.
|
||||
These functions handle:
|
||||
- Validating safetensors files for corruption
|
||||
- Checking for missing shards in sharded models
|
||||
- Cleaning up corrupted files (selective or full cache deletion)
|
||||
- Automatic retry logic for corrupted downloads
|
||||
- Validating config/tokenizer files completeness to enable offline mode
|
||||
|
||||
For regular users, weight_utils.py provides simple download functionality without
|
||||
the overhead of validation and automatic cleanup. The CI-specific behavior is
|
||||
gated by is_in_ci() checks in weight_utils.py.
|
||||
Validates safetensors/bin files and shard completeness, and repairs the HF cache
|
||||
by deleting what has to be re-downloaded. `weight_utils.py` gates every entry
|
||||
point here behind `is_in_ci()`; regular users take the plain download path.
|
||||
"""
|
||||
|
||||
import glob as glob_module
|
||||
@@ -33,17 +24,7 @@ logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def _get_per_run_marker_dir() -> str:
|
||||
"""
|
||||
Get the directory for per-run validation markers.
|
||||
|
||||
These markers are specific to the current CI run and are not shared across
|
||||
runners. They are stored in a temporary directory that is cleaned up after
|
||||
the run completes.
|
||||
|
||||
Returns:
|
||||
Path to per-run marker directory
|
||||
"""
|
||||
# Prefer RUNNER_TEMP (GitHub Actions) or TMPDIR, fallback to /tmp
|
||||
# Markers are per CI run; sharing them across runners leaks cache state.
|
||||
base_dir = os.environ.get("RUNNER_TEMP", os.environ.get("TMPDIR", "/tmp"))
|
||||
marker_dir = os.path.join(base_dir, "sglang_ci_offline_markers")
|
||||
os.makedirs(marker_dir, exist_ok=True)
|
||||
@@ -51,18 +32,6 @@ def _get_per_run_marker_dir() -> str:
|
||||
|
||||
|
||||
def _get_per_run_marker_path(snapshot_dir: str) -> Optional[str]:
|
||||
"""
|
||||
Get the path to per-run validation marker file for a snapshot.
|
||||
|
||||
Per-run markers are specific to the current CI run and are not shared
|
||||
across runners. This prevents cross-runner cache state pollution.
|
||||
|
||||
Args:
|
||||
snapshot_dir: Path to snapshot directory
|
||||
|
||||
Returns:
|
||||
Path to per-run marker file or None if snapshot_dir is invalid
|
||||
"""
|
||||
if not snapshot_dir or not os.path.isdir(snapshot_dir):
|
||||
return None
|
||||
|
||||
@@ -74,15 +43,6 @@ def _get_per_run_marker_path(snapshot_dir: str) -> Optional[str]:
|
||||
|
||||
|
||||
def _read_per_run_marker(snapshot_dir: str) -> Optional[dict]:
|
||||
"""
|
||||
Read per-run validation marker for a snapshot.
|
||||
|
||||
Args:
|
||||
snapshot_dir: Path to snapshot directory
|
||||
|
||||
Returns:
|
||||
Marker dict if exists and valid, None otherwise
|
||||
"""
|
||||
marker_path = _get_per_run_marker_path(snapshot_dir)
|
||||
if not marker_path or not os.path.exists(marker_path):
|
||||
return None
|
||||
@@ -91,7 +51,6 @@ def _read_per_run_marker(snapshot_dir: str) -> Optional[dict]:
|
||||
with open(marker_path, "r", encoding="utf-8") as f:
|
||||
marker = json.load(f)
|
||||
|
||||
# Validate marker structure
|
||||
if not isinstance(marker, dict):
|
||||
return None
|
||||
|
||||
@@ -112,14 +71,6 @@ def _read_per_run_marker(snapshot_dir: str) -> Optional[dict]:
|
||||
def _write_per_run_marker(
|
||||
snapshot_dir: str, model_id: str, required_files: Optional[list] = None
|
||||
) -> None:
|
||||
"""
|
||||
Write per-run validation marker for a snapshot.
|
||||
|
||||
Args:
|
||||
snapshot_dir: Path to snapshot directory
|
||||
model_id: Model identifier
|
||||
required_files: List of required files that were validated
|
||||
"""
|
||||
marker_path = _get_per_run_marker_path(snapshot_dir)
|
||||
if not marker_path:
|
||||
logger.debug("Cannot write per-run marker: invalid snapshot_dir")
|
||||
@@ -165,22 +116,7 @@ def _write_per_run_marker(
|
||||
def validate_cache_lightweight(
|
||||
snapshot_dir: str, requires_hf_quant_config: bool = False
|
||||
) -> bool:
|
||||
"""
|
||||
Lightweight runtime validation for cache completeness.
|
||||
|
||||
This is used during test runs to ensure the current runner's cache
|
||||
is complete before enabling offline mode. Much faster than full validation
|
||||
as it only checks file existence, not corruption.
|
||||
|
||||
Args:
|
||||
snapshot_dir: Path to the model snapshot directory
|
||||
requires_hf_quant_config: If True, hf_quant_config.json must exist
|
||||
(required for modelopt quantization)
|
||||
|
||||
Returns:
|
||||
True if cache is complete, False otherwise
|
||||
"""
|
||||
# Check required config files
|
||||
"""Existence-only cache check: no corruption reads, cheap enough to run per test."""
|
||||
required_files = [
|
||||
"config.json",
|
||||
"tokenizer_config.json",
|
||||
@@ -190,7 +126,6 @@ def validate_cache_lightweight(
|
||||
if not os.path.exists(os.path.join(snapshot_dir, fname)):
|
||||
return False
|
||||
|
||||
# Check tokenizer files (at least one must exist)
|
||||
tokenizer_files = [
|
||||
"tokenizer.json",
|
||||
"tokenizer.model",
|
||||
@@ -203,7 +138,6 @@ def validate_cache_lightweight(
|
||||
if not has_tokenizer:
|
||||
return False
|
||||
|
||||
# Check for trust_remote_code dynamic module files if needed
|
||||
# When auto_map exists in config.json, the model requires custom Python files
|
||||
# These files must be present for offline mode to work
|
||||
config_path = os.path.join(snapshot_dir, "config.json")
|
||||
@@ -214,17 +148,13 @@ def validate_cache_lightweight(
|
||||
|
||||
auto_map = config.get("auto_map", {})
|
||||
if auto_map and isinstance(auto_map, dict):
|
||||
# Extract Python module files from auto_map
|
||||
# auto_map format: {"AutoConfig": "configuration_xxx.ConfigClass", ...}
|
||||
# We need to check if the .py files exist
|
||||
custom_files = set()
|
||||
for key, value in auto_map.items():
|
||||
if isinstance(value, str) and "." in value:
|
||||
# Extract module name (e.g., "configuration_xxx" from "configuration_xxx.ConfigClass")
|
||||
module_name = value.split(".")[0]
|
||||
custom_files.add(f"{module_name}.py")
|
||||
|
||||
# Check if all custom files exist in snapshot directory
|
||||
for custom_file in custom_files:
|
||||
custom_file_path = os.path.join(snapshot_dir, custom_file)
|
||||
if not os.path.exists(custom_file_path):
|
||||
@@ -249,13 +179,11 @@ def validate_cache_lightweight(
|
||||
has_index = os.path.exists(index_path)
|
||||
|
||||
if has_index:
|
||||
# If index exists, validate that all shards listed in it exist
|
||||
try:
|
||||
with open(index_path, "r", encoding="utf-8") as f:
|
||||
index_data = json.load(f)
|
||||
weight_map = index_data.get("weight_map", {})
|
||||
if weight_map:
|
||||
# Check that all shard files referenced in index exist
|
||||
required_shards = set(weight_map.values())
|
||||
for shard_name in required_shards:
|
||||
shard_path = os.path.join(snapshot_dir, shard_name)
|
||||
@@ -270,7 +198,6 @@ def validate_cache_lightweight(
|
||||
logger.debug("Failed to validate index file %s: %s", index_path, e)
|
||||
return False
|
||||
else:
|
||||
# No index file - check for weight files and validate shard completeness
|
||||
safetensors_files = glob_module.glob(
|
||||
os.path.join(snapshot_dir, "*.safetensors")
|
||||
)
|
||||
@@ -278,7 +205,6 @@ def validate_cache_lightweight(
|
||||
return False
|
||||
|
||||
# Check shard completeness for sharded models (e.g., model-00001-of-00047.safetensors)
|
||||
# Pattern: prefix-NNNNN-of-NNNNN.safetensors
|
||||
shard_pattern = re.compile(r"(.*?)-(\d+)-of-(\d+)\.safetensors$")
|
||||
shard_groups = {}
|
||||
|
||||
@@ -298,7 +224,6 @@ def validate_cache_lightweight(
|
||||
}
|
||||
shard_groups[group_key]["found_shards"].add(shard_id)
|
||||
|
||||
# Validate each shard group has all expected shards
|
||||
for group_key, group_info in shard_groups.items():
|
||||
total_shards = group_info["total"]
|
||||
found_shards = group_info["found_shards"]
|
||||
@@ -324,20 +249,9 @@ def validate_cache_lightweight(
|
||||
|
||||
|
||||
def _validate_safetensors_file(file_path: str) -> bool:
|
||||
"""
|
||||
Validate that a safetensors file is readable and not corrupted.
|
||||
|
||||
Args:
|
||||
file_path: Path to the safetensors file
|
||||
|
||||
Returns:
|
||||
True if the file is valid, False if corrupted
|
||||
"""
|
||||
try:
|
||||
# Attempt to open and read the header
|
||||
# This will fail if the file is corrupted or incomplete
|
||||
with safetensors.safe_open(file_path, framework="pt", device="cpu") as f:
|
||||
# Just accessing the keys validates the header is readable
|
||||
# Listing keys forces the header parse; open() alone does not.
|
||||
_ = list(f.keys())
|
||||
return True
|
||||
except Exception as e:
|
||||
@@ -351,20 +265,7 @@ def _validate_safetensors_file(file_path: str) -> bool:
|
||||
|
||||
|
||||
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
|
||||
"""
|
||||
# Truncated archives surface as a PytorchStreamReader "invalid header" error.
|
||||
try:
|
||||
import torch
|
||||
|
||||
@@ -383,19 +284,6 @@ def _validate_pytorch_bin_file(file_path: str) -> bool:
|
||||
|
||||
|
||||
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.
|
||||
|
||||
This catches cases where the snapshot directory exists but files are missing
|
||||
(e.g., due to incomplete downloads or corrupted cache).
|
||||
|
||||
Args:
|
||||
snapshot_dir: Path to the model snapshot directory
|
||||
|
||||
Returns:
|
||||
Tuple of (all_exist, error_message)
|
||||
"""
|
||||
# Find all safetensors index files
|
||||
index_files = [
|
||||
f for f in os.listdir(snapshot_dir) if f.endswith(".safetensors.index.json")
|
||||
]
|
||||
@@ -416,7 +304,6 @@ def _check_index_files_exist(snapshot_dir: str) -> Tuple[bool, Optional[str]]:
|
||||
logger.warning(
|
||||
"Removed broken index symlink: %s (blob missing)", index_file
|
||||
)
|
||||
# Also try to remove dangling blob reference if it somehow exists
|
||||
if os.path.exists(blob_path):
|
||||
os.remove(blob_path)
|
||||
except Exception as e:
|
||||
@@ -434,7 +321,6 @@ def _check_index_files_exist(snapshot_dir: str) -> Tuple[bool, Optional[str]]:
|
||||
if not weight_map:
|
||||
continue
|
||||
|
||||
# Check that all files in weight_map exist
|
||||
required_files = set(weight_map.values())
|
||||
missing_files = []
|
||||
|
||||
@@ -467,18 +353,6 @@ def _check_index_files_exist(snapshot_dir: str) -> Tuple[bool, Optional[str]]:
|
||||
def _validate_sharded_model(
|
||||
snapshot_dir: str, weight_files: List[str]
|
||||
) -> Tuple[bool, Optional[str], List[str]]:
|
||||
"""
|
||||
Validate that all model shards are present and not corrupted.
|
||||
|
||||
Args:
|
||||
snapshot_dir: Path to the model snapshot directory
|
||||
weight_files: List of weight file paths
|
||||
|
||||
Returns:
|
||||
Tuple of (is_valid, error_message, corrupted_files)
|
||||
- corrupted_files: List of file paths that are corrupted (for selective cleanup)
|
||||
"""
|
||||
# First, check if all files from the index actually exist
|
||||
# This catches missing files that wouldn't be found by glob
|
||||
index_check_valid, index_error = _check_index_files_exist(snapshot_dir)
|
||||
if not index_check_valid:
|
||||
@@ -487,7 +361,6 @@ def _validate_sharded_model(
|
||||
# Pattern for sharded files: model-00001-of-00009.safetensors
|
||||
shard_pattern = re.compile(r"(.*?)-(\d+)-of-(\d+)\.(safetensors|bin)")
|
||||
|
||||
# Group files by shard pattern (prefix-*-of-N)
|
||||
shard_groups = {}
|
||||
for f in weight_files:
|
||||
base_name = os.path.basename(f)
|
||||
@@ -511,10 +384,8 @@ def _validate_sharded_model(
|
||||
shard_groups[group_key]["found_shards"].append(shard_id)
|
||||
shard_groups[group_key]["files"].append(f)
|
||||
|
||||
# Track corrupted files for selective cleanup
|
||||
corrupted_files = []
|
||||
|
||||
# Validate each shard group
|
||||
for group_key, group_info in shard_groups.items():
|
||||
total_shards = group_info["total"]
|
||||
found_shards = set(group_info["found_shards"])
|
||||
@@ -523,7 +394,6 @@ def _validate_sharded_model(
|
||||
min_idx = min(found_shards) if found_shards else 1
|
||||
expected_shards = set(range(min_idx, min_idx + total_shards))
|
||||
|
||||
# Check for missing shards
|
||||
missing_shards = expected_shards - found_shards
|
||||
if missing_shards:
|
||||
return (
|
||||
@@ -532,7 +402,6 @@ def _validate_sharded_model(
|
||||
[],
|
||||
)
|
||||
|
||||
# Validate weight files for corruption
|
||||
if group_info["suffix"] == "safetensors":
|
||||
for f in group_info["files"]:
|
||||
if not _validate_safetensors_file(f):
|
||||
@@ -542,7 +411,6 @@ def _validate_sharded_model(
|
||||
if not _validate_pytorch_bin_file(f):
|
||||
corrupted_files.append(f)
|
||||
|
||||
# Check for required index file for safetensors shards
|
||||
if group_info["suffix"] == "safetensors":
|
||||
index_file = os.path.join(
|
||||
snapshot_dir, f"{group_info['prefix']}.safetensors.index.json"
|
||||
@@ -567,19 +435,6 @@ def _validate_sharded_model(
|
||||
def _cleanup_corrupted_files_selective(
|
||||
model_name_or_path: str, corrupted_files: List[str]
|
||||
) -> int:
|
||||
"""
|
||||
Selectively remove corrupted files and their blobs to force re-download.
|
||||
|
||||
This is more efficient than removing the entire model cache as it only
|
||||
re-downloads corrupted files rather than the entire model.
|
||||
|
||||
Args:
|
||||
model_name_or_path: Model identifier
|
||||
corrupted_files: List of corrupted file paths (symlinks in snapshot)
|
||||
|
||||
Returns:
|
||||
Number of files successfully cleaned up
|
||||
"""
|
||||
cleaned_count = 0
|
||||
|
||||
for file_path in corrupted_files:
|
||||
@@ -588,13 +443,11 @@ def _cleanup_corrupted_files_selective(
|
||||
if os.path.islink(file_path):
|
||||
blob_path = os.path.realpath(file_path)
|
||||
|
||||
# Delete the symlink
|
||||
os.remove(file_path)
|
||||
logger.info(
|
||||
"Removed corrupted symlink: %s", os.path.basename(file_path)
|
||||
)
|
||||
|
||||
# Delete the blob (the actual corrupted data)
|
||||
if os.path.exists(blob_path):
|
||||
os.remove(blob_path)
|
||||
logger.info(
|
||||
@@ -603,7 +456,6 @@ def _cleanup_corrupted_files_selective(
|
||||
|
||||
cleaned_count += 1
|
||||
elif os.path.exists(file_path):
|
||||
# Not a symlink, just delete the file
|
||||
os.remove(file_path)
|
||||
logger.info("Removed corrupted file: %s", os.path.basename(file_path))
|
||||
cleaned_count += 1
|
||||
@@ -629,17 +481,7 @@ def _cleanup_corrupted_files_selective(
|
||||
def _cleanup_corrupted_model_cache(
|
||||
model_name_or_path: str, snapshot_dir: str, reason: str
|
||||
) -> None:
|
||||
"""
|
||||
Remove entire corrupted model cache directory to force a clean re-download.
|
||||
|
||||
This is used when we cannot selectively clean (e.g., missing shards, incomplete
|
||||
downloads with unknown affected files).
|
||||
|
||||
Args:
|
||||
model_name_or_path: Model identifier
|
||||
snapshot_dir: Path to the snapshot directory
|
||||
reason: Reason for cleanup
|
||||
"""
|
||||
# Full-cache delete: for when the affected files are unknown.
|
||||
# Navigate up to the model root directory: snapshots/hash -> snapshots -> model_root
|
||||
repo_folder = os.path.abspath(os.path.join(snapshot_dir, "..", ".."))
|
||||
|
||||
@@ -666,27 +508,10 @@ def ci_validate_and_cleanup_local_snapshot(
|
||||
found_local_snapshot_dir: str,
|
||||
local_weight_files: List[str],
|
||||
) -> bool:
|
||||
"""
|
||||
CI-specific validation and cleanup for local model snapshots.
|
||||
|
||||
This function validates the local snapshot and performs automatic cleanup
|
||||
if corruption or missing files are detected. This behavior is only appropriate
|
||||
for CI environments where we want automatic recovery.
|
||||
|
||||
Args:
|
||||
model_name_or_path: Model identifier for logging
|
||||
found_local_snapshot_dir: Path to the local snapshot directory
|
||||
local_weight_files: List of weight file paths found in the snapshot
|
||||
|
||||
Returns:
|
||||
True if the snapshot is valid and can be used, False if it was invalid
|
||||
and cleanup was performed (caller should re-download)
|
||||
"""
|
||||
# Check for incomplete files and clean up if found
|
||||
"""Validate a local snapshot, cleaning it up on failure; False means re-download."""
|
||||
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_module.glob(os.path.join(blobs_dir, "*.incomplete"))
|
||||
@@ -704,14 +529,12 @@ def ci_validate_and_cleanup_local_snapshot(
|
||||
)
|
||||
return False
|
||||
|
||||
# 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 "
|
||||
@@ -722,8 +545,8 @@ def ci_validate_and_cleanup_local_snapshot(
|
||||
return False
|
||||
else:
|
||||
# Missing shards (not corruption) - let snapshot_download handle it.
|
||||
# IMPORTANT: Do NOT delete the entire cache here, as other processes
|
||||
# (TP/EP ranks) may already be loading weights from these files.
|
||||
# Other processes (TP/EP ranks) may already be loading these
|
||||
# files, so the whole cache must not be deleted here.
|
||||
log_info_on_rank0(
|
||||
logger,
|
||||
f"Validation failed for {model_name_or_path}: {error_msg}. "
|
||||
@@ -731,10 +554,8 @@ def ci_validate_and_cleanup_local_snapshot(
|
||||
)
|
||||
return False
|
||||
|
||||
# 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)
|
||||
# Include adapter_model.safetensors for LoRA adapters
|
||||
if base_name in [
|
||||
"model.safetensors",
|
||||
@@ -747,10 +568,8 @@ def ci_validate_and_cleanup_local_snapshot(
|
||||
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
|
||||
# Also validate single PyTorch .bin files
|
||||
elif base_name in [
|
||||
"pytorch_model.bin",
|
||||
"model.bin",
|
||||
@@ -774,23 +593,6 @@ def _validate_weights_after_download(
|
||||
allow_patterns: List[str],
|
||||
model_name_or_path: str,
|
||||
) -> bool:
|
||||
"""
|
||||
Validate downloaded weight files to catch corruption early.
|
||||
|
||||
This function validates safetensors files after download to catch
|
||||
corruption issues (truncated downloads, network errors, etc.) before
|
||||
model loading fails with cryptic errors. If corruption is found,
|
||||
the corrupted files are automatically cleaned up.
|
||||
|
||||
Args:
|
||||
hf_folder: Path to the downloaded model folder
|
||||
allow_patterns: Patterns used to match weight files
|
||||
model_name_or_path: Model identifier for error messages
|
||||
|
||||
Returns:
|
||||
True if all files are valid, False if corrupted files were found and cleaned up
|
||||
"""
|
||||
# Find all weight files that were downloaded
|
||||
weight_files: List[str] = []
|
||||
for pattern in allow_patterns:
|
||||
weight_files.extend(glob_module.glob(os.path.join(hf_folder, pattern)))
|
||||
@@ -798,7 +600,6 @@ def _validate_weights_after_download(
|
||||
if not weight_files:
|
||||
return True # No weight files to validate
|
||||
|
||||
# Validate weight files (safetensors and .bin)
|
||||
corrupted_files = []
|
||||
for f in weight_files:
|
||||
if f.endswith(".safetensors") and os.path.exists(f):
|
||||
@@ -828,24 +629,6 @@ def _validate_weights_after_download(
|
||||
def _get_lock_file_path(
|
||||
model_name_or_path: str, cache_dir: Optional[str] = None
|
||||
) -> str:
|
||||
"""
|
||||
Generate a unique lock file path for download coordination.
|
||||
|
||||
In CI environments where multiple containers share an NFS-mounted HF cache,
|
||||
the lock file is placed on the shared cache directory so ALL containers
|
||||
coordinate on the same lock. This prevents cross-container .incomplete
|
||||
file race conditions.
|
||||
|
||||
Falls back to /dev/shm (container-local) for non-CI or when the cache
|
||||
dir is not accessible.
|
||||
|
||||
Args:
|
||||
model_name_or_path: Model identifier
|
||||
cache_dir: HF cache directory (None to use default)
|
||||
|
||||
Returns:
|
||||
Path to the lock file
|
||||
"""
|
||||
key_hash = hashlib.sha256(model_name_or_path.encode()).hexdigest()[:16]
|
||||
|
||||
# In CI, place lock on the shared HF cache directory so that ALL containers
|
||||
@@ -862,27 +645,13 @@ def _get_lock_file_path(
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
# Fallback to container-local lock
|
||||
if os.path.isdir("/dev/shm"):
|
||||
return f"/dev/shm/sglang_download_lock_{key_hash}"
|
||||
return f"/tmp/sglang_download_lock_{key_hash}"
|
||||
|
||||
|
||||
def _cleanup_incomplete_blobs(model_name_or_path: str, cache_dir: Optional[str]) -> int:
|
||||
"""
|
||||
Remove stale .incomplete files from the model's blobs directory.
|
||||
|
||||
This is lighter than _cleanup_corrupted_model_cache (which deletes the
|
||||
entire cache). We only remove .incomplete files so snapshot_download
|
||||
starts fresh on retry, preserving any successfully downloaded blobs.
|
||||
|
||||
Args:
|
||||
model_name_or_path: Model identifier (e.g., "meta-llama/Llama-2-7b-hf")
|
||||
cache_dir: HF cache directory (None to use default)
|
||||
|
||||
Returns:
|
||||
Number of .incomplete files removed
|
||||
"""
|
||||
# Only .incomplete files go, so retries keep the blobs already downloaded.
|
||||
try:
|
||||
import huggingface_hub.constants
|
||||
|
||||
@@ -929,30 +698,10 @@ def ci_download_with_validation_and_retry(
|
||||
revision: Optional[str],
|
||||
max_retries: int = 3,
|
||||
) -> str:
|
||||
"""
|
||||
CI-specific download with validation and automatic retry on corruption.
|
||||
"""Download weights, validating each attempt and retrying on corruption.
|
||||
|
||||
This function handles the download of model weights in CI environments,
|
||||
with automatic validation and retry logic for handling corrupted downloads.
|
||||
|
||||
Uses filelock.FileLock on the shared HF cache directory to coordinate
|
||||
downloads across all processes AND all containers sharing the same
|
||||
NFS-mounted cache. Only one process downloads at a time; others wait
|
||||
for the lock then use the cached result.
|
||||
|
||||
Args:
|
||||
model_name_or_path: The model name or path
|
||||
allow_patterns: The allowed patterns for weight files
|
||||
ignore_patterns: The patterns to filter out weight files
|
||||
cache_dir: The cache directory to store model weights
|
||||
revision: The revision of the model
|
||||
max_retries: Maximum number of download retries if corruption is detected
|
||||
|
||||
Returns:
|
||||
str: The path to the downloaded model weights
|
||||
|
||||
Raises:
|
||||
RuntimeError: If download fails after max_retries attempts
|
||||
Holds a filelock on the shared HF cache so that processes and containers on
|
||||
the same NFS mount take turns; the rest wait and reuse the cached result.
|
||||
"""
|
||||
import filelock
|
||||
import huggingface_hub.constants
|
||||
@@ -1104,18 +853,10 @@ def ci_download_with_validation_and_retry(
|
||||
|
||||
|
||||
def ci_validate_and_clean_hf_cache(model_path: str) -> None:
|
||||
"""
|
||||
Validate and clean corrupted safetensors files in HF cache before loading.
|
||||
"""Drop corrupted safetensors from the HF cache before a non-SGLang load.
|
||||
|
||||
This function is needed because HFRunner (used in tests) calls transformers'
|
||||
from_pretrained() directly, which bypasses SGLang's weight validation.
|
||||
Corrupted cached files can cause cryptic errors like "EOF while parsing"
|
||||
from safetensors.
|
||||
|
||||
Only runs in CI to avoid overhead for regular users.
|
||||
|
||||
Args:
|
||||
model_path: Model identifier (e.g., "meta-llama/Llama-2-7b")
|
||||
HFRunner calls transformers' from_pretrained() directly, which bypasses the
|
||||
validation in this module; a corrupted cache surfaces as "EOF while parsing".
|
||||
"""
|
||||
from sglang.utils import is_in_ci
|
||||
|
||||
|
||||
Reference in New Issue
Block a user