fix: restrict cache validation behaviors to CI only (#14849)
This commit is contained in:
@@ -47,6 +47,7 @@ from sglang.srt.model_loader.weight_validation import (
|
|||||||
_validate_sharded_model,
|
_validate_sharded_model,
|
||||||
)
|
)
|
||||||
from sglang.srt.utils import find_local_repo_dir, log_info_on_rank0, print_warning_once
|
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__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
@@ -326,27 +327,32 @@ def _find_local_hf_snapshot_dir_unlocked(
|
|||||||
if not os.path.isdir(found_local_snapshot_dir):
|
if not os.path.isdir(found_local_snapshot_dir):
|
||||||
return None
|
return None
|
||||||
|
|
||||||
# Check for incomplete files and clean up if found
|
# Only perform cache validation and cleanup in CI to avoid
|
||||||
repo_folder = os.path.abspath(os.path.join(found_local_snapshot_dir, "..", ".."))
|
# unnecessary overhead for regular users
|
||||||
blobs_dir = os.path.join(repo_folder, "blobs")
|
if is_in_ci():
|
||||||
|
# Check for incomplete files and clean up if found
|
||||||
# Check for incomplete download markers
|
repo_folder = os.path.abspath(
|
||||||
incomplete_files = []
|
os.path.join(found_local_snapshot_dir, "..", "..")
|
||||||
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(
|
blobs_dir = os.path.join(repo_folder, "blobs")
|
||||||
model_name_or_path,
|
|
||||||
found_local_snapshot_dir,
|
# Check for incomplete download markers
|
||||||
f"Incomplete download detected ({len(incomplete_files)} incomplete files)",
|
incomplete_files = []
|
||||||
)
|
if os.path.isdir(blobs_dir):
|
||||||
return None
|
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] = []
|
local_weight_files: List[str] = []
|
||||||
try:
|
try:
|
||||||
@@ -366,53 +372,57 @@ def _find_local_hf_snapshot_dir_unlocked(
|
|||||||
)
|
)
|
||||||
local_weight_files = []
|
local_weight_files = []
|
||||||
|
|
||||||
# Validate sharded models and check for corruption
|
# Only perform cache validation and cleanup in CI
|
||||||
if local_weight_files:
|
if is_in_ci():
|
||||||
is_valid, error_msg, corrupted_files = _validate_sharded_model(
|
# Validate sharded models and check for corruption
|
||||||
found_local_snapshot_dir, local_weight_files
|
if local_weight_files:
|
||||||
)
|
is_valid, error_msg, corrupted_files = _validate_sharded_model(
|
||||||
if not is_valid:
|
found_local_snapshot_dir, local_weight_files
|
||||||
if corrupted_files:
|
)
|
||||||
# Selective cleanup: only remove corrupted files
|
if not is_valid:
|
||||||
log_info_on_rank0(
|
if corrupted_files:
|
||||||
logger,
|
# Selective cleanup: only remove corrupted files
|
||||||
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(
|
log_info_on_rank0(
|
||||||
logger,
|
logger,
|
||||||
f"Corrupted model file {base_name} for {model_name_or_path}. "
|
f"Found {len(corrupted_files)} corrupted file(s) for "
|
||||||
"Will selectively clean and re-download this file.",
|
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
|
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:
|
if len(local_weight_files) > 0:
|
||||||
log_info_on_rank0(
|
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}")
|
log_info_on_rank0(logger, f"Using model weights format {allow_patterns}")
|
||||||
|
|
||||||
# Retry loop for handling corrupted downloads
|
# Only perform validation and retry in CI to avoid overhead for regular users
|
||||||
for attempt in range(max_retries):
|
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(
|
hf_folder = snapshot_download(
|
||||||
model_name_or_path,
|
model_name_or_path,
|
||||||
allow_patterns=allow_patterns,
|
allow_patterns=allow_patterns,
|
||||||
@@ -580,32 +629,7 @@ def download_weights_from_hf(
|
|||||||
revision=revision,
|
revision=revision,
|
||||||
local_files_only=huggingface_hub.constants.HF_HUB_OFFLINE,
|
local_files_only=huggingface_hub.constants.HF_HUB_OFFLINE,
|
||||||
)
|
)
|
||||||
|
return hf_folder
|
||||||
# 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(
|
def download_safetensors_index_file_from_hf(
|
||||||
|
|||||||
Reference in New Issue
Block a user