refactor: remove TempFileManager from ConfigBuilder; enhance TempFileManager with temp_dir and cleanup_dir methods for improved temporary directory management
This commit is contained in:
@@ -273,7 +273,6 @@ class ConfigBuilder:
|
||||
|
||||
def __init__(self):
|
||||
self.manager = ModularConfigManager()
|
||||
self.temp_manager = TempFileManager()
|
||||
|
||||
def combine_configs(
|
||||
self,
|
||||
|
||||
+35
-25
@@ -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()
|
||||
|
||||
|
||||
|
||||
+1
-1
Submodule lightx2v updated: 411dd37a0e...fd53d98563
@@ -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)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user