[diffusion] CI: fix LoRA downloading issues and respect offline flag (#15813)
This commit is contained in:
@@ -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")
|
||||||
|
|||||||
@@ -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"
|
||||||
|
|||||||
Reference in New Issue
Block a user