[diffusion] CI: add validation and cleanup for corrupted safetensors in multimodal loader (#13870)
Co-authored-by: Mick <mickjagger19@icloud.com> Co-authored-by: gemini-code-assist[bot] <176961590+gemini-code-assist[bot]@users.noreply.github.com>
This commit is contained in:
co-authored by
Mick
gemini-code-assist[bot]
parent
a2c388ba11
commit
83e7207763
@@ -115,6 +115,30 @@ def filter_files_not_needed_for_inference(hf_weights_files: list[str]) -> list[s
|
|||||||
_BAR_FORMAT = "{desc}: {percentage:3.0f}% Completed | {n_fmt}/{total_fmt} [{elapsed}<{remaining}, {rate_fmt}]\n" # noqa: E501
|
_BAR_FORMAT = "{desc}: {percentage:3.0f}% Completed | {n_fmt}/{total_fmt} [{elapsed}<{remaining}, {rate_fmt}]\n" # noqa: E501
|
||||||
|
|
||||||
|
|
||||||
|
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 file is valid, False if corrupted
|
||||||
|
"""
|
||||||
|
try:
|
||||||
|
with safe_open(file_path, framework="pt", device="cpu") as f:
|
||||||
|
_ = list(f.keys())
|
||||||
|
return True
|
||||||
|
except Exception as e:
|
||||||
|
logger.error(
|
||||||
|
"Corrupted safetensors file detected: %s - %s: %s",
|
||||||
|
file_path,
|
||||||
|
type(e).__name__,
|
||||||
|
str(e),
|
||||||
|
)
|
||||||
|
return False
|
||||||
|
|
||||||
|
|
||||||
def safetensors_weights_iterator(
|
def safetensors_weights_iterator(
|
||||||
hf_weights_files: list[str],
|
hf_weights_files: list[str],
|
||||||
to_cpu: bool = True,
|
to_cpu: bool = True,
|
||||||
@@ -124,6 +148,39 @@ def safetensors_weights_iterator(
|
|||||||
not torch.distributed.is_initialized() or torch.distributed.get_rank() == 0
|
not torch.distributed.is_initialized() or torch.distributed.get_rank() == 0
|
||||||
)
|
)
|
||||||
device = "cpu" if to_cpu else str(get_local_torch_device())
|
device = "cpu" if to_cpu else str(get_local_torch_device())
|
||||||
|
|
||||||
|
# Validate files before loading
|
||||||
|
corrupted_files = [st_file for st_file in hf_weights_files if not _validate_safetensors_file(st_file)]
|
||||||
|
|
||||||
|
if corrupted_files:
|
||||||
|
# Delete corrupted files (both symlink and blob if applicable)
|
||||||
|
for file_path in corrupted_files:
|
||||||
|
try:
|
||||||
|
if os.path.islink(file_path):
|
||||||
|
blob_path = os.path.realpath(file_path)
|
||||||
|
os.remove(file_path)
|
||||||
|
logger.info(
|
||||||
|
"Removed corrupted symlink: %s", os.path.basename(file_path)
|
||||||
|
)
|
||||||
|
if os.path.exists(blob_path):
|
||||||
|
os.remove(blob_path)
|
||||||
|
logger.info(
|
||||||
|
"Removed corrupted blob: %s", os.path.basename(blob_path)
|
||||||
|
)
|
||||||
|
elif os.path.isfile(file_path):
|
||||||
|
os.remove(file_path)
|
||||||
|
logger.info(
|
||||||
|
"Removed corrupted file: %s", os.path.basename(file_path)
|
||||||
|
)
|
||||||
|
except Exception as e:
|
||||||
|
logger.warning("Failed to remove corrupted file %s: %s", file_path, e)
|
||||||
|
|
||||||
|
raise RuntimeError(
|
||||||
|
f"Found {len(corrupted_files)} corrupted safetensors file(s). "
|
||||||
|
f"Files have been removed: {[os.path.basename(f) for f in corrupted_files]}. "
|
||||||
|
"Please retry - the files will be re-downloaded automatically."
|
||||||
|
)
|
||||||
|
|
||||||
for st_file in tqdm(
|
for st_file in tqdm(
|
||||||
hf_weights_files,
|
hf_weights_files,
|
||||||
desc="Loading safetensors checkpoint shards",
|
desc="Loading safetensors checkpoint shards",
|
||||||
|
|||||||
Reference in New Issue
Block a user