| """ |
| Download all v2 models to local cache. |
| |
| Total disk: ~35 GB |
| - SDXL 1.0 base (DreamShaper-XL turbo) ~7 GB |
| - SDXL refiner ~6 GB |
| - SDXL inpainting ~6 GB |
| - IP-Adapter XL + image encoder ~3 GB |
| - Stable Video Diffusion XT 1.1 ~10 GB |
| - ControlNet Canny / Depth / OpenPose (SDXL) ~3 GB (3 × 1 GB) |
| - controlnet-aux Annotators ~1 GB |
| """ |
| from __future__ import annotations |
|
|
| import logging |
| import os |
| from pathlib import Path |
|
|
| log = logging.getLogger(__name__) |
|
|
|
|
| def download_all() -> None: |
| from huggingface_hub import snapshot_download |
| from config.settings import settings |
|
|
| models = [ |
| ("SDXL DreamShaper-XL turbo", settings.t2i_model_id), |
| ("SDXL Inpainting", settings.inpaint_model_id), |
| ("IP-Adapter XL (image encoder)", settings.image_encoder_id), |
| ("Stable Video Diffusion XT 1.1", settings.svd_model_id), |
| ("ControlNet Canny SDXL", settings.controlnet_canny_id), |
| ("ControlNet Depth SDXL", settings.controlnet_depth_id), |
| ("ControlNet OpenPose SDXL", settings.controlnet_openpose_id), |
| ] |
| if settings.t2i_refiner_id: |
| models.append(("SDXL Refiner", settings.t2i_refiner_id)) |
|
|
| cache_dir = os.getenv("HF_HOME", str(Path.home() / ".cache" / "huggingface")) |
| log.info("Downloading %d models to cache: %s", len(models), cache_dir) |
|
|
| for name, model_id in models: |
| log.info("--- Downloading %s (%s) ---", name, model_id) |
| try: |
| snapshot_download( |
| repo_id=model_id, |
| cache_dir=cache_dir, |
| max_workers=4, |
| ) |
| log.info("✓ %s", name) |
| except Exception as exc: |
| log.error("✗ %s: %s", name, exc) |
|
|
| |
| log.info("--- Downloading IP-Adapter XL weights (base ViT-H) ---") |
| try: |
| from huggingface_hub import hf_hub_download |
| |
| |
| hf_hub_download( |
| repo_id=settings.ip_adapter_model_id, |
| filename=f"{settings.ip_adapter_subfolder}/{settings.ip_adapter_weight_name}", |
| cache_dir=cache_dir, |
| ) |
| log.info("✓ IP-Adapter XL base weights (ip-adapter_sdxl_vit-h.safetensors)") |
| except Exception as exc: |
| log.error("✗ IP-Adapter weights: %s", exc) |
|
|
| log.info("All downloads complete.") |
|
|
|
|
| if __name__ == "__main__": |
| logging.basicConfig(level=logging.INFO, |
| format="%(asctime)s | %(levelname)s | %(message)s") |
| download_all() |
|
|