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:
+10
-2
@@ -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',
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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
@@ -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}")
|
||||
Reference in New Issue
Block a user