From d0d24a18cbc319d29e76d4c7d58758bdbb64bc38 Mon Sep 17 00:00:00 2001 From: Adrien Toupet Date: Sat, 20 Sep 2025 01:53:25 -0400 Subject: [PATCH] Refactor: Replace dynamic relative imports with model class registry for cross-platform compatibility. Minor: Adjust ASCII logo width for better screen fit --- configs_3b/main.yaml | 10 ++------- configs_7b/main.yaml | 10 ++------- src/common/config.py | 39 ++++++++++++---------------------- src/interfaces/comfyui_node.py | 14 ++++++------ src/utils/model_registry.py | 14 +++++++++++- 5 files changed, 37 insertions(+), 50 deletions(-) diff --git a/configs_3b/main.yaml b/configs_3b/main.yaml index 0f02859..ebd5adf 100644 --- a/configs_3b/main.yaml +++ b/configs_3b/main.yaml @@ -5,10 +5,7 @@ __object__: dit: model: __object__: - path: - - "custom_nodes.ComfyUI-SeedVR2_VideoUpscaler.src.models.dit_3b.nadit" - - "ComfyUI.custom_nodes.ComfyUI-SeedVR2_VideoUpscaler.src.models.dit_3b.nadit" - - "src.models.dit_3b.nadit" + path: "dit_3b.nadit" name: "NaDiT" args: "as_params" vid_in_channels: 33 @@ -48,10 +45,7 @@ ema: vae: model: __object__: - path: - - "custom_nodes.ComfyUI-SeedVR2_VideoUpscaler.src.models.video_vae_v3.modules.attn_video_vae" - - "ComfyUI.custom_nodes.ComfyUI-SeedVR2_VideoUpscaler.src.models.video_vae_v3.modules.attn_video_vae" - - "src.models.video_vae_v3.modules.attn_video_vae" + path: "video_vae_v3.modules.attn_video_vae" name: "VideoAutoencoderKLWrapper" args: "as_params" freeze_encoder: False diff --git a/configs_7b/main.yaml b/configs_7b/main.yaml index fda7dc3..ac66e9a 100644 --- a/configs_7b/main.yaml +++ b/configs_7b/main.yaml @@ -5,10 +5,7 @@ __object__: dit: model: __object__: - path: - - "custom_nodes.ComfyUI-SeedVR2_VideoUpscaler.src.models.dit_7b.nadit" - - "ComfyUI.custom_nodes.ComfyUI-SeedVR2_VideoUpscaler.src.models.dit_7b.nadit" - - "src.models.dit_7b.nadit" + path: "dit_7b.nadit" name: "NaDiT" args: "as_params" vid_in_channels: 33 @@ -45,10 +42,7 @@ ema: vae: model: __object__: - path: - - "custom_nodes.ComfyUI-SeedVR2_VideoUpscaler.src.models.video_vae_v3.modules.attn_video_vae" - - "ComfyUI.custom_nodes.ComfyUI-SeedVR2_VideoUpscaler.src.models.video_vae_v3.modules.attn_video_vae" - - "src.models.video_vae_v3.modules.attn_video_vae" + path: "video_vae_v3.modules.attn_video_vae" name: "VideoAutoencoderKLWrapper" args: "as_params" freeze_encoder: False diff --git a/src/common/config.py b/src/common/config.py index 9a07b76..1020e55 100644 --- a/src/common/config.py +++ b/src/common/config.py @@ -19,6 +19,7 @@ Configuration utility functions import importlib from typing import Any, Callable, List, Union from omegaconf import DictConfig, ListConfig, OmegaConf +from ..utils.model_registry import MODEL_CLASSES try: OmegaConf.register_new_resolver("eval", eval) @@ -87,40 +88,26 @@ def resolve_inheritance(config: Union[DictConfig, ListConfig]) -> Any: return config -def import_item(path: Union[str, List[str]], name: str) -> Any: +def import_item(path: str, name: str) -> Any: """ - Import a python item with fallback support. + Import a python item, checking model registry first. Args: - path: Single path string or list of paths to try (fallback order) + path: Module path name: Class/function name to import Returns: Imported object - - Example: - import_item("path.to.file", "MyClass") -> MyClass - import_item(["path1.to.file", "path2.to.file"], "MyClass") -> MyClass (first working path) - """ - if isinstance(path, str): - # Single path - original behavior + """ + # Simple lookup with path as key + if path in MODEL_CLASSES: + return MODEL_CLASSES[path] + + # Fallback to dynamic import for everything else + try: return getattr(importlib.import_module(path), name) - - elif isinstance(path, (list, ListConfig)): - # Multiple paths - try each until one works - last_error = None - for single_path in path: - try: - return getattr(importlib.import_module(single_path), name) - except ImportError as e: - last_error = e - continue - - # If we get here, none of the paths worked - raise ImportError(f"Could not import '{name}' from any of the paths: {path}. Last error: {last_error}") - - else: - raise ValueError(f"Path must be string or list of strings, got: {type(path)}") + except (ImportError, AttributeError) as e: + raise ImportError(f"Could not import '{name}' from '{path}': {e}") def create_object(config: DictConfig) -> Any: diff --git a/src/interfaces/comfyui_node.py b/src/interfaces/comfyui_node.py index 5838c1d..4d5fa6b 100644 --- a/src/interfaces/comfyui_node.py +++ b/src/interfaces/comfyui_node.py @@ -222,13 +222,13 @@ class SeedVR2: debug.start_timer("total_execution", force=True) debug.log("", category="none", force=True) - debug.log(" ╔══════════════════════════════════════════════════════════════╗", category="none", force=True) - debug.log(" ║ ███████ ███████ ███████ ██████ ██ ██ ██████ ██████ ║", category="none", force=True) - debug.log(" ║ ██ ██ ██ ██ ██ ██ ██ ██ ██ ██ ║", category="none", force=True) - debug.log(" ║ ███████ █████ █████ ██ ██ ██ ██ ██████ █████ ║", category="none", force=True) - debug.log(" ║ ██ ██ ██ ██ ██ ██ ██ ██ ██ ██ ║", category="none", force=True) - debug.log(" ║ ███████ ███████ ███████ ██████ ████ ██ ██ ███████ ║", category="none", force=True) - debug.log(" ╚══════════════════════════════════════════════════════════════╝", category="none", force=True) + debug.log(" ╔══════════════════════════════════════════════════════════╗", category="none", force=True) + debug.log(" ║ ███████ ███████ ███████ ██████ ██ ██ ██████ ███████ ║", category="none", force=True) + debug.log(" ║ ██ ██ ██ ██ ██ ██ ██ ██ ██ ██ ║", category="none", force=True) + debug.log(" ║ ███████ █████ █████ ██ ██ ██ ██ ██████ █████ ║", category="none", force=True) + debug.log(" ║ ██ ██ ██ ██ ██ ██ ██ ██ ██ ██ ║", category="none", force=True) + debug.log(" ║ ███████ ███████ ███████ ██████ ████ ██ ██ ███████ ║", category="none", force=True) + debug.log(" ╚══════════════════════════════════════════════════════════╝", category="none", force=True) debug.log("", category="none", force=True) debug.log("━━━━━━━━━ Model Preparation ━━━━━━━━━", category="none") diff --git a/src/utils/model_registry.py b/src/utils/model_registry.py index f7901ac..f698515 100644 --- a/src/utils/model_registry.py +++ b/src/utils/model_registry.py @@ -3,10 +3,22 @@ Model Registry for SeedVR2 Central registry for model definitions, repositories, and metadata """ -from typing import Dict, List, Optional +from typing import Dict, List, Optional, Any from dataclasses import dataclass from .constants import SEEDVR2_MODEL_TYPE, is_supported_model_file, get_base_cache_dir +# Model class imports using relative imports +from ..models.dit_3b.nadit import NaDiT as NaDiT3B +from ..models.dit_7b.nadit import NaDiT as NaDiT7B +from ..models.video_vae_v3.modules.attn_video_vae import VideoAutoencoderKLWrapper + +# Model classes - simple registry with clear keys +MODEL_CLASSES = { + "dit_3b.nadit": NaDiT3B, + "dit_7b.nadit": NaDiT7B, + "video_vae_v3.modules.attn_video_vae": VideoAutoencoderKLWrapper, +} + @dataclass class ModelInfo: """Model metadata"""