feat: implement HTTPFileDownloader for downloading files from URLs; integrate downloader into LightX2VModularInference and LightX2VConfigCombinerV2 for enhanced audio and mask handling
This commit is contained in:
+136
-2
@@ -1,8 +1,7 @@
|
||||
"""Unified file handlers for LightX2V ComfyUI wrapper."""
|
||||
|
||||
import logging
|
||||
import os
|
||||
import tempfile
|
||||
import urllib.parse
|
||||
from abc import ABC, abstractmethod
|
||||
from contextlib import contextmanager
|
||||
from typing import Any, Dict, List, Optional, Tuple, Union
|
||||
@@ -322,6 +321,141 @@ class TempFileManager:
|
||||
self.cleanup_all()
|
||||
|
||||
|
||||
class HTTPFileDownloader:
|
||||
"""Handler for downloading files from HTTP/HTTPS URLs."""
|
||||
|
||||
def __init__(self):
|
||||
self.temp_manager = TempFileManager()
|
||||
|
||||
@staticmethod
|
||||
def is_url(path: str) -> bool:
|
||||
"""Check if the path is an HTTP/HTTPS URL.
|
||||
|
||||
Args:
|
||||
path: Path to check
|
||||
|
||||
Returns:
|
||||
True if path is HTTP/HTTPS URL, False otherwise
|
||||
"""
|
||||
if not path:
|
||||
return False
|
||||
|
||||
parsed = urllib.parse.urlparse(path)
|
||||
return parsed.scheme in ("http", "https")
|
||||
|
||||
def download_to_input(self, url: str, filename: Optional[str] = None) -> str:
|
||||
"""Download file from URL to ComfyUI input directory.
|
||||
|
||||
Args:
|
||||
url: URL to download from
|
||||
filename: Target filename (optional, will be generated if not provided)
|
||||
|
||||
Returns:
|
||||
Absolute path to downloaded file
|
||||
|
||||
Raises:
|
||||
Exception: If download fails
|
||||
"""
|
||||
try:
|
||||
import requests
|
||||
except ImportError:
|
||||
logging.error("requests module not available for HTTP download")
|
||||
raise ImportError("requests module is required for HTTP download")
|
||||
|
||||
# Generate filename if not provided
|
||||
if not filename:
|
||||
# Extract filename from URL
|
||||
parsed_url = urllib.parse.urlparse(url)
|
||||
url_filename = os.path.basename(parsed_url.path)
|
||||
|
||||
if url_filename:
|
||||
# Use URL filename but add a unique suffix to avoid conflicts
|
||||
import uuid
|
||||
|
||||
name, ext = os.path.splitext(url_filename)
|
||||
filename = f"{name}_{uuid.uuid4().hex[:8]}{ext}"
|
||||
else:
|
||||
# Generate a completely new filename
|
||||
import uuid
|
||||
|
||||
filename = f"downloaded_{uuid.uuid4().hex[:8]}"
|
||||
|
||||
# Get input directory
|
||||
input_dir = ComfyUIFileResolver.get_input_directory()
|
||||
full_path = os.path.join(input_dir, filename)
|
||||
|
||||
# Create directory if needed
|
||||
os.makedirs(input_dir, exist_ok=True)
|
||||
|
||||
try:
|
||||
logging.info(f"Downloading file from {url} to {full_path}")
|
||||
|
||||
# Download with streaming to handle large files
|
||||
response = requests.get(url, stream=True, timeout=30)
|
||||
response.raise_for_status()
|
||||
|
||||
# Get total size for progress reporting
|
||||
total_size = int(response.headers.get("content-length", 0))
|
||||
downloaded_size = 0
|
||||
|
||||
# Write to file
|
||||
with open(full_path, "wb") as f:
|
||||
for chunk in response.iter_content(chunk_size=8192):
|
||||
if chunk:
|
||||
f.write(chunk)
|
||||
downloaded_size += len(chunk)
|
||||
|
||||
# Log progress for large files
|
||||
if total_size > 0 and total_size > 1024 * 1024: # > 1MB
|
||||
progress = (downloaded_size / total_size) * 100
|
||||
if downloaded_size % (1024 * 1024) == 0: # Log every 1MB
|
||||
logging.debug(f"Download progress: {progress:.1f}%")
|
||||
|
||||
logging.info(f"Successfully downloaded file to {full_path}")
|
||||
return full_path
|
||||
|
||||
except requests.exceptions.RequestException as e:
|
||||
# Clean up partial file if download failed
|
||||
if os.path.exists(full_path):
|
||||
try:
|
||||
os.unlink(full_path)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
logging.error(f"Failed to download file from {url}: {e}")
|
||||
raise Exception(f"Failed to download file from {url}: {e}")
|
||||
except Exception as e:
|
||||
# Clean up partial file if download failed
|
||||
if os.path.exists(full_path):
|
||||
try:
|
||||
os.unlink(full_path)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
logging.error(f"Error downloading file: {e}")
|
||||
raise
|
||||
|
||||
def download_if_url(self, path: str, prefix: str = "downloaded") -> str:
|
||||
"""Download file if path is URL, otherwise return path as-is.
|
||||
|
||||
Args:
|
||||
path: Path or URL to process
|
||||
prefix: Prefix for downloaded filename
|
||||
|
||||
Returns:
|
||||
Absolute path to local file
|
||||
"""
|
||||
if self.is_url(path):
|
||||
# Generate filename with prefix
|
||||
import uuid
|
||||
|
||||
ext = os.path.splitext(urllib.parse.urlparse(path).path)[1] or ".bin"
|
||||
filename = f"{prefix}_{uuid.uuid4().hex[:8]}{ext}"
|
||||
return self.download_to_input(path, filename)
|
||||
|
||||
return path
|
||||
|
||||
|
||||
class ComfyUIFileResolver:
|
||||
"""Resolve file paths for ComfyUI input/output directories."""
|
||||
|
||||
|
||||
@@ -1,5 +1,3 @@
|
||||
"""Modular ComfyUI nodes for LightX2V without presets."""
|
||||
|
||||
import gc
|
||||
import json
|
||||
import logging
|
||||
@@ -27,6 +25,7 @@ from .data_models import (
|
||||
from .file_handlers import (
|
||||
AudioFileHandler,
|
||||
ComfyUIFileResolver,
|
||||
HTTPFileDownloader,
|
||||
ImageFileHandler,
|
||||
TempFileManager,
|
||||
)
|
||||
@@ -695,10 +694,6 @@ class LightX2VConfigCombiner:
|
||||
{"tooltip": "Memory optimization configuration"},
|
||||
),
|
||||
"lora_chain": ("LORA_CHAIN", {"tooltip": "LoRA chain configuration"}),
|
||||
"talk_objects_config": (
|
||||
"TALK_OBJECTS_CONFIG",
|
||||
{"tooltip": "Multi-person talk objects configuration"},
|
||||
),
|
||||
},
|
||||
}
|
||||
|
||||
@@ -752,6 +747,7 @@ class LightX2VModularInference:
|
||||
self.image_handler = ImageFileHandler()
|
||||
self.audio_handler = AudioFileHandler()
|
||||
self.resolver = ComfyUIFileResolver()
|
||||
self.http_downloader = HTTPFileDownloader()
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
@@ -840,21 +836,47 @@ class LightX2VModularInference:
|
||||
if "audio" in processed_obj:
|
||||
processed_talk_objects.append(processed_obj)
|
||||
|
||||
# Resolve paths
|
||||
# Resolve paths and download URLs
|
||||
for obj in processed_talk_objects:
|
||||
if "audio" in obj and obj["audio"]:
|
||||
audio_path = obj["audio"]
|
||||
if not os.path.isabs(audio_path) and not audio_path.startswith("/tmp"):
|
||||
|
||||
# Check if it's a URL and download if needed
|
||||
if self.http_downloader.is_url(audio_path):
|
||||
try:
|
||||
downloaded_path = self.http_downloader.download_if_url(audio_path, prefix="audio")
|
||||
obj["audio"] = downloaded_path
|
||||
logging.info(f"Downloaded audio from URL: {audio_path} -> {downloaded_path}")
|
||||
except Exception as e:
|
||||
logging.error(f"Failed to download audio from {audio_path}: {e}")
|
||||
continue
|
||||
# Handle relative paths
|
||||
elif not os.path.isabs(audio_path) and not audio_path.startswith("/tmp"):
|
||||
obj["audio"] = self.resolver.resolve_input_path(audio_path)
|
||||
logging.info(f"Resolved audio path: {audio_path} -> {obj['audio']}")
|
||||
|
||||
# Check if file exists
|
||||
if not os.path.exists(obj["audio"]):
|
||||
logging.warning(f"Audio file not found: {obj['audio']}")
|
||||
|
||||
if "mask" in obj and obj["mask"]:
|
||||
mask_path = obj["mask"]
|
||||
if not os.path.isabs(mask_path) and not mask_path.startswith("/tmp"):
|
||||
|
||||
# Check if it's a URL and download if needed
|
||||
if self.http_downloader.is_url(mask_path):
|
||||
try:
|
||||
downloaded_path = self.http_downloader.download_if_url(mask_path, prefix="mask")
|
||||
obj["mask"] = downloaded_path
|
||||
logging.info(f"Downloaded mask from URL: {mask_path} -> {downloaded_path}")
|
||||
except Exception as e:
|
||||
logging.error(f"Failed to download mask from {mask_path}: {e}")
|
||||
# Don't skip the object if mask download fails (mask is optional)
|
||||
# Handle relative paths
|
||||
elif not os.path.isabs(mask_path) and not mask_path.startswith("/tmp"):
|
||||
obj["mask"] = self.resolver.resolve_input_path(mask_path)
|
||||
logging.info(f"Resolved mask path: {mask_path} -> {obj['mask']}")
|
||||
|
||||
# Check if file exists
|
||||
if not os.path.exists(obj["mask"]):
|
||||
logging.warning(f"Mask file not found: {obj['mask']}")
|
||||
|
||||
@@ -921,6 +943,286 @@ class LightX2VModularInference:
|
||||
pass
|
||||
|
||||
|
||||
class LightX2VConfigCombinerV2:
|
||||
"""Config combiner that also handles data preparation (image/audio/prompts)."""
|
||||
|
||||
def __init__(self):
|
||||
self.config_builder = ConfigBuilder()
|
||||
self.temp_manager = TempFileManager()
|
||||
self.image_handler = ImageFileHandler()
|
||||
self.audio_handler = AudioFileHandler()
|
||||
self.resolver = ComfyUIFileResolver()
|
||||
self.http_downloader = HTTPFileDownloader()
|
||||
|
||||
@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": {
|
||||
"teacache_config": (
|
||||
"TEACACHE_CONFIG",
|
||||
{"tooltip": "TeaCache configuration"},
|
||||
),
|
||||
"quantization_config": (
|
||||
"QUANT_CONFIG",
|
||||
{"tooltip": "Quantization configuration"},
|
||||
),
|
||||
"memory_config": (
|
||||
"MEMORY_CONFIG",
|
||||
{"tooltip": "Memory optimization configuration"},
|
||||
),
|
||||
"lora_chain": ("LORA_CHAIN", {"tooltip": "LoRA chain configuration"}),
|
||||
"talk_objects_config": ("TALK_OBJECTS", {"tooltip": "Talk objects configuration"}),
|
||||
"image": ("IMAGE", {"tooltip": "Input image for i2v task"}),
|
||||
"audio": (
|
||||
"AUDIO",
|
||||
{"tooltip": "Input audio for audio-driven generation"},
|
||||
),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("PREPARED_CONFIG",)
|
||||
RETURN_NAMES = ("prepared_config",)
|
||||
FUNCTION = "prepare_config"
|
||||
CATEGORY = "LightX2V/ConfigV2"
|
||||
|
||||
def prepare_config(
|
||||
self,
|
||||
inference_config,
|
||||
prompt,
|
||||
negative_prompt,
|
||||
teacache_config=None,
|
||||
quantization_config=None,
|
||||
memory_config=None,
|
||||
lora_chain=None,
|
||||
talk_objects_config=None,
|
||||
image=None,
|
||||
audio=None,
|
||||
):
|
||||
"""Combine configurations and prepare data for inference."""
|
||||
|
||||
# Convert dict configs back to objects if needed
|
||||
inf_config = InferenceConfig(**inference_config) if isinstance(inference_config, dict) else inference_config
|
||||
tea_config = TeaCacheConfig(**teacache_config) if teacache_config and isinstance(teacache_config, dict) else teacache_config
|
||||
quant_config = (
|
||||
QuantizationConfig(**quantization_config) if quantization_config and isinstance(quantization_config, dict) else quantization_config
|
||||
)
|
||||
mem_config = MemoryOptimizationConfig(**memory_config) if memory_config and isinstance(memory_config, dict) else memory_config
|
||||
|
||||
# Build combined config
|
||||
config = self.config_builder.combine_configs(
|
||||
inference_config=inf_config,
|
||||
teacache_config=tea_config,
|
||||
quantization_config=quant_config,
|
||||
memory_config=mem_config,
|
||||
lora_chain=lora_chain,
|
||||
talk_objects_config=talk_objects_config,
|
||||
)
|
||||
|
||||
# Add prompts to config
|
||||
config.prompt = prompt
|
||||
config.negative_prompt = negative_prompt
|
||||
|
||||
# Validate task requirements
|
||||
if config.task == "i2v" and image is None:
|
||||
raise ValueError("i2v task requires input image")
|
||||
|
||||
# Handle image input
|
||||
if config.task == "i2v" and image is not None:
|
||||
image_np = (image[0].cpu().numpy() * 255).astype(np.uint8)
|
||||
pil_image = Image.fromarray(image_np)
|
||||
|
||||
temp_path = self.temp_manager.create_temp_file(suffix=".png")
|
||||
pil_image.save(temp_path)
|
||||
config.image_path = temp_path
|
||||
logging.info(f"Image saved to {temp_path}")
|
||||
|
||||
# Handle audio input for seko models
|
||||
if audio is not None and hasattr(config, "model_cls") and "seko" in config.model_cls:
|
||||
temp_path = self.temp_manager.create_temp_file(suffix=".wav")
|
||||
self.audio_handler.save(audio, temp_path)
|
||||
config.audio_path = temp_path
|
||||
logging.info(f"Audio saved to {temp_path}")
|
||||
|
||||
# Handle talk objects
|
||||
if hasattr(config, "talk_objects") and config.talk_objects:
|
||||
talk_objects = config.talk_objects
|
||||
processed_talk_objects = []
|
||||
|
||||
for talk_obj in talk_objects:
|
||||
processed_obj = {}
|
||||
|
||||
if "audio" in talk_obj:
|
||||
processed_obj["audio"] = talk_obj["audio"]
|
||||
|
||||
if "mask" in talk_obj:
|
||||
processed_obj["mask"] = talk_obj["mask"]
|
||||
|
||||
if "audio" in processed_obj:
|
||||
processed_talk_objects.append(processed_obj)
|
||||
|
||||
# Resolve paths and download URLs
|
||||
for obj in processed_talk_objects:
|
||||
if "audio" in obj and obj["audio"]:
|
||||
audio_path = obj["audio"]
|
||||
|
||||
# Check if it's a URL and download if needed
|
||||
if self.http_downloader.is_url(audio_path):
|
||||
try:
|
||||
downloaded_path = self.http_downloader.download_if_url(audio_path, prefix="audio")
|
||||
obj["audio"] = downloaded_path
|
||||
logging.info(f"Downloaded audio from URL: {audio_path} -> {downloaded_path}")
|
||||
except Exception as e:
|
||||
logging.error(f"Failed to download audio from {audio_path}: {e}")
|
||||
continue
|
||||
# Handle relative paths
|
||||
elif not os.path.isabs(audio_path) and not audio_path.startswith("/tmp"):
|
||||
obj["audio"] = self.resolver.resolve_input_path(audio_path)
|
||||
logging.info(f"Resolved audio path: {audio_path} -> {obj['audio']}")
|
||||
|
||||
# Check if file exists
|
||||
if not os.path.exists(obj["audio"]):
|
||||
logging.warning(f"Audio file not found: {obj['audio']}")
|
||||
|
||||
if "mask" in obj and obj["mask"]:
|
||||
mask_path = obj["mask"]
|
||||
|
||||
# Check if it's a URL and download if needed
|
||||
if self.http_downloader.is_url(mask_path):
|
||||
try:
|
||||
downloaded_path = self.http_downloader.download_if_url(mask_path, prefix="mask")
|
||||
obj["mask"] = downloaded_path
|
||||
logging.info(f"Downloaded mask from URL: {mask_path} -> {downloaded_path}")
|
||||
except Exception as e:
|
||||
logging.error(f"Failed to download mask from {mask_path}: {e}")
|
||||
# Don't skip the object if mask download fails (mask is optional)
|
||||
# Handle relative paths
|
||||
elif not os.path.isabs(mask_path) and not mask_path.startswith("/tmp"):
|
||||
obj["mask"] = self.resolver.resolve_input_path(mask_path)
|
||||
logging.info(f"Resolved mask path: {mask_path} -> {obj['mask']}")
|
||||
|
||||
# Check if file exists
|
||||
if not os.path.exists(obj["mask"]):
|
||||
logging.warning(f"Mask file not found: {obj['mask']}")
|
||||
|
||||
if processed_talk_objects:
|
||||
config.talk_objects = processed_talk_objects
|
||||
logging.info(f"Processed {len(processed_talk_objects)} talk objects")
|
||||
|
||||
logging.info("lightx2v prepared config: " + json.dumps(config, indent=2, ensure_ascii=False))
|
||||
|
||||
return (config,)
|
||||
|
||||
|
||||
class LightX2VModularInferenceV2:
|
||||
"""Pure inference node that takes prepared config and runs inference."""
|
||||
|
||||
_current_runner = None
|
||||
_current_config_hash = None
|
||||
|
||||
def __init__(self):
|
||||
if not hasattr(self.__class__, "_current_runner"):
|
||||
self.__class__._current_runner = None
|
||||
if not hasattr(self.__class__, "_current_config_hash"):
|
||||
self.__class__._current_config_hash = None
|
||||
|
||||
self.config_builder = ConfigBuilder()
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"prepared_config": (
|
||||
"PREPARED_CONFIG",
|
||||
{"tooltip": "Fully prepared configuration from ConfigCombinerV2"},
|
||||
),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("IMAGE", "AUDIO")
|
||||
RETURN_NAMES = ("images", "audio")
|
||||
FUNCTION = "generate"
|
||||
CATEGORY = "LightX2V/InferenceV2"
|
||||
|
||||
def _get_config_hash(self, config) -> str:
|
||||
"""Get hash of configuration to detect changes."""
|
||||
return self.config_builder.get_config_hash(config)
|
||||
|
||||
def generate(self, prepared_config):
|
||||
"""Run inference with prepared configuration."""
|
||||
config = prepared_config
|
||||
|
||||
try:
|
||||
config_hash = self._get_config_hash(config)
|
||||
|
||||
current_runner = getattr(self.__class__, "_current_runner", None)
|
||||
current_config_hash = getattr(self.__class__, "_current_config_hash", None)
|
||||
|
||||
needs_reinit = current_runner is None or current_config_hash != config_hash or getattr(config, "lazy_load", False)
|
||||
|
||||
logging.info(f"Needs reinit: {needs_reinit}, old config hash: {current_config_hash}, new config hash: {config_hash}")
|
||||
if needs_reinit:
|
||||
if current_runner is not None:
|
||||
# current_runner.end_run()
|
||||
del self.__class__._current_runner
|
||||
torch.cuda.empty_cache()
|
||||
gc.collect()
|
||||
|
||||
self.__class__._current_runner = init_runner(config)
|
||||
self.__class__._current_config_hash = config_hash
|
||||
else:
|
||||
if hasattr(current_runner, "config"):
|
||||
current_runner.config = config
|
||||
current_runner.model.config = config
|
||||
current_runner.model.scheduler.config = config
|
||||
|
||||
progress = ProgressBar(100)
|
||||
|
||||
def update_progress(current_step, _total):
|
||||
progress.update_absolute(current_step)
|
||||
|
||||
current_runner = getattr(self.__class__, "_current_runner", None)
|
||||
|
||||
if hasattr(current_runner, "set_progress_callback"):
|
||||
current_runner.set_progress_callback(update_progress)
|
||||
|
||||
result_dict = current_runner.run_pipeline()
|
||||
images = result_dict.get("video", None)
|
||||
audio = result_dict.get("audio", None)
|
||||
|
||||
if getattr(config, "unload_after_inference", False):
|
||||
if hasattr(self.__class__, "_current_runner"):
|
||||
del self.__class__._current_runner
|
||||
self.__class__._current_runner = None
|
||||
self.__class__._current_config_hash = None
|
||||
|
||||
torch.cuda.empty_cache()
|
||||
gc.collect()
|
||||
|
||||
return (images, audio)
|
||||
|
||||
except Exception as e:
|
||||
logging.error(f"Error during inference: {e}")
|
||||
raise
|
||||
|
||||
finally:
|
||||
# Cleanup is handled by TempFileManager destructor
|
||||
pass
|
||||
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"LightX2VInferenceConfig": LightX2VInferenceConfig,
|
||||
"LightX2VTeaCache": LightX2VTeaCache,
|
||||
@@ -929,6 +1231,8 @@ NODE_CLASS_MAPPINGS = {
|
||||
"LightX2VLoRALoader": LightX2VLoRALoader,
|
||||
"LightX2VConfigCombiner": LightX2VConfigCombiner,
|
||||
"LightX2VModularInference": LightX2VModularInference,
|
||||
"LightX2VConfigCombinerV2": LightX2VConfigCombinerV2,
|
||||
"LightX2VModularInferenceV2": LightX2VModularInferenceV2,
|
||||
"LightX2VTalkObjectInput": TalkObjectInput,
|
||||
"LightX2VTalkObjectsCombiner": TalkObjectsCombiner,
|
||||
"LightX2VTalkObjectsFromJSON": TalkObjectsFromJSON,
|
||||
@@ -943,6 +1247,8 @@ NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"LightX2VLoRALoader": "LightX2V LoRA Loader",
|
||||
"LightX2VConfigCombiner": "LightX2V Config Combiner",
|
||||
"LightX2VModularInference": "LightX2V Modular Inference",
|
||||
"LightX2VConfigCombinerV2": "LightX2V Config Combiner V2",
|
||||
"LightX2VModularInferenceV2": "LightX2V Modular Inference V2",
|
||||
"LightX2VTalkObjectInput": "LightX2V Talk Object Input (Single)",
|
||||
"LightX2VTalkObjectsCombiner": "LightX2V Talk Objects Combiner",
|
||||
"LightX2VTalkObjectsFromFiles": "LightX2V Talk Objects From Files",
|
||||
|
||||
Reference in New Issue
Block a user