feat: introduce model_utils.py for model and LoRA management, enhancing model path retrieval and scanning functionality in nodes.py
This commit is contained in:
@@ -0,0 +1,82 @@
|
||||
# coding: utf-8
|
||||
|
||||
import json
|
||||
from pathlib import Path
|
||||
from typing import Dict, List
|
||||
|
||||
import folder_paths
|
||||
|
||||
|
||||
def get_model_base_path() -> Path:
|
||||
models_base = folder_paths.models_dir
|
||||
lightx2v_path = Path(models_base) / "lightx2v"
|
||||
if not lightx2v_path.exists():
|
||||
lightx2v_path.mkdir(parents=True, exist_ok=True)
|
||||
return lightx2v_path
|
||||
|
||||
|
||||
def scan_models() -> List[str]:
|
||||
models = []
|
||||
base_path = get_model_base_path()
|
||||
|
||||
if base_path.exists():
|
||||
for item in base_path.iterdir():
|
||||
if item.is_dir() and item.name != "loras":
|
||||
models.append(item.name)
|
||||
|
||||
models.sort()
|
||||
|
||||
return ["None"] + models if models else ["None"]
|
||||
|
||||
|
||||
def scan_loras() -> List[str]:
|
||||
loras = []
|
||||
base_path = get_model_base_path()
|
||||
loras_path = base_path / "loras"
|
||||
|
||||
if loras_path.exists():
|
||||
for item in loras_path.iterdir():
|
||||
if item.is_file():
|
||||
if item.suffix.lower() in [".safetensors", ".pt", ".pth", ".ckpt"]:
|
||||
loras.append(item.name)
|
||||
|
||||
loras.sort()
|
||||
|
||||
return ["None"] + loras if loras else ["None"]
|
||||
|
||||
|
||||
def get_model_full_path(model_name: str) -> str:
|
||||
if model_name == "None" or not model_name:
|
||||
return ""
|
||||
|
||||
base_path = get_model_base_path()
|
||||
model_path = base_path / model_name
|
||||
|
||||
if model_path.exists():
|
||||
return str(model_path)
|
||||
return ""
|
||||
|
||||
|
||||
def get_lora_full_path(lora_name: str) -> str:
|
||||
if lora_name == "None" or not lora_name:
|
||||
return ""
|
||||
|
||||
base_path = get_model_base_path()
|
||||
lora_path = base_path / "loras" / lora_name
|
||||
|
||||
if lora_path.exists():
|
||||
return str(lora_path)
|
||||
return ""
|
||||
|
||||
|
||||
def get_model_info(model_name: str) -> Dict:
|
||||
if model_name == "None" or not model_name:
|
||||
return {}
|
||||
|
||||
base_path = get_model_base_path()
|
||||
config_path = base_path / model_name / "config.json"
|
||||
|
||||
if config_path.exists():
|
||||
with open(config_path, "r") as f:
|
||||
return json.load(f)
|
||||
return {}
|
||||
@@ -2,10 +2,11 @@
|
||||
|
||||
import asyncio
|
||||
import gc
|
||||
import hashlib
|
||||
import json
|
||||
import logging
|
||||
import os
|
||||
import tempfile
|
||||
from typing import Any, Dict
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
@@ -14,17 +15,18 @@ from PIL import Image
|
||||
|
||||
from .bridge import ModularConfigManager, get_available_attn_ops, get_available_quant_ops
|
||||
from .lightx2v.lightx2v.infer import init_runner
|
||||
from .model_utils import get_lora_full_path, get_model_full_path, scan_loras, scan_models
|
||||
|
||||
|
||||
class LightX2VInferenceConfig:
|
||||
"""Basic inference configuration node."""
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
available_models = scan_models()
|
||||
|
||||
return {
|
||||
"required": {
|
||||
"model_cls": (["wan2.1", "wan2.1_audio", "wan2.1_distill", "hunyuan"], {"default": "wan2.1", "tooltip": "Model type"}),
|
||||
"model_path": ("STRING", {"default": "", "tooltip": "Model path"}),
|
||||
"model_name": (available_models, {"default": available_models[0], "tooltip": "Select model from available models"}),
|
||||
"task": (["t2v", "i2v"], {"default": "t2v", "tooltip": "Task type: text-to-video or image-to-video"}),
|
||||
"infer_steps": ("INT", {"default": 40, "min": 1, "max": 100, "tooltip": "Inference steps"}),
|
||||
"seed": ("INT", {"default": 42, "min": -1, "max": 2**32 - 1, "tooltip": "Random seed, -1 for random"}),
|
||||
@@ -52,9 +54,11 @@ class LightX2VInferenceConfig:
|
||||
CATEGORY = "LightX2V/Config"
|
||||
|
||||
def create_config(
|
||||
self, model_cls, model_path, task, infer_steps, seed, cfg_scale, sample_shift, height, width, video_length, fps, denoising_steps=""
|
||||
self, model_cls, model_name, task, infer_steps, seed, cfg_scale, sample_shift, height, width, video_length, fps, denoising_steps=""
|
||||
):
|
||||
"""Create basic inference configuration."""
|
||||
model_path = get_model_full_path(model_name)
|
||||
|
||||
config = {
|
||||
"model_cls": model_cls,
|
||||
"model_path": model_path,
|
||||
@@ -108,7 +112,6 @@ class LightX2VTeaCache:
|
||||
CATEGORY = "LightX2V/Config"
|
||||
|
||||
def create_config(self, enable, threshold, use_ret_steps):
|
||||
"""Create TeaCache configuration."""
|
||||
config = {
|
||||
"enable": enable,
|
||||
"threshold": threshold,
|
||||
@@ -118,11 +121,8 @@ class LightX2VTeaCache:
|
||||
|
||||
|
||||
class LightX2VQuantization:
|
||||
"""Quantization configuration node."""
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
# Get available quantization backends
|
||||
available_ops = get_available_quant_ops()
|
||||
quant_backends = []
|
||||
|
||||
@@ -169,7 +169,6 @@ class LightX2VMemoryOptimization:
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
# Get available attention types
|
||||
available_attn = get_available_attn_ops()
|
||||
attn_types = []
|
||||
|
||||
@@ -177,7 +176,6 @@ class LightX2VMemoryOptimization:
|
||||
if is_available:
|
||||
attn_types.append(op_name)
|
||||
|
||||
# Always include fallback
|
||||
if "torch_sdpa" not in attn_types:
|
||||
attn_types.append("torch_sdpa")
|
||||
|
||||
@@ -222,7 +220,6 @@ class LightX2VMemoryOptimization:
|
||||
lazy_load=False,
|
||||
unload_after_inference=False,
|
||||
):
|
||||
"""Create memory optimization configuration."""
|
||||
config = {
|
||||
"optimization_level": optimization_level,
|
||||
"attention_type": attention_type,
|
||||
@@ -239,8 +236,6 @@ class LightX2VMemoryOptimization:
|
||||
|
||||
|
||||
class LightX2VLightweightVAE:
|
||||
"""Lightweight VAE configuration node."""
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
@@ -256,7 +251,6 @@ class LightX2VLightweightVAE:
|
||||
CATEGORY = "LightX2V/Config"
|
||||
|
||||
def create_config(self, use_tiny_vae, use_tiling_vae):
|
||||
"""Create VAE configuration."""
|
||||
config = {
|
||||
"use_tiny_vae": use_tiny_vae,
|
||||
"use_tiling_vae": use_tiling_vae,
|
||||
@@ -265,13 +259,13 @@ class LightX2VLightweightVAE:
|
||||
|
||||
|
||||
class LightX2VLoRALoader:
|
||||
"""LoRA loader node that can be chained."""
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
available_loras = scan_loras()
|
||||
|
||||
return {
|
||||
"required": {
|
||||
"lora_path": ("STRING", {"default": "", "tooltip": "Path to the LoRA file"}),
|
||||
"lora_name": (available_loras, {"default": available_loras[0], "tooltip": "Select LoRA from available LoRAs"}),
|
||||
"strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 2.0, "step": 0.1, "tooltip": "LoRA strength"}),
|
||||
},
|
||||
"optional": {
|
||||
@@ -284,26 +278,22 @@ class LightX2VLoRALoader:
|
||||
FUNCTION = "load_lora"
|
||||
CATEGORY = "LightX2V/LoRA"
|
||||
|
||||
def load_lora(self, lora_path, strength, lora_chain=None):
|
||||
"""Load LoRA and chain with previous LoRAs."""
|
||||
# Initialize or extend the LoRA chain
|
||||
def load_lora(self, lora_name, strength, lora_chain=None):
|
||||
if lora_chain is None:
|
||||
lora_chain = []
|
||||
else:
|
||||
# Make a copy to avoid modifying the input
|
||||
lora_chain = lora_chain.copy()
|
||||
|
||||
# Add new LoRA configuration
|
||||
if lora_path and lora_path.strip():
|
||||
lora_config = {"path": lora_path.strip(), "strength": strength}
|
||||
lora_path = get_lora_full_path(lora_name)
|
||||
|
||||
if lora_path:
|
||||
lora_config = {"path": lora_path, "strength": strength}
|
||||
lora_chain.append(lora_config)
|
||||
|
||||
return (lora_chain,)
|
||||
|
||||
|
||||
class LightX2VConfigCombiner:
|
||||
"""Combines all configuration nodes into a single config object."""
|
||||
|
||||
def __init__(self):
|
||||
self.config_manager = ModularConfigManager()
|
||||
|
||||
@@ -336,8 +326,6 @@ class LightX2VConfigCombiner:
|
||||
vae_config=None,
|
||||
lora_chain=None,
|
||||
):
|
||||
"""Combine all configurations into a single config object."""
|
||||
# Collect all configurations
|
||||
configs = {
|
||||
"inference": inference_config,
|
||||
}
|
||||
@@ -351,10 +339,8 @@ class LightX2VConfigCombiner:
|
||||
if vae_config:
|
||||
configs["vae"] = vae_config
|
||||
|
||||
# Build final configuration
|
||||
config = self.config_manager.build_final_config(configs)
|
||||
|
||||
# Add LoRA configurations if provided
|
||||
if lora_chain:
|
||||
config.lora_configs = lora_chain
|
||||
|
||||
@@ -362,8 +348,6 @@ class LightX2VConfigCombiner:
|
||||
|
||||
|
||||
class LightX2VModularInference:
|
||||
"""Modular inference node that uses a combined configuration."""
|
||||
|
||||
def __init__(self):
|
||||
self._current_runner = None
|
||||
self._current_config_hash = None
|
||||
@@ -388,11 +372,6 @@ class LightX2VModularInference:
|
||||
CATEGORY = "LightX2V/Inference"
|
||||
|
||||
def _get_config_hash(self, config) -> str:
|
||||
"""Generate a hash for configuration to detect changes."""
|
||||
import hashlib
|
||||
import json
|
||||
|
||||
# Only hash model-related configs
|
||||
relevant_configs = {
|
||||
"model_cls": getattr(config, "model_cls", None),
|
||||
"model_path": getattr(config, "model_path", None),
|
||||
@@ -413,9 +392,6 @@ class LightX2VModularInference:
|
||||
audio=None,
|
||||
**kwargs,
|
||||
):
|
||||
"""Generate video using combined configuration."""
|
||||
|
||||
# Set environment variables
|
||||
os.environ["TOKENIZERS_PARALLELISM"] = "false"
|
||||
if "DTYPE" not in os.environ:
|
||||
os.environ["DTYPE"] = "BF16"
|
||||
|
||||
Reference in New Issue
Block a user