[misc] Add a comment style rule to .claude/rules (#35597)

This commit is contained in:
Liangsheng Yin
2026-08-19 18:52:48 -07:00
committed by GitHub
parent f65961844a
commit a49560ce50
2 changed files with 183 additions and 278 deletions
@@ -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