[diffusion] CI: fix LoRA downloading issues and respect offline flag (#15813)

This commit is contained in:
Prozac614
2025-12-30 11:39:27 +08:00
committed by GitHub
parent 1e45320198
commit f253f43c9d
4 changed files with 136 additions and 28 deletions
+1 -1
View File
@@ -105,7 +105,7 @@ test = [
] ]
diffusion = [ diffusion = [
"diffusers @ git+https://github.com/huggingface/diffusers.git@6290fdfda40610ce7b99920146853614ba529c6e", "diffusers @ git+https://github.com/huggingface/diffusers.git@6290fdfda40610ce7b99920146853614ba529c6e",
"opencv-python==4.10.0.84", "opencv-python-headless==4.10.0.84",
"imageio==2.36.0", "imageio==2.36.0",
"imageio-ffmpeg==0.5.1", "imageio-ffmpeg==0.5.1",
"PyYAML==6.0.1", "PyYAML==6.0.1",
@@ -21,6 +21,7 @@
import contextlib import contextlib
import json import json
import os import os
import time
from functools import reduce from functools import reduce
from pathlib import Path from pathlib import Path
from typing import Any, Optional, cast from typing import Any, Optional, cast
@@ -29,6 +30,13 @@ from diffusers.loaders.lora_base import (
_best_guess_weight_name, # watch out for potetential removal from diffusers _best_guess_weight_name, # watch out for potetential removal from diffusers
) )
from huggingface_hub import snapshot_download from huggingface_hub import snapshot_download
from huggingface_hub.errors import (
LocalEntryNotFoundError,
RepositoryNotFoundError,
RevisionNotFoundError,
)
from requests.exceptions import ConnectionError as RequestsConnectionError
from requests.exceptions import RequestException
from transformers import AutoConfig, PretrainedConfig from transformers import AutoConfig, PretrainedConfig
from transformers.models.auto.modeling_auto import MODEL_FOR_CAUSAL_LM_MAPPING_NAMES from transformers.models.auto.modeling_auto import MODEL_FOR_CAUSAL_LM_MAPPING_NAMES
@@ -208,15 +216,20 @@ def maybe_download_lora(
Returns: Returns:
Local path to the model Local path to the model
""" """
allow_patterns = ["*.json", "*.safetensors", "*.bin"]
local_path = maybe_download_model(model_name_or_path, local_dir, download) local_path = maybe_download_model(
model_name_or_path,
local_dir,
download,
is_lora=True,
allow_patterns=allow_patterns,
)
# return directly if local_path is a file # return directly if local_path is a file
if os.path.isfile(local_path): if os.path.isfile(local_path):
return local_path return local_path
weight_name = _best_guess_weight_name( weight_name = _best_guess_weight_name(local_path, file_extension=".safetensors")
model_name_or_path, file_extension=".safetensors"
)
return os.path.join(local_path, weight_name) return os.path.join(local_path, weight_name)
@@ -349,7 +362,11 @@ def maybe_download_model_index(model_name_or_path: str) -> dict[str, Any]:
def maybe_download_model( def maybe_download_model(
model_name_or_path: str, local_dir: str | None = None, download: bool = True model_name_or_path: str,
local_dir: str | None = None,
download: bool = True,
is_lora: bool = False,
allow_patterns: list[str] | None = None,
) -> str: ) -> str:
""" """
Check if the model path is a Hugging Face Hub model ID and download it if needed. Check if the model path is a Hugging Face Hub model ID and download it if needed.
@@ -358,6 +375,7 @@ def maybe_download_model(
model_name_or_path: Local path or Hugging Face Hub model ID model_name_or_path: Local path or Hugging Face Hub model ID
local_dir: Local directory to save the model local_dir: Local directory to save the model
download: Whether to download the model from Hugging Face Hub download: Whether to download the model from Hugging Face Hub
is_lora: If True, skip model completeness verification (LoRA models don't have transformer/vae directories)
Returns: Returns:
Local path to the model Local path to the model
@@ -374,9 +392,9 @@ def maybe_download_model(
and os.path.exists(vae_dir) and os.path.exists(vae_dir)
) )
# If the path exists locally, verify it's complete # 1. Local path check: if path exists locally, verify it's complete (skip for LoRA)
if os.path.exists(model_name_or_path): if os.path.exists(model_name_or_path):
if _verify_model_complete(model_name_or_path): if is_lora or _verify_model_complete(model_name_or_path):
logger.info("Model already exists locally and is complete") logger.info("Model already exists locally and is complete")
return model_name_or_path return model_name_or_path
else: else:
@@ -386,19 +404,73 @@ def maybe_download_model(
model_name_or_path, model_name_or_path,
) )
# Otherwise, assume it's a HF Hub model ID and try to download it # 2. Cache-first strategy (Fast Path)
# Try to read from HF cache without network access
try: try:
logger.info( logger.info(
"Downloading model snapshot from HF Hub for %s...", model_name_or_path "Checking for cached model in HF Hub cache for %s...", model_name_or_path
)
local_path = snapshot_download(
repo_id=model_name_or_path,
ignore_patterns=["*.onnx", "*.msgpack"],
local_dir=local_dir,
local_files_only=True,
resume_download=True,
max_workers=8,
etag_timeout=60,
)
if is_lora or _verify_model_complete(local_path):
logger.info("Found complete model in cache at %s", local_path)
return str(local_path)
else:
# Model found in cache but incomplete
if not download:
raise ValueError(
f"Model {model_name_or_path} found in cache but is incomplete and download=False."
)
logger.info(
"Model found in cache but incomplete, will download from HF Hub"
)
except LocalEntryNotFoundError:
if not download:
raise ValueError(
f"Model {model_name_or_path} not found in local cache and download=False."
)
logger.info("Model not found in cache, will download from HF Hub")
except Exception as e:
logger.warning(
"Unexpected error while checking cache for %s: %s, will attempt download",
model_name_or_path,
e,
)
if not download:
raise ValueError(
f"Error checking cache for {model_name_or_path} and download=False: {e}"
) from e
# 3. Download strategy (with retry mechanism)
MAX_RETRIES = 5
for attempt in range(MAX_RETRIES):
try:
logger.info(
"Downloading model snapshot from HF Hub for %s (attempt %d/%d)...",
model_name_or_path,
attempt + 1,
MAX_RETRIES,
) )
with get_lock(model_name_or_path).acquire(poll_interval=2): with get_lock(model_name_or_path).acquire(poll_interval=2):
local_path = snapshot_download( local_path = snapshot_download(
repo_id=model_name_or_path, repo_id=model_name_or_path,
ignore_patterns=["*.onnx", "*.msgpack"], ignore_patterns=["*.onnx", "*.msgpack"],
allow_patterns=allow_patterns,
local_dir=local_dir, local_dir=local_dir,
resume_download=True,
max_workers=8,
etag_timeout=120,
) )
# Verify downloaded model is complete
if not _verify_model_complete(local_path): # Verify downloaded model is complete (skip for LoRA)
if not is_lora and not _verify_model_complete(local_path):
logger.warning( logger.warning(
"Downloaded model at %s is incomplete, retrying with force_download=True", "Downloaded model at %s is incomplete, retrying with force_download=True",
local_path, local_path,
@@ -408,11 +480,41 @@ def maybe_download_model(
repo_id=model_name_or_path, repo_id=model_name_or_path,
ignore_patterns=["*.onnx", "*.msgpack"], ignore_patterns=["*.onnx", "*.msgpack"],
local_dir=local_dir, local_dir=local_dir,
resume_download=True,
max_workers=8,
etag_timeout=60,
force_download=True, force_download=True,
) )
if not _verify_model_complete(local_path):
raise ValueError(
f"Downloaded model at {local_path} is still incomplete after forced re-download. "
"The model repository may be missing required components (model_index.json, transformer/, or vae/)."
)
logger.info("Downloaded model to %s", local_path) logger.info("Downloaded model to %s", local_path)
return str(local_path) return str(local_path)
except (RepositoryNotFoundError, RevisionNotFoundError) as e:
raise ValueError(
f"Model or revision not found at {model_name_or_path}. "
f"Please check the model ID or ensure you have access to the repository. Error: {e}"
) from e
except (RequestException, RequestsConnectionError) as e:
if attempt == MAX_RETRIES - 1:
raise ValueError(
f"Could not find model at {model_name_or_path} and failed to download from HF Hub "
f"after {MAX_RETRIES} attempts due to network error: {e}"
) from e
wait_time = 2**attempt
logger.warning(
"Download failed (attempt %d/%d) due to network error: %s. "
"Retrying in %d seconds...",
attempt + 1,
MAX_RETRIES,
e,
wait_time,
)
time.sleep(wait_time)
except Exception as e: except Exception as e:
raise ValueError( raise ValueError(
f"Could not find model at {model_name_or_path} and failed to download from HF Hub: {e}" f"Could not find model at {model_name_or_path} and failed to download from HF Hub: {e}"
@@ -198,7 +198,12 @@ def run_pytest(files, filter_expr=None):
is_flaky_ci_assertion = "SafetensorError" in full_output is_flaky_ci_assertion = "SafetensorError" in full_output
if not (is_perf_assertion or is_flaky_ci_assertion): is_oom_error = (
"out of memory" in full_output.lower()
or "oom killer" in full_output.lower()
)
if not (is_perf_assertion or is_flaky_ci_assertion or is_oom_error):
return returncode return returncode
logger.info(f"Max retry exceeded") logger.info(f"Max retry exceeded")
+1
View File
@@ -96,6 +96,7 @@ fi
# Clean up existing installations # Clean up existing installations
$PIP_UNINSTALL_CMD sgl-kernel sglang $PIP_UNINSTALL_SUFFIX || true $PIP_UNINSTALL_CMD sgl-kernel sglang $PIP_UNINSTALL_SUFFIX || true
$PIP_UNINSTALL_CMD flashinfer-python flashinfer-cubin flashinfer-jit-cache $PIP_UNINSTALL_SUFFIX || true $PIP_UNINSTALL_CMD flashinfer-python flashinfer-cubin flashinfer-jit-cache $PIP_UNINSTALL_SUFFIX || true
$PIP_UNINSTALL_CMD opencv-python opencv-python-headless $PIP_UNINSTALL_SUFFIX || true
# Install the main package # Install the main package
EXTRAS="dev" EXTRAS="dev"