refactor: remove TempFileManager from ConfigBuilder; enhance TempFileManager with temp_dir and cleanup_dir methods for improved temporary directory management

This commit is contained in:
gaclove
2025-10-14 07:29:35 +00:00
parent e68b02da93
commit e234a2b167
4 changed files with 75 additions and 57 deletions
-1
View File
@@ -273,7 +273,6 @@ class ConfigBuilder:
def __init__(self):
self.manager = ModularConfigManager()
self.temp_manager = TempFileManager()
def combine_configs(
self,
+35 -25
View File
@@ -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()
+39 -30
View File
@@ -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)