diff --git a/config_builder.py b/config_builder.py index d9706ea..c9ec0d3 100644 --- a/config_builder.py +++ b/config_builder.py @@ -273,7 +273,6 @@ class ConfigBuilder: def __init__(self): self.manager = ModularConfigManager() - self.temp_manager = TempFileManager() def combine_configs( self, diff --git a/file_handlers.py b/file_handlers.py index 24c02ad..54a797b 100644 --- a/file_handlers.py +++ b/file_handlers.py @@ -252,23 +252,12 @@ class MaskFileHandler(ImageFileHandler): class TempFileManager: - """Manager for temporary files with automatic cleanup.""" - def __init__(self): self.temp_files: List[str] = [] + self.temp_dirs: List[str] = [] @contextmanager def temp_file(self, suffix: str = "", prefix: str = "lightx2v_", delete: bool = True): - """Context manager for temporary file creation. - - Args: - suffix: File suffix - prefix: File prefix - delete: Whether to delete on exit - - Yields: - Path to temporary file - """ temp_file = tempfile.NamedTemporaryFile(suffix=suffix, prefix=prefix, delete=False) temp_path = temp_file.name temp_file.close() @@ -282,15 +271,6 @@ class TempFileManager: self.cleanup_file(temp_path) def create_temp_file(self, suffix: str = "", prefix: str = "lightx2v_") -> str: - """Create a temporary file that will be tracked for cleanup. - - Args: - suffix: File suffix - prefix: File prefix - - Returns: - Path to temporary file - """ with tempfile.NamedTemporaryFile(suffix=suffix, prefix=prefix, delete=False) as tmp: temp_path = tmp.name @@ -298,7 +278,6 @@ class TempFileManager: return temp_path def cleanup_file(self, path: str): - """Clean up a specific file.""" if path in self.temp_files: self.temp_files.remove(path) @@ -309,15 +288,46 @@ class TempFileManager: except Exception as e: logging.warning(f"Failed to clean up {path}: {e}") + @contextmanager + def temp_dir(self, suffix: str = "", prefix: str = "lightx2v_", delete: bool = True): + temp_dir = tempfile.mkdtemp(suffix=suffix, prefix=prefix) + self.temp_dirs.append(temp_dir) + + try: + yield temp_dir + finally: + if delete: + self.cleanup_dir(temp_dir) + + + def create_temp_dir(self, suffix: str = "", prefix: str = "lightx2v_") -> str: + temp_dir = tempfile.mkdtemp(suffix=suffix, prefix=prefix) + self.temp_dirs.append(temp_dir) + return temp_dir + + + def cleanup_dir(self, path: str): + if path in self.temp_dirs: + self.temp_dirs.remove(path) + + if os.path.exists(path): + try: + import shutil + shutil.rmtree(path) + logging.debug(f"Cleaned up temp directory: {path}") + except Exception as e: + logging.warning(f"Failed to clean up directory {path}: {e}") + def cleanup_all(self): - """Clean up all tracked temporary files.""" for temp_file in self.temp_files[:]: self.cleanup_file(temp_file) - self.temp_files.clear() + for temp_dir in self.temp_dirs[:]: + self.cleanup_dir(temp_dir) + self.temp_dirs.clear() + def __del__(self): - """Cleanup on deletion.""" self.cleanup_all() diff --git a/lightx2v b/lightx2v index 411dd37..fd53d98 160000 --- a/lightx2v +++ b/lightx2v @@ -1 +1 @@ -Subproject commit 411dd37a0e728d95c74d0a7457149a3ec5ac63d2 +Subproject commit fd53d98563b2e647bdeeea0cb82aa31ea2ae9e8f diff --git a/nodes.py b/nodes.py index bc708ae..bcdbdc4 100644 --- a/nodes.py +++ b/nodes.py @@ -30,6 +30,8 @@ from .file_handlers import ( TempFileManager, ) from .lightx2v.lightx2v.infer import init_runner +from .lightx2v.lightx2v.utils.input_info import set_input_info +from .lightx2v.lightx2v.utils.set_config import set_config from .model_utils import scan_loras, scan_models, support_model_cls_list @@ -62,10 +64,10 @@ class LightX2VInferenceConfig: }, ), "task": ( - ["t2v", "i2v"], + ["t2v", "i2v", "s2v"], { "default": "i2v", - "tooltip": "Task type: text-to-video or image-to-video", + "tooltip": "Task type: text-to-video or image-to-video or audio-to-video", }, ), "infer_steps": ( @@ -767,10 +769,10 @@ class LightX2VModularInference: ), }, "optional": { - "image": ("IMAGE", {"tooltip": "Input image for i2v task"}), + "image": ("IMAGE", {"tooltip": "Input image for i2v or s2v task"}), "audio": ( "AUDIO", - {"tooltip": "Input audio for audio-driven generation"}, + {"tooltip": "Input audio for audio-driven generation for s2v task"}, ), }, } @@ -798,12 +800,12 @@ class LightX2VModularInference: config.prompt = prompt config.negative_prompt = negative_prompt - if config.task == "i2v" and image is None: - raise ValueError("i2v task requires input image") + if config.task in ["i2v", "s2v"] and image is None: + raise ValueError(f"{config.task} task requires input image") try: # Handle image input - if config.task == "i2v" and image is not None: + if config.task in ["i2v", "s2v"] and image is not None: image_np = (image[0].cpu().numpy() * 255).astype(np.uint8) pil_image = Image.fromarray(image_np) @@ -900,14 +902,9 @@ class LightX2VModularInference: del self.__class__._current_runner torch.cuda.empty_cache() gc.collect() - - self.__class__._current_runner = init_runner(config) + formated_config = set_config(config) + self.__class__._current_runner = init_runner(formated_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) @@ -919,7 +916,12 @@ class LightX2VModularInference: if hasattr(current_runner, "set_progress_callback"): current_runner.set_progress_callback(update_progress) - result_dict = current_runner.run_pipeline(save_video=False) + config["return_result_tensor"] = True + config["save_result_path"] = "" + config["negative_prompt"] = config.get("negative_prompt", "") + input_info = set_input_info(config) + current_runner.set_config(config) + result_dict = current_runner.run_pipeline(input_info) images = result_dict.get("video", None) audio = result_dict.get("audio", None) @@ -991,10 +993,10 @@ class LightX2VConfigCombinerV2: ), "lora_chain": ("LORA_CHAIN", {"tooltip": "LoRA chain configuration"}), "talk_objects_config": ("TALK_OBJECTS_CONFIG", {"tooltip": "Talk objects configuration"}), - "image": ("IMAGE", {"tooltip": "Input image for i2v task"}), + "image": ("IMAGE", {"tooltip": "Input image for i2v or s2v task"}), "audio": ( "AUDIO", - {"tooltip": "Input audio for audio-driven generation"}, + {"tooltip": "Input audio for audio-driven generation for s2v task"}, ), }, } @@ -1042,11 +1044,11 @@ class LightX2VConfigCombinerV2: config.negative_prompt = negative_prompt # Validate task requirements - if config.task == "i2v" and image is None: - raise ValueError("i2v task requires input image") + if config.task in ["i2v", "s2v"] and image is None: + raise ValueError("i2v or s2v task requires input image") # Handle image input - if config.task == "i2v" and image is not None: + if config.task in ["i2v", "s2v"] and image is not None: image_np = (image[0].cpu().numpy() * 255).astype(np.uint8) pil_image = Image.fromarray(image_np) @@ -1124,8 +1126,13 @@ class LightX2VConfigCombinerV2: logging.warning(f"Mask file not found: {obj['mask']}") if processed_talk_objects: - config.talk_objects = processed_talk_objects + # config.talk_objects = processed_talk_objects + temp_dir = self.temp_manager.create_temp_dir() + with open(os.path.join(temp_dir, "config.json"), "w") as f: + json.dump({"talk_objects": processed_talk_objects}, f) + config.audio_path = temp_dir logging.info(f"Processed {len(processed_talk_objects)} talk objects") + logging.info("lightx2v prepared config: " + json.dumps(config, indent=2, ensure_ascii=False)) @@ -1185,26 +1192,28 @@ class LightX2VModularInferenceV2: del self.__class__._current_runner torch.cuda.empty_cache() gc.collect() - - self.__class__._current_runner = init_runner(config) + formatted_config =set_config(config) + self.__class__._current_runner = init_runner(formatted_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) + current_runner = self.__class__._current_runner if hasattr(current_runner, "set_progress_callback"): current_runner.set_progress_callback(update_progress) - result_dict = current_runner.run_pipeline(save_video=False) + config["return_result_tensor"] = True + config["save_result_path"] = "" + config["negative_prompt"] = config.get("negative_prompt", "") + input_info = set_input_info(config) + current_runner.set_config(config) + + result_dict = current_runner.run_pipeline(input_info) + images = result_dict.get("video", None) audio = result_dict.get("audio", None)