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:
gaclove
2025-09-26 16:16:27 +08:00
parent 436a83340d
commit 6fff546770
2 changed files with 451 additions and 11 deletions
+136 -2
View File
@@ -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."""
+315 -9
View File
@@ -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",