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:
gaclove
2025-07-16 15:41:37 +08:00
parent 4c7c11b2d6
commit a257203136
2 changed files with 99 additions and 41 deletions
+82
View File
@@ -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 {}
+17 -41
View File
@@ -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"