refactor: centralize model registry and improve dynamic model discovery

- Add model registry with metadata for all SeedVR2 models in different HF repos

- Support dynamic discovery of .safetensors and .gguf files

- Centralize constants and remove duplicate model lists

- Simplify download logic with clear manual fallback instructions
This commit is contained in:
Adrien Toupet
2025-07-24 11:54:43 -04:00
parent bc8f0aa6d0
commit 5c65c446e6
5 changed files with 65 additions and 73 deletions
+10 -2
View File
@@ -37,11 +37,16 @@ if parent_dir not in sys.path:
sys.path.insert(0, parent_dir)
'''
# Progressive import system with fallback
# ===== MODULE 0: Constants =====
if MODULES_AVAILABLE['downloads']:
from src.utils.constants import (
get_base_cache_dir,
)
# ===== MODULE 1: Downloads =====
if MODULES_AVAILABLE['downloads']:
from src.utils.downloads import (
download_weight,
get_base_cache_dir
)
@@ -111,8 +116,11 @@ if MODULES_AVAILABLE['comfyui_node']:
# Export all available functions
__all__ = [
# Constants
'get_base_cache_dir',
# Utils
'download_weight', 'get_base_cache_dir',
'download_weight',
# Memory Management
'get_vram_usage', 'clear_vram_cache', 'reset_vram_peak',
+2 -1
View File
@@ -20,6 +20,7 @@ import os
import gc
import torch
import time
from src.utils.constants import get_script_directory
from torchvision.transforms import Compose, Lambda, Normalize
@@ -37,7 +38,7 @@ except:
COMFYUI_AVAILABLE = False
pass
# Get script directory for embeddings
script_directory = os.path.dirname(os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
script_directory = get_script_directory()
# Import transforms and color fix
+2 -1
View File
@@ -18,6 +18,7 @@ Key Features:
import os
import time
import torch
from src.utils.constants import get_script_directory
from omegaconf import DictConfig, OmegaConf
# Import SafeTensors with fallback
@@ -36,7 +37,7 @@ from src.core.infer import VideoDiffusionInfer
from src.optimization.blockswap import apply_block_swap_to_dit
# Get script directory for config paths
script_directory = os.path.dirname(os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
script_directory = get_script_directory()
def configure_runner(model, base_cache_dir, preserve_vram=False, debug=False, block_swap_config=None, cached_runner=None):
+9 -12
View File
@@ -7,7 +7,10 @@ import time
import torch
from typing import Tuple, Dict, Any
from src.utils.downloads import download_weight, get_base_cache_dir
from src.utils.constants import get_base_cache_dir
from src.utils.downloads import download_weight
from src.utils.model_registry import get_available_models, DEFAULT_MODEL
from src.utils.constants import get_script_directory
from src.core.model_manager import configure_runner
from src.core.generation import generation_loop
from src.optimization.memory_manager import fast_model_cleanup, fast_ram_cleanup
@@ -22,7 +25,7 @@ from src.optimization.memory_manager import (
# Import ComfyUI progress reporting
from server import PromptServer
script_directory = os.path.dirname(os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
script_directory = get_script_directory()
class SeedVR2:
"""
@@ -53,17 +56,11 @@ class SeedVR2:
Dictionary defining input parameters, types, and validation
"""
return {
"required": {
"required": {
"images": ("IMAGE", ),
"model": ([
"seedvr2_ema_3b_fp16.safetensors",
"seedvr2_ema_3b_fp8_e4m3fn.safetensors",
"seedvr2_ema_7b_fp16.safetensors",
"seedvr2_ema_7b_fp8_e4m3fn.safetensors",
"seedvr2_ema_7b_sharp_fp16.safetensors",
"seedvr2_ema_7b_sharp_fp8_e4m3fn.safetensors"
], {
"default": "seedvr2_ema_3b_fp8_e4m3fn.safetensors"
"model": (get_available_models(), {
"default": DEFAULT_MODEL,
"tooltip": "Model variants with different sizes and precisions. Models will automatically download on first use. Additional models can be added to the ComfyUI models folder."
}),
"seed": ("INT", {
"default": 100,
+42 -57
View File
@@ -1,73 +1,58 @@
"""
Downloads utility module for SeedVR2
Handles model and VAE downloads from HuggingFace Hub
Extracted from: seedvr2.py (line 968-1015)
Handles model and VAE downloads from HuggingFace repositories
"""
import os
import urllib.error
from typing import Optional
from torchvision.datasets.utils import download_url
try:
import folder_paths
# Configuration des chemins
base_cache_dir = os.path.join(folder_paths.models_dir, "SEEDVR2")
# S'assurer que le dossier de cache existe
folder_paths.add_model_folder_path("seedvr2", os.path.join(folder_paths.models_dir, "SEEDVR2"))
except:
base_cache_dir = "./seedvr2_models"
from src.utils.model_registry import MODEL_REGISTRY, get_model_repo, DEFAULT_VAE
from src.utils.constants import get_base_cache_dir
def download_weight(model, model_dir=None):
# HuggingFace URL template
HUGGINGFACE_BASE_URL = "https://huggingface.co/{repo}/resolve/main/{filename}"
def download_weight(model: str, model_dir: Optional[str] = None) -> None:
"""
Télécharge un modèle SeedVR2 et son VAE associé depuis HuggingFace Hub
Download a SeedVR2 model and its associated VAE from HuggingFace Hub
Args:
model (str): Nom du fichier modèle à télécharger
(ex: "seedvr2_ema_3b_fp16.safetensors")
Gère automatiquement:
- Téléchargement du modèle principal
- Téléchargement du VAE avec fallbacks:
1. ema_vae_fp16.safetensors (priorité)
2. ema_vae_fp8_e4m3fn.safetensors (fallback)
3. ema_vae.pth (legacy fallback)
model: Model filename to download
model_dir: Optional custom directory for models
"""
if model_dir is None:
model_path = os.path.join(base_cache_dir, model)
vae_fp16_path = os.path.join(base_cache_dir, "ema_vae_fp16.safetensors")
cache_dir = base_cache_dir
else:
model_path = os.path.join(model_dir, model)
vae_fp16_path = os.path.join(model_dir, "ema_vae_fp16.safetensors")
cache_dir = model_dir
# Configuration HuggingFace
repo_id = "numz/SeedVR2_comfyUI"
base_url = f"https://huggingface.co/{repo_id}/resolve/main"
# Setup paths
cache_dir = model_dir or get_base_cache_dir()
model_path = os.path.join(cache_dir, model)
# 🚀 Téléchargement du modèle principal
# Download main model if not exists
if not os.path.exists(model_path):
print(f"📥 Downloading model: {model}")
download_url(f"{base_url}/{model}", cache_dir, filename=model)
print(f"✅ Downloaded: {model}")
# 🚀 Téléchargement du VAE avec stratégie de fallback
if not os.path.exists(vae_fp16_path):
print("📥 Downloading FP16 VAE SafeTensors...")
repo = get_model_repo(model)
url = HUGGINGFACE_BASE_URL.format(repo=repo, filename=model)
print(f"📥 Downloading {model} from {repo}...")
try:
download_url(f"{base_url}/ema_vae_fp16.safetensors", cache_dir, filename="ema_vae_fp16.safetensors")
print("✅ Downloaded: ema_vae_fp16.safetensors (FP16 SafeTensors)")
except Exception as e:
print(f"⚠️ FP16 SafeTensors VAE not available: {e}")
download_url(url, cache_dir, filename=model)
print(f"✅ Downloaded: {model}")
except (urllib.error.HTTPError, urllib.error.URLError) as e:
print(f"❌ Download failed: {e}")
print(f"📎 Please download manually from: https://huggingface.co/{repo}")
print(f" and place it in: {cache_dir}")
return
return
def get_base_cache_dir():
"""
Retourne le répertoire de cache base pour les modèles SeedVR2
Returns:
str: Chemin du répertoire de cache
"""
return base_cache_dir
# Download VAE if model is in registry
if model in MODEL_REGISTRY:
vae_path = os.path.join(cache_dir, DEFAULT_VAE)
if not os.path.exists(vae_path):
vae_repo = get_model_repo(DEFAULT_VAE)
vae_url = HUGGINGFACE_BASE_URL.format(repo=vae_repo, filename=DEFAULT_VAE)
print(f"📥 Downloading VAE: {DEFAULT_VAE}")
try:
download_url(vae_url, cache_dir, filename=DEFAULT_VAE)
print(f"✅ Downloaded: {DEFAULT_VAE}")
except (urllib.error.HTTPError, urllib.error.URLError) as e:
print(f"⚠️ VAE download failed: {e}")
print(f"📎 Please download VAE from: https://huggingface.co/{vae_repo}")
print(f" and place it in: {cache_dir}")