refactor: enhance configuration management by introducing LightX2VConfigCombiner for streamlined config handling and updating LightX2VModularInference to utilize combined configurations
This commit is contained in:
@@ -222,7 +222,6 @@ class ModularConfigManager:
|
||||
if is_available:
|
||||
available.append(op_name)
|
||||
|
||||
# Always include fallback
|
||||
if "torch_sdpa" not in available:
|
||||
available.append("torch_sdpa")
|
||||
|
||||
@@ -245,7 +244,6 @@ class ModularConfigManager:
|
||||
"""Apply basic inference configuration."""
|
||||
updates = {}
|
||||
|
||||
# Model settings
|
||||
if "model_cls" in config:
|
||||
updates["model_cls"] = config["model_cls"]
|
||||
if "model_path" in config:
|
||||
@@ -253,7 +251,6 @@ class ModularConfigManager:
|
||||
if "task" in config:
|
||||
updates["task"] = config["task"]
|
||||
|
||||
# Inference parameters
|
||||
if "infer_steps" in config:
|
||||
updates["infer_steps"] = config["infer_steps"]
|
||||
if "seed" in config and config["seed"] != -1:
|
||||
@@ -264,7 +261,6 @@ class ModularConfigManager:
|
||||
if "sample_shift" in config:
|
||||
updates["sample_shift"] = config["sample_shift"]
|
||||
|
||||
# Video parameters
|
||||
if "height" in config:
|
||||
updates["target_height"] = config["height"]
|
||||
if "width" in config:
|
||||
@@ -300,13 +296,11 @@ class ModularConfigManager:
|
||||
"""Apply quantization configuration."""
|
||||
updates = {}
|
||||
|
||||
# DIT quantization
|
||||
dit_scheme = config.get("dit_precision", "bf16")
|
||||
updates["dit_quant_scheme"] = dit_scheme
|
||||
if dit_scheme != "bf16":
|
||||
updates["dit_quantized_ckpt"] = os.path.join(model_path, dit_scheme)
|
||||
|
||||
# T5 quantization
|
||||
t5_scheme = config.get("t5_precision", "bf16")
|
||||
updates["t5_quant_scheme"] = t5_scheme
|
||||
updates["t5_quantized"] = t5_scheme != "bf16"
|
||||
@@ -314,7 +308,6 @@ class ModularConfigManager:
|
||||
t5_path = os.path.join(model_path, t5_scheme)
|
||||
updates["t5_quantized_ckpt"] = os.path.join(t5_path, f"models_t5_umt5-xxl-enc-{t5_scheme}.pth")
|
||||
|
||||
# CLIP quantization
|
||||
clip_scheme = config.get("clip_precision", "fp16")
|
||||
updates["clip_quant_scheme"] = clip_scheme
|
||||
updates["clip_quantized"] = clip_scheme != "fp16"
|
||||
@@ -322,11 +315,9 @@ class ModularConfigManager:
|
||||
clip_path = os.path.join(model_path, clip_scheme)
|
||||
updates["clip_quantized_ckpt"] = os.path.join(clip_path, f"clip-{clip_scheme}.pth")
|
||||
|
||||
# Quantization backend
|
||||
quant_backend = config.get("quant_backend", "vllm")
|
||||
updates["quant_op"] = quant_backend
|
||||
|
||||
# Determine mm_type based on quantization settings
|
||||
if dit_scheme != "bf16":
|
||||
if quant_backend == "vllm":
|
||||
mm_type = f"W-{dit_scheme}-channel-sym-A-{dit_scheme}-channel-sym-dynamic-Vllm"
|
||||
@@ -344,7 +335,6 @@ class ModularConfigManager:
|
||||
else:
|
||||
updates["mm_config"] = {"mm_type": "Default"}
|
||||
|
||||
# Precision mode
|
||||
updates["precision_mode"] = config.get("sensitive_layers_precision", "fp32")
|
||||
|
||||
return updates
|
||||
@@ -408,14 +398,13 @@ class ModularConfigManager:
|
||||
"""Build final configuration from module configs."""
|
||||
final_config = copy.deepcopy(self.base_config)
|
||||
|
||||
# Apply configurations in order
|
||||
if "inference" in configs:
|
||||
final_config.update(self.apply_inference_config(configs["inference"]))
|
||||
|
||||
if "teacache" in configs:
|
||||
teacache_updates = self.apply_teacache_config(
|
||||
configs["teacache"],
|
||||
final_config, # Pass current config for coefficient calculation
|
||||
final_config,
|
||||
)
|
||||
final_config.update(teacache_updates)
|
||||
|
||||
@@ -431,13 +420,11 @@ class ModularConfigManager:
|
||||
model_path = final_config.get("model_path", "")
|
||||
final_config.update(self.apply_vae_config(configs["vae"], model_path))
|
||||
|
||||
# Load model config if exists
|
||||
model_config_path = os.path.join(final_config["model_path"], "config.json")
|
||||
if os.path.exists(model_config_path):
|
||||
try:
|
||||
with open(model_config_path, "r") as f:
|
||||
model_config = json.load(f)
|
||||
# Model config has lower priority than user configs
|
||||
for key, value in model_config.items():
|
||||
if key not in final_config or final_config[key] is None:
|
||||
final_config[key] = value
|
||||
|
||||
@@ -281,24 +281,19 @@ class LightX2VLoRALoader:
|
||||
return (lora_chain,)
|
||||
|
||||
|
||||
class LightX2VModularInference:
|
||||
"""Modular inference node that combines all configurations."""
|
||||
class LightX2VConfigCombiner:
|
||||
"""Combines all configuration nodes into a single config object."""
|
||||
|
||||
def __init__(self):
|
||||
self.config_manager = ModularConfigManager()
|
||||
self._current_runner = None
|
||||
self._current_config_hash = None
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"inference_config": ("INFERENCE_CONFIG", {"tooltip": "Basic inference configuration"}),
|
||||
"prompt": ("STRING", {"multiline": True, "default": "", "tooltip": "Generation prompt"}),
|
||||
"negative_prompt": ("STRING", {"multiline": True, "default": "", "tooltip": "Negative prompt"}),
|
||||
},
|
||||
"optional": {
|
||||
"image": ("IMAGE", {"tooltip": "Input image for i2v task"}),
|
||||
"teacache_config": ("TEACACHE_CONFIG", {"tooltip": "TeaCache configuration"}),
|
||||
"quantization_config": ("QUANT_CONFIG", {"tooltip": "Quantization configuration"}),
|
||||
"memory_config": ("MEMORY_CONFIG", {"tooltip": "Memory optimization configuration"}),
|
||||
@@ -307,51 +302,21 @@ class LightX2VModularInference:
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("IMAGE",)
|
||||
RETURN_NAMES = ("images",)
|
||||
FUNCTION = "generate"
|
||||
CATEGORY = "LightX2V/Inference"
|
||||
RETURN_TYPES = ("COMBINED_CONFIG",)
|
||||
RETURN_NAMES = ("combined_config",)
|
||||
FUNCTION = "combine_configs"
|
||||
CATEGORY = "LightX2V/Config"
|
||||
|
||||
def _get_config_hash(self, configs: Dict[str, Any]) -> str:
|
||||
"""Generate a hash for configuration to detect changes."""
|
||||
import hashlib
|
||||
import json
|
||||
|
||||
# Only hash model-related configs
|
||||
relevant_configs = {
|
||||
"model_cls": configs.get("inference", {}).get("model_cls"),
|
||||
"model_path": configs.get("inference", {}).get("model_path"),
|
||||
"quantization": configs.get("quantization"),
|
||||
"memory_lazy_load": configs.get("memory", {}).get("lazy_load"),
|
||||
}
|
||||
|
||||
config_str = json.dumps(relevant_configs, sort_keys=True)
|
||||
return hashlib.md5(config_str.encode()).hexdigest()
|
||||
|
||||
def generate(
|
||||
def combine_configs(
|
||||
self,
|
||||
inference_config,
|
||||
prompt,
|
||||
negative_prompt,
|
||||
image=None,
|
||||
teacache_config=None,
|
||||
quantization_config=None,
|
||||
memory_config=None,
|
||||
vae_config=None,
|
||||
lora_chain=None,
|
||||
**kwargs,
|
||||
):
|
||||
"""Generate video using modular configuration."""
|
||||
|
||||
# Set environment variables
|
||||
os.environ["TOKENIZERS_PARALLELISM"] = "false"
|
||||
if "DTYPE" not in os.environ:
|
||||
os.environ["DTYPE"] = "BF16"
|
||||
if "ENABLE_GRAPH_MODE" not in os.environ:
|
||||
os.environ["ENABLE_GRAPH_MODE"] = "false"
|
||||
if "ENABLE_PROFILING_DEBUG" not in os.environ:
|
||||
os.environ["ENABLE_PROFILING_DEBUG"] = "true"
|
||||
|
||||
"""Combine all configurations into a single config object."""
|
||||
# Collect all configurations
|
||||
configs = {
|
||||
"inference": inference_config,
|
||||
@@ -373,50 +338,126 @@ class LightX2VModularInference:
|
||||
if lora_chain:
|
||||
config.lora_configs = lora_chain
|
||||
|
||||
# Add prompt and negative prompt
|
||||
return (config,)
|
||||
|
||||
|
||||
class LightX2VModularInference:
|
||||
"""Modular inference node that uses a combined configuration."""
|
||||
|
||||
def __init__(self):
|
||||
self._current_runner = None
|
||||
self._current_config_hash = None
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"combined_config": ("COMBINED_CONFIG", {"tooltip": "Combined configuration from config combiner"}),
|
||||
"prompt": ("STRING", {"multiline": True, "default": "", "tooltip": "Generation prompt"}),
|
||||
"negative_prompt": ("STRING", {"multiline": True, "default": "", "tooltip": "Negative prompt"}),
|
||||
},
|
||||
"optional": {
|
||||
"image": ("IMAGE", {"tooltip": "Input image for i2v task"}),
|
||||
"audio": ("AUDIO", {"tooltip": "Input audio for audio-driven generation"}),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("IMAGE",)
|
||||
RETURN_NAMES = ("images",)
|
||||
FUNCTION = "generate"
|
||||
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),
|
||||
"dit_quantized": getattr(config, "dit_quantized", False),
|
||||
"t5_quantized": getattr(config, "t5_quantized", False),
|
||||
"lazy_load": getattr(config, "lazy_load", False),
|
||||
}
|
||||
|
||||
config_str = json.dumps(relevant_configs, sort_keys=True)
|
||||
return hashlib.md5(config_str.encode()).hexdigest()
|
||||
|
||||
def generate(
|
||||
self,
|
||||
combined_config,
|
||||
prompt,
|
||||
negative_prompt,
|
||||
image=None,
|
||||
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"
|
||||
if "ENABLE_GRAPH_MODE" not in os.environ:
|
||||
os.environ["ENABLE_GRAPH_MODE"] = "false"
|
||||
if "ENABLE_PROFILING_DEBUG" not in os.environ:
|
||||
os.environ["ENABLE_PROFILING_DEBUG"] = "true"
|
||||
|
||||
config = combined_config
|
||||
|
||||
config.prompt = prompt
|
||||
config.negative_prompt = negative_prompt
|
||||
|
||||
# Check if task requires image
|
||||
if config.task == "i2v" and image is None:
|
||||
raise ValueError("i2v task requires input image")
|
||||
|
||||
temp_files = []
|
||||
|
||||
try:
|
||||
# Handle image input for i2v
|
||||
if config.task == "i2v" and image is not None:
|
||||
# Convert ComfyUI image to PIL
|
||||
image_np = (image[0].cpu().numpy() * 255).astype(np.uint8)
|
||||
pil_image = Image.fromarray(image_np)
|
||||
|
||||
# Save to temporary file
|
||||
with tempfile.NamedTemporaryFile(suffix=".png", delete=False) as tmp:
|
||||
pil_image.save(tmp.name)
|
||||
config.image_path = tmp.name
|
||||
temp_files.append(tmp.name)
|
||||
|
||||
# Check if we need to reinitialize runner
|
||||
config_hash = self._get_config_hash(configs)
|
||||
needs_reinit = (
|
||||
self._current_runner is None or self._current_config_hash != config_hash or configs.get("memory", {}).get("lazy_load", False)
|
||||
)
|
||||
if audio is not None and hasattr(config, "model_cls") and "audio" in config.model_cls:
|
||||
if isinstance(audio, tuple) and len(audio) == 2:
|
||||
waveform, sample_rate = audio
|
||||
|
||||
if isinstance(waveform, torch.Tensor):
|
||||
waveform = waveform.cpu().numpy()
|
||||
|
||||
with tempfile.NamedTemporaryFile(suffix=".wav", delete=False) as tmp:
|
||||
import scipy.io.wavfile as wavfile
|
||||
|
||||
if waveform.ndim == 1:
|
||||
wavfile.write(tmp.name, sample_rate, waveform)
|
||||
else:
|
||||
if waveform.shape[0] < waveform.shape[1]:
|
||||
waveform = waveform.T
|
||||
wavfile.write(tmp.name, sample_rate, waveform)
|
||||
|
||||
config.audio_path = tmp.name
|
||||
temp_files.append(tmp.name)
|
||||
|
||||
config_hash = self._get_config_hash(config)
|
||||
needs_reinit = self._current_runner is None or self._current_config_hash != config_hash or getattr(config, "lazy_load", False)
|
||||
|
||||
if needs_reinit:
|
||||
# Clear old runner
|
||||
if self._current_runner is not None:
|
||||
del self._current_runner
|
||||
torch.cuda.empty_cache()
|
||||
gc.collect()
|
||||
|
||||
# Initialize new runner
|
||||
self._current_runner = init_runner(config)
|
||||
self._current_config_hash = config_hash
|
||||
else:
|
||||
# Update config for existing runner
|
||||
self._current_runner.config = config
|
||||
|
||||
# Set up progress callback
|
||||
total_steps = config.get("infer_steps", 40)
|
||||
progress = ProgressBar(total_steps)
|
||||
|
||||
@@ -425,11 +466,9 @@ class LightX2VModularInference:
|
||||
|
||||
self._current_runner.set_progress_callback(update_progress)
|
||||
|
||||
# Run inference
|
||||
images = asyncio.run(self._current_runner.run_pipeline(save_video=False))
|
||||
|
||||
# Clean up if requested
|
||||
if configs.get("memory", {}).get("unload_after_inference", False):
|
||||
if getattr(config, "unload_after_inference", False):
|
||||
del self._current_runner
|
||||
self._current_runner = None
|
||||
self._current_config_hash = None
|
||||
@@ -437,7 +476,6 @@ class LightX2VModularInference:
|
||||
torch.cuda.empty_cache()
|
||||
gc.collect()
|
||||
|
||||
# Convert output to ComfyUI format
|
||||
images = (images + 1) / 2
|
||||
images = images.squeeze(0).permute(1, 2, 3, 0).cpu()
|
||||
images = torch.clamp(images, 0, 1)
|
||||
@@ -448,17 +486,15 @@ class LightX2VModularInference:
|
||||
logging.error(f"Error during inference: {e}")
|
||||
raise
|
||||
|
||||
finally:
|
||||
# Clean up temporary files
|
||||
for temp_file in temp_files:
|
||||
if os.path.exists(temp_file):
|
||||
try:
|
||||
os.unlink(temp_file)
|
||||
except Exception:
|
||||
pass
|
||||
# finally:
|
||||
# for temp_file in temp_files:
|
||||
# if os.path.exists(temp_file):
|
||||
# try:
|
||||
# os.unlink(temp_file)
|
||||
# except Exception:
|
||||
# pass
|
||||
|
||||
|
||||
# Node mappings
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"LightX2VInferenceConfig": LightX2VInferenceConfig,
|
||||
"LightX2VTeaCache": LightX2VTeaCache,
|
||||
@@ -466,6 +502,7 @@ NODE_CLASS_MAPPINGS = {
|
||||
"LightX2VMemoryOptimization": LightX2VMemoryOptimization,
|
||||
"LightX2VLightweightVAE": LightX2VLightweightVAE,
|
||||
"LightX2VLoRALoader": LightX2VLoRALoader,
|
||||
"LightX2VConfigCombiner": LightX2VConfigCombiner,
|
||||
"LightX2VModularInference": LightX2VModularInference,
|
||||
}
|
||||
|
||||
@@ -476,5 +513,6 @@ NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"LightX2VMemoryOptimization": "LightX2V Memory Optimization",
|
||||
"LightX2VLightweightVAE": "LightX2V Lightweight VAE",
|
||||
"LightX2VLoRALoader": "LightX2V LoRA Loader",
|
||||
"LightX2VConfigCombiner": "LightX2V Config Combiner",
|
||||
"LightX2VModularInference": "LightX2V Modular Inference",
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user