Refactor code structure and update class names in nodes
- Added structured logging - Removed SingleGPUSaverNode and MultiGPUSaverNode classes - Added VideoSaverNode class to handle video saving - Updated imports and mappings in __init__.py #refactor #codestructure #classnames #nodes
This commit is contained in:
+4
-6
@@ -4,17 +4,16 @@ from .nodes.single_gpu_loader import SingleGPUVEnhancerLoader
|
||||
from .nodes.multi_gpu_loader import MultiGPUVEnhancerLoader
|
||||
from .nodes.single_gpu_inference import SingleGPUInferenceNode
|
||||
from .nodes.multigpu_inference import MultiGPUInferenceNode
|
||||
from .nodes.single_gpu_saver import SingleGPUSaverNode
|
||||
from .nodes.multigpu_saver import MultiGPUSaverNode
|
||||
from .nodes.video_saver import VideoSaverNode
|
||||
from .nodes.video_loader import VideoLoaderNode
|
||||
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"SingleGPUVEnhancerLoader": SingleGPUVEnhancerLoader,
|
||||
"MultiGPUVEnhancerLoader": MultiGPUVEnhancerLoader,
|
||||
"SingleGPUInference": SingleGPUInferenceNode,
|
||||
"MultiGPUInference": MultiGPUInferenceNode,
|
||||
"SingleGPUSaver": SingleGPUSaverNode,
|
||||
"MultiGPUSaver": MultiGPUSaverNode,
|
||||
"VideoSaver": VideoSaverNode,
|
||||
"VideoLoader": VideoLoaderNode,
|
||||
}
|
||||
|
||||
@@ -23,8 +22,7 @@ NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"MultiGPUVEnhancerLoader": "Load VEnhancer (Multi-GPU)",
|
||||
"SingleGPUInference": "Enhance Video (Single GPU)",
|
||||
"MultiGPUInference": "Enhance Video (Multi-GPU)",
|
||||
"SingleGPUSaver": "Save Video (Single GPU)",
|
||||
"MultiGPUSaver": "Save Video (Multi-GPU)",
|
||||
"VideoSaver": "Save Video",
|
||||
"VideoLoader": "Load Video",
|
||||
}
|
||||
|
||||
|
||||
@@ -1,11 +1,20 @@
|
||||
"""Node for loading VEnhancer model in multi-GPU distributed mode."""
|
||||
|
||||
from typing import Dict, Any, Tuple
|
||||
import time
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
from loguru import logger
|
||||
from VEnhancer.configs.distributred_venhancer_config import DistributedConfig
|
||||
from VEnhancer.enhance_a_video_MultiGPU import DistributedVEnhancer
|
||||
|
||||
|
||||
class MultiGPUVEnhancerLoader:
|
||||
"""ComfyUI node for loading VEnhancer model in distributed multi-GPU mode."""
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls) -> Dict[str, Any]:
|
||||
"""Define input parameters for distributed model loading."""
|
||||
return {
|
||||
"required": {
|
||||
"version": (["v1", "v2"], {"default": "v1"}),
|
||||
@@ -27,13 +36,68 @@ class MultiGPUVEnhancerLoader:
|
||||
FUNCTION = "load_model"
|
||||
CATEGORY = "loaders/venhancer"
|
||||
|
||||
def load_model(self, world_size: int, rank: int, local_rank: int, **kwargs) -> Tuple[Any]:
|
||||
|
||||
dist_config = DistributedConfig(
|
||||
world_size=world_size,
|
||||
rank=rank,
|
||||
local_rank=local_rank
|
||||
)
|
||||
model = DistributedVEnhancer(dist_config)
|
||||
return (model,)
|
||||
def __init__(self):
|
||||
"""Initialize MultiGPUVEnhancerLoader with logging."""
|
||||
self.logger = logger.bind(node="MultiGPULoader")
|
||||
|
||||
def load_model(self, world_size: int, rank: int, local_rank: int, **kwargs) -> Tuple[Any]:
|
||||
"""
|
||||
Load VEnhancer model in distributed mode.
|
||||
|
||||
Args:
|
||||
world_size: Total number of GPUs
|
||||
rank: Global rank of current process
|
||||
local_rank: Local GPU ID
|
||||
**kwargs: Additional configuration parameters
|
||||
|
||||
Returns:
|
||||
Tuple[Any]: Tuple containing initialized distributed VEnhancer model
|
||||
|
||||
Raises:
|
||||
RuntimeError: If CUDA is not available or distributed setup fails
|
||||
Exception: If model loading fails
|
||||
"""
|
||||
try:
|
||||
start_time = time.time()
|
||||
init_mem = torch.cuda.memory_allocated(local_rank)
|
||||
|
||||
self.logger.info(f"Initializing distributed setup on rank {rank}/{world_size-1}",
|
||||
extra={
|
||||
"config": {
|
||||
"world_size": world_size,
|
||||
"rank": rank,
|
||||
"local_rank": local_rank,
|
||||
"gpu_name": torch.cuda.get_device_name(local_rank)
|
||||
}
|
||||
})
|
||||
|
||||
dist_config = DistributedConfig(
|
||||
world_size=world_size,
|
||||
rank=rank,
|
||||
local_rank=local_rank
|
||||
)
|
||||
model = DistributedVEnhancer(dist_config)
|
||||
|
||||
load_time = time.time() - start_time
|
||||
mem_used = torch.cuda.memory_allocated(local_rank) - init_mem
|
||||
|
||||
self.logger.success(f"Model loaded on rank {rank}", extra={
|
||||
"metrics": {
|
||||
"load_time": f"{load_time:.2f}s",
|
||||
"gpu_memory": f"{mem_used/1024/1024:.1f}MB",
|
||||
"gpu_utilization": f"{torch.cuda.utilization(local_rank)}%",
|
||||
"version": kwargs.get("version"),
|
||||
"solver_mode": kwargs.get("solver_mode")
|
||||
}
|
||||
})
|
||||
|
||||
dist.barrier() # Synchronize all processes
|
||||
return (model,)
|
||||
|
||||
except Exception as e:
|
||||
self.logger.exception(f"Failed to load model on rank {rank}: {str(e)}",
|
||||
extra={
|
||||
"gpu_state": torch.cuda.memory_summary(local_rank)
|
||||
})
|
||||
dist.destroy_process_group()
|
||||
raise
|
||||
@@ -1,9 +1,18 @@
|
||||
"""Node for running VEnhancer inference in distributed multi-GPU mode."""
|
||||
|
||||
from typing import Dict, Any, Tuple
|
||||
import time
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
from loguru import logger
|
||||
|
||||
|
||||
class MultiGPUInferenceNode:
|
||||
"""ComfyUI node for running VEnhancer inference across multiple GPUs."""
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls) -> Dict[str, Any]:
|
||||
"""Define input parameters for distributed video enhancement."""
|
||||
return {
|
||||
"required": {
|
||||
"model": ("MULTI_GPU_MODEL",),
|
||||
@@ -20,8 +29,89 @@ class MultiGPUInferenceNode:
|
||||
FUNCTION = "enhance_video"
|
||||
CATEGORY = "generators/venhancer"
|
||||
|
||||
def __init__(self):
|
||||
"""Initialize MultiGPUInferenceNode with logging."""
|
||||
self.logger = logger.bind(node="MultiGPUInference")
|
||||
|
||||
def enhance_video(self, model: Any, video: str, sync_mode: str, **kwargs) -> Tuple[str]:
|
||||
output_path = model.enhance_a_video(video_path=video, **kwargs)
|
||||
if sync_mode == "barrier":
|
||||
dist.barrier() # Synchronize all processes
|
||||
return (output_path,)
|
||||
"""
|
||||
Enhance video using distributed VEnhancer model.
|
||||
|
||||
Args:
|
||||
model: Distributed VEnhancer model instance
|
||||
video: Path to input video
|
||||
sync_mode: Synchronization mode ('barrier' or 'gather')
|
||||
**kwargs: Enhancement parameters including:
|
||||
- prompt: Text prompt for enhancement
|
||||
- up_scale: Upscaling factor
|
||||
- target_fps: Target frame rate
|
||||
- noise_aug: Noise augmentation level
|
||||
|
||||
Returns:
|
||||
Tuple[str]: Path to enhanced video
|
||||
|
||||
Raises:
|
||||
RuntimeError: If GPU memory is insufficient or synchronization fails
|
||||
Exception: If enhancement fails
|
||||
"""
|
||||
try:
|
||||
start_time = time.time()
|
||||
rank = dist.get_rank()
|
||||
world_size = dist.get_world_size()
|
||||
local_rank = model.dist_config.local_rank
|
||||
init_mem = torch.cuda.memory_allocated(local_rank)
|
||||
|
||||
self.logger.info(f"Starting distributed enhancement on rank {rank}/{world_size-1}",
|
||||
extra={
|
||||
"config": {
|
||||
"video": video,
|
||||
"prompt": kwargs.get("prompt"),
|
||||
"up_scale": kwargs.get("up_scale"),
|
||||
"target_fps": kwargs.get("target_fps"),
|
||||
"noise_aug": kwargs.get("noise_aug"),
|
||||
"sync_mode": sync_mode,
|
||||
"gpu": torch.cuda.get_device_name(local_rank)
|
||||
}
|
||||
})
|
||||
|
||||
# Run distributed enhancement
|
||||
output_path = model.enhance_a_video(video_path=video, **kwargs)
|
||||
|
||||
# Synchronize based on mode
|
||||
if sync_mode == "barrier":
|
||||
dist.barrier()
|
||||
sync_status = "barrier_sync"
|
||||
else:
|
||||
# Gather results if needed
|
||||
dist.gather(torch.tensor([1], device=f"cuda:{local_rank}"),
|
||||
dst=0 if rank == 0 else None)
|
||||
sync_status = "gather_sync"
|
||||
|
||||
# Log metrics
|
||||
enhance_time = time.time() - start_time
|
||||
peak_mem = torch.cuda.max_memory_allocated(local_rank)
|
||||
mem_used = peak_mem - init_mem
|
||||
|
||||
self.logger.success(f"Enhancement completed on rank {rank}", extra={
|
||||
"metrics": {
|
||||
"total_time": f"{enhance_time:.2f}s",
|
||||
"gpu_memory_used": f"{mem_used/1024/1024:.1f}MB",
|
||||
"peak_gpu_memory": f"{peak_mem/1024/1024:.1f}MB",
|
||||
"gpu_utilization": f"{torch.cuda.utilization(local_rank)}%",
|
||||
"sync_mode": sync_status
|
||||
}
|
||||
})
|
||||
|
||||
return (output_path,)
|
||||
|
||||
except Exception as e:
|
||||
self.logger.exception(f"Enhancement failed on rank {rank}", extra={
|
||||
"error": str(e),
|
||||
"gpu_state": torch.cuda.memory_summary(local_rank)
|
||||
})
|
||||
# Try to clean up distributed resources
|
||||
try:
|
||||
dist.barrier()
|
||||
except:
|
||||
pass
|
||||
raise
|
||||
@@ -1,51 +0,0 @@
|
||||
|
||||
"""Node for saving videos from multi-GPU inference."""
|
||||
|
||||
import os
|
||||
import shutil
|
||||
from typing import Dict, Any
|
||||
import torch.distributed as dist
|
||||
from VEnhancer.video_to_video.context_parallel import get_context_parallel_rank
|
||||
|
||||
class MultiGPUSaverNode:
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls) -> Dict[str, Any]:
|
||||
return {
|
||||
"required": {
|
||||
"video": ("MULTI_GPU_VIDEO",),
|
||||
"filename": ("STRING", {"default": "enhanced.mp4"}),
|
||||
},
|
||||
"optional": {
|
||||
"output_dir": ("STRING", {"default": "ComfyUI/output"}),
|
||||
"overwrite": ("BOOLEAN", {"default": False}),
|
||||
"only_rank_zero": ("BOOLEAN", {"default": True}),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ()
|
||||
FUNCTION = "save_video"
|
||||
CATEGORY = "savers/venhancer"
|
||||
OUTPUT_NODE = True
|
||||
|
||||
def save_video(
|
||||
self,
|
||||
video: str,
|
||||
filename: str,
|
||||
output_dir: str = "ComfyUI/output",
|
||||
overwrite: bool = False,
|
||||
only_rank_zero: bool = True
|
||||
) -> None:
|
||||
if only_rank_zero and get_context_parallel_rank() != 0:
|
||||
dist.barrier() # Wait for rank 0 to finish saving
|
||||
return
|
||||
|
||||
os.makedirs(output_dir, exist_ok=True)
|
||||
output_path = os.path.join(output_dir, filename)
|
||||
|
||||
if os.path.exists(output_path) and not overwrite:
|
||||
raise FileExistsError(f"Output file already exists: {output_path}")
|
||||
|
||||
shutil.copy2(video, output_path)
|
||||
|
||||
if only_rank_zero:
|
||||
dist.barrier()
|
||||
@@ -1,10 +1,17 @@
|
||||
"""Node for running VEnhancer inference on single GPU."""
|
||||
|
||||
from typing import Dict, Any, Tuple
|
||||
import time
|
||||
import torch
|
||||
from loguru import logger
|
||||
|
||||
|
||||
class SingleGPUInferenceNode:
|
||||
"""ComfyUI node for running VEnhancer inference on single GPU."""
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls) -> Dict[str, Any]:
|
||||
"""Define input parameters for video enhancement."""
|
||||
return {
|
||||
"required": {
|
||||
"model": ("SINGLE_GPU_MODEL",),
|
||||
@@ -20,5 +27,66 @@ class SingleGPUInferenceNode:
|
||||
FUNCTION = "enhance_video"
|
||||
CATEGORY = "generators/venhancer"
|
||||
|
||||
def __init__(self):
|
||||
"""Initialize SingleGPUInferenceNode with logging."""
|
||||
self.logger = logger.bind(node="SingleGPUInference")
|
||||
|
||||
def enhance_video(self, model: Any, video: str, **kwargs) -> Tuple[str]:
|
||||
return (model.enhance_a_video(video_path=video, **kwargs),)
|
||||
"""
|
||||
Enhance video using loaded VEnhancer model.
|
||||
|
||||
Args:
|
||||
model: Loaded VEnhancer model instance
|
||||
video: Path to input video
|
||||
**kwargs: Enhancement parameters including:
|
||||
- prompt: Text prompt for enhancement
|
||||
- up_scale: Upscaling factor
|
||||
- target_fps: Target frame rate
|
||||
- noise_aug: Noise augmentation level
|
||||
|
||||
Returns:
|
||||
Tuple[str]: Path to enhanced video
|
||||
|
||||
Raises:
|
||||
RuntimeError: If GPU memory is insufficient
|
||||
Exception: If enhancement fails
|
||||
"""
|
||||
try:
|
||||
start_time = time.time()
|
||||
init_mem = torch.cuda.memory_allocated()
|
||||
|
||||
self.logger.info("Starting video enhancement", extra={
|
||||
"config": {
|
||||
"video": video,
|
||||
"prompt": kwargs.get("prompt"),
|
||||
"up_scale": kwargs.get("up_scale"),
|
||||
"target_fps": kwargs.get("target_fps"),
|
||||
"noise_aug": kwargs.get("noise_aug")
|
||||
}
|
||||
})
|
||||
|
||||
# Run enhancement
|
||||
output_path = model.enhance_a_video(video_path=video, **kwargs)
|
||||
|
||||
# Log metrics
|
||||
enhance_time = time.time() - start_time
|
||||
peak_mem = torch.cuda.max_memory_allocated()
|
||||
mem_used = peak_mem - init_mem
|
||||
|
||||
self.logger.success("Enhancement completed", extra={
|
||||
"metrics": {
|
||||
"total_time": f"{enhance_time:.2f}s",
|
||||
"gpu_memory_used": f"{mem_used/1024/1024:.1f}MB",
|
||||
"peak_gpu_memory": f"{peak_mem/1024/1024:.1f}MB",
|
||||
"gpu_utilization": f"{torch.cuda.utilization()}%"
|
||||
}
|
||||
})
|
||||
|
||||
return (output_path,)
|
||||
|
||||
except Exception as e:
|
||||
self.logger.exception("Enhancement failed", extra={
|
||||
"error": str(e),
|
||||
"gpu_state": torch.cuda.memory_summary()
|
||||
})
|
||||
raise
|
||||
@@ -1,12 +1,19 @@
|
||||
"""Node for loading VEnhancer model in single GPU mode."""
|
||||
|
||||
from typing import Dict, Any, Tuple
|
||||
import time
|
||||
import torch
|
||||
from loguru import logger
|
||||
from VEnhancer.configs.venhnacer_config import VEnhancerConfig
|
||||
from VEnhancer.enhance_a_video import VEnhancer
|
||||
|
||||
|
||||
class SingleGPUVEnhancerLoader:
|
||||
"""ComfyUI node for loading VEnhancer model in single GPU mode."""
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls) -> Dict[str, Any]:
|
||||
"""Define input parameters for model loading."""
|
||||
return {
|
||||
"required": {
|
||||
"version": (["v1", "v2"], {"default": "v1"}),
|
||||
@@ -25,9 +32,48 @@ class SingleGPUVEnhancerLoader:
|
||||
FUNCTION = "load_model"
|
||||
CATEGORY = "loaders/venhancer"
|
||||
|
||||
def load_model(self, **kwargs) -> Tuple[Any]:
|
||||
model_config = VEnhancerConfig(**kwargs)
|
||||
model = VEnhancer(model_config)
|
||||
return (model,)
|
||||
|
||||
def __init__(self):
|
||||
"""Initialize SingleGPUVEnhancerLoader with logging."""
|
||||
self.logger = logger.bind(node="SingleGPULoader")
|
||||
|
||||
def load_model(self, **kwargs) -> Tuple[Any]:
|
||||
"""
|
||||
Load VEnhancer model with specified configuration.
|
||||
|
||||
Args:
|
||||
**kwargs: Configuration parameters for VEnhancer model
|
||||
|
||||
Returns:
|
||||
Tuple[Any]: Tuple containing initialized VEnhancer model
|
||||
|
||||
Raises:
|
||||
RuntimeError: If CUDA is not available
|
||||
Exception: If model loading fails
|
||||
"""
|
||||
try:
|
||||
start_time = time.time()
|
||||
init_mem = torch.cuda.memory_allocated() if torch.cuda.is_available() else 0
|
||||
|
||||
self.logger.info("Loading VEnhancer model", extra={"config": kwargs})
|
||||
|
||||
model_config = VEnhancerConfig(**kwargs)
|
||||
model = VEnhancer(model_config)
|
||||
|
||||
load_time = time.time() - start_time
|
||||
final_mem = torch.cuda.memory_allocated() if torch.cuda.is_available() else 0
|
||||
mem_used = final_mem - init_mem
|
||||
|
||||
self.logger.success("Model loaded successfully", extra={
|
||||
"metrics": {
|
||||
"load_time": f"{load_time:.2f}s",
|
||||
"gpu_memory": f"{mem_used/1024/1024:.1f}MB",
|
||||
"version": kwargs.get("version"),
|
||||
"solver_mode": kwargs.get("solver_mode")
|
||||
}
|
||||
})
|
||||
|
||||
return (model,)
|
||||
|
||||
except Exception as e:
|
||||
self.logger.exception(f"Failed to load model: {str(e)}")
|
||||
raise
|
||||
@@ -1,32 +0,0 @@
|
||||
|
||||
import os
|
||||
import shutil
|
||||
from typing import Dict, Any
|
||||
|
||||
class SingleGPUSaverNode:
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls) -> Dict[str, Any]:
|
||||
return {
|
||||
"required": {
|
||||
"video": ("SINGLE_GPU_VIDEO",),
|
||||
"filename": ("STRING", {"default": "enhanced.mp4"}),
|
||||
},
|
||||
"optional": {
|
||||
"output_dir": ("STRING", {"default": "ComfyUI/output"}),
|
||||
"overwrite": ("BOOLEAN", {"default": False}),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ()
|
||||
FUNCTION = "save_video"
|
||||
CATEGORY = "savers/venhancer"
|
||||
OUTPUT_NODE = True
|
||||
|
||||
def save_video(self, video: str, filename: str, output_dir: str = "ComfyUI/output", overwrite: bool = False) -> None:
|
||||
os.makedirs(output_dir, exist_ok=True)
|
||||
output_path = os.path.join(output_dir, filename)
|
||||
|
||||
if os.path.exists(output_path) and not overwrite:
|
||||
raise FileExistsError(f"Output file already exists: {output_path}")
|
||||
|
||||
shutil.copy2(video, output_path)
|
||||
+194
-21
@@ -1,52 +1,225 @@
|
||||
"""Node for loading video files into ComfyUI."""
|
||||
"""
|
||||
Video loader node module for VEnhancer ComfyUI integration.
|
||||
|
||||
This module provides functionality for loading, validating, and previewing video files
|
||||
in the ComfyUI interface. It includes support for multi-frame preview generation,
|
||||
video validation, and detailed logging of video properties and system metrics.
|
||||
|
||||
usage:
|
||||
loader = VideoLoaderNode()
|
||||
video_path, preview = loader.load_video("input.mp4")
|
||||
"""
|
||||
|
||||
import os
|
||||
import cv2
|
||||
import torch
|
||||
from typing import Dict, Any, Tuple
|
||||
from PIL import Image
|
||||
import numpy as np
|
||||
from loguru import logger
|
||||
|
||||
|
||||
class VideoLoaderNode:
|
||||
"""Node for loading and validating video files."""
|
||||
"""
|
||||
A ComfyUI node for loading and validating video files with preview generation.
|
||||
|
||||
This node handles video file loading, validation, and preview generation for the
|
||||
VEnhancer workflow. It supports various video formats and provides detailed logging
|
||||
of video properties and system resource usage.
|
||||
|
||||
Attributes:
|
||||
logger: Loguru logger instance with node context
|
||||
RETURN_TYPES: Tuple specifying return types (VIDEO, IMAGE)
|
||||
FUNCTION: Name of the primary function to execute
|
||||
CATEGORY: Node category in ComfyUI interface
|
||||
|
||||
Example:
|
||||
node = VideoLoaderNode()
|
||||
video_path, preview = node.load_video(
|
||||
video_path="input.mp4",
|
||||
preview_frames=4,
|
||||
check_video=True
|
||||
)
|
||||
"""
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls) -> Dict[str, Any]:
|
||||
"""
|
||||
Define the input parameters for the node.
|
||||
|
||||
Returns:
|
||||
Dictionary containing input parameter specifications:
|
||||
- video_path: Path to input video file
|
||||
- preview_frames: Number of frames to show in preview grid
|
||||
- check_video: Whether to perform video validation
|
||||
|
||||
Example:
|
||||
{
|
||||
"required": {"video_path": ("STRING", {"default": ""})},
|
||||
"optional": {
|
||||
"preview_frames": ("INT", {"default": 4, "min": 1, "max": 10}),
|
||||
"check_video": ("BOOLEAN", {"default": True})
|
||||
}
|
||||
}
|
||||
"""
|
||||
return {
|
||||
"required": {
|
||||
"video_path": ("STRING", {"default": ""}),
|
||||
},
|
||||
"optional": {
|
||||
"check_video": ("BOOLEAN", {"default": True}),
|
||||
"preview_frames": ("INT", {"default": 4, "min": 1, "max": 10}),
|
||||
"check_video": ("BOOLEAN", {"default": True})
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("VIDEO",)
|
||||
RETURN_TYPES = ("VIDEO", "IMAGE")
|
||||
FUNCTION = "load_video"
|
||||
CATEGORY = "loaders/venhancer"
|
||||
|
||||
def load_video(self, video_path: str, check_video: bool = True) -> Tuple[str]:
|
||||
"""Load and validate video file.
|
||||
|
||||
def __init__(self):
|
||||
"""
|
||||
Initialize the VideoLoaderNode.
|
||||
|
||||
Sets up logging with node-specific context and initializes the node.
|
||||
The logger is configured to track node-specific operations and metrics.
|
||||
"""
|
||||
self.logger = logger.bind(node="VideoLoader")
|
||||
self.logger.info("Initializing VideoLoader node")
|
||||
|
||||
def load_video(self, video_path: str, preview_frames: int = 4, check_video: bool = True) -> Tuple[str, Image.Image]:
|
||||
"""
|
||||
Load and validate a video file, generating preview frames.
|
||||
|
||||
This method handles the complete video loading process, including validation,
|
||||
preview generation, and resource tracking. It provides detailed logging of
|
||||
video properties and system resource usage.
|
||||
|
||||
Args:
|
||||
video_path: Path to video file
|
||||
check_video: Whether to verify video can be opened
|
||||
video_path: Path to the input video file
|
||||
preview_frames: Number of frames to extract for preview grid (default: 4)
|
||||
check_video: Whether to verify video file integrity (default: True)
|
||||
|
||||
Returns:
|
||||
Tuple containing validated video path
|
||||
Tuple containing:
|
||||
- str: Validated video file path
|
||||
- Image: Preview grid of selected frames
|
||||
|
||||
Raises:
|
||||
AssertionError: If video file doesn't exist or can't be opened
|
||||
ValueError: If video path is empty
|
||||
FileNotFoundError: If video file doesn't exist
|
||||
ValueError: If video file cannot be opened
|
||||
Exception: For other video processing errors
|
||||
|
||||
Example:
|
||||
loader = VideoLoaderNode()
|
||||
path, preview = loader.load_video("input.mp4", preview_frames=6)
|
||||
"""
|
||||
if not video_path:
|
||||
raise ValueError("Video path cannot be empty")
|
||||
try:
|
||||
if not os.path.exists(video_path):
|
||||
raise FileNotFoundError(f"Video not found: {video_path}")
|
||||
|
||||
if not os.path.exists(video_path):
|
||||
raise AssertionError(f"Video not found: {video_path}")
|
||||
|
||||
if check_video:
|
||||
import cv2
|
||||
cap = cv2.VideoCapture(video_path)
|
||||
if not cap.isOpened():
|
||||
raise AssertionError(f"Cannot open video: {video_path}")
|
||||
raise ValueError(f"Cannot open video: {video_path}")
|
||||
|
||||
total_frames = int(cap.get(cv2.CAP_PROP_FRAME_COUNT))
|
||||
fps = cap.get(cv2.CAP_PROP_FPS)
|
||||
width = int(cap.get(cv2.CAP_PROP_FRAME_WIDTH))
|
||||
height = int(cap.get(cv2.CAP_PROP_FRAME_HEIGHT))
|
||||
|
||||
self.logger.info(f"Loading video: {video_path}",
|
||||
extra={
|
||||
"frames": total_frames,
|
||||
"fps": fps,
|
||||
"resolution": f"{width}x{height}"
|
||||
})
|
||||
|
||||
frames = []
|
||||
frame_indices = np.linspace(0, total_frames-1, preview_frames, dtype=int)
|
||||
|
||||
for idx in frame_indices:
|
||||
cap.set(cv2.CAP_PROP_POS_FRAMES, idx)
|
||||
ret, frame = cap.read()
|
||||
if ret:
|
||||
frame = cv2.cvtColor(frame, cv2.COLOR_BGR2RGB)
|
||||
frames.append(Image.fromarray(frame))
|
||||
|
||||
grid_size = int(np.ceil(np.sqrt(preview_frames)))
|
||||
grid_width = grid_size * width
|
||||
grid_height = grid_size * height
|
||||
preview_grid = Image.new('RGB', (grid_width, grid_height))
|
||||
|
||||
for idx, frame in enumerate(frames):
|
||||
x = (idx % grid_size) * width
|
||||
y = (idx // grid_size) * height
|
||||
preview_grid.paste(frame, (x, y))
|
||||
|
||||
cap.release()
|
||||
|
||||
return (video_path,)
|
||||
gpu_mem = torch.cuda.memory_allocated() if torch.cuda.is_available() else 0
|
||||
self.logger.success(
|
||||
"Video loaded successfully",
|
||||
extra={
|
||||
"video_info": {
|
||||
"path": video_path,
|
||||
"frames": total_frames,
|
||||
"fps": fps,
|
||||
"resolution": f"{width}x{height}",
|
||||
"preview_frames": preview_frames
|
||||
},
|
||||
"system_info": {
|
||||
"gpu_memory": f"{gpu_mem/1024/1024:.1f}MB"
|
||||
}
|
||||
}
|
||||
)
|
||||
|
||||
return video_path, preview_grid
|
||||
|
||||
except Exception as e:
|
||||
self.logger.exception(f"Failed to load video: {str(e)}")
|
||||
raise
|
||||
|
||||
@classmethod
|
||||
def IS_CHANGED(cls, video_path: str) -> bool:
|
||||
"""
|
||||
Check if the video file has been modified since last load.
|
||||
|
||||
This method monitors video file changes for ComfyUI's caching system.
|
||||
It verifies file existence and tracks modification timestamps.
|
||||
|
||||
Args:
|
||||
video_path: Path to video file to check
|
||||
|
||||
Returns:
|
||||
bool: True if file has changed or doesn't exist, False otherwise
|
||||
|
||||
Example:
|
||||
changed = VideoLoaderNode.IS_CHANGED("input.mp4")
|
||||
"""
|
||||
try:
|
||||
if not os.path.exists(video_path):
|
||||
return True
|
||||
return float(os.path.getmtime(video_path))
|
||||
except Exception as e:
|
||||
logger.error(f"Error checking video modification: {str(e)}")
|
||||
return True
|
||||
|
||||
@classmethod
|
||||
def VALIDATE_INPUTS(cls, **kwargs) -> bool:
|
||||
"""
|
||||
Validate input parameters before processing.
|
||||
|
||||
This method performs validation of all input parameters to ensure
|
||||
they meet the required criteria before video processing begins.
|
||||
|
||||
Args:
|
||||
**kwargs: Keyword arguments containing input parameters
|
||||
|
||||
Returns:
|
||||
bool: True if inputs are valid, False otherwise
|
||||
|
||||
Example:
|
||||
valid = VideoLoaderNode.VALIDATE_INPUTS(video_path="input.mp4")
|
||||
"""
|
||||
if not kwargs.get("video_path"):
|
||||
return False
|
||||
return True
|
||||
@@ -0,0 +1,229 @@
|
||||
"""
|
||||
Video saver node module for VEnhancer ComfyUI integration.
|
||||
|
||||
This module handles saving enhanced videos and generating previews in the ComfyUI interface.
|
||||
It includes support for various output formats, preview generation, and detailed logging
|
||||
of save operations and system metrics.
|
||||
|
||||
Typical usage:
|
||||
saver = VideoSaverNode()
|
||||
output_path = saver.save_video("enhanced.mp4", frames)
|
||||
"""
|
||||
|
||||
import os
|
||||
import cv2
|
||||
import torch
|
||||
from typing import Dict, Any, Tuple, Optional
|
||||
from PIL import Image
|
||||
import numpy as np
|
||||
from loguru import logger
|
||||
|
||||
|
||||
class VideoSaverNode:
|
||||
"""
|
||||
A ComfyUI node for saving enhanced videos with preview generation.
|
||||
|
||||
This node handles video saving operations with support for different formats,
|
||||
preview generation, and detailed metrics logging. It manages output directories
|
||||
and provides feedback through the ComfyUI interface.
|
||||
|
||||
Attributes:
|
||||
logger: Loguru logger instance with node context
|
||||
output_dir: Base directory for saving enhanced videos
|
||||
RETURN_TYPES: Tuple specifying return types (STRING, IMAGE)
|
||||
FUNCTION: Name of the primary function to execute
|
||||
CATEGORY: Node category in ComfyUI interface
|
||||
OUTPUT_NODE: Boolean indicating this is an output node
|
||||
"""
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls) -> Dict[str, Any]:
|
||||
"""
|
||||
Define the input parameters for the video saver node.
|
||||
|
||||
Returns:
|
||||
Dictionary containing input parameter specifications:
|
||||
- video: Enhanced video data to save
|
||||
- filename: Output filename
|
||||
- format: Output video format
|
||||
- generate_preview: Whether to create preview image
|
||||
|
||||
Example:
|
||||
{
|
||||
"required": {
|
||||
"video": ("VIDEO",),
|
||||
"filename": ("STRING", {"default": "enhanced.mp4"})
|
||||
},
|
||||
"optional": {
|
||||
"format": (["mp4", "avi"], {"default": "mp4"}),
|
||||
"generate_preview": ("BOOLEAN", {"default": True})
|
||||
}
|
||||
}
|
||||
"""
|
||||
return {
|
||||
"required": {
|
||||
"video": ("VIDEO",),
|
||||
"filename": ("STRING", {"default": "enhanced.mp4"}),
|
||||
},
|
||||
"optional": {
|
||||
"format": (["mp4", "avi"], {"default": "mp4"}),
|
||||
"generate_preview": ("BOOLEAN", {"default": True}),
|
||||
"preview_frames": ("INT", {"default": 4, "min": 1, "max": 10}),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("STRING", "IMAGE")
|
||||
FUNCTION = "save_video"
|
||||
CATEGORY = "savers/venhancer"
|
||||
OUTPUT_NODE = True
|
||||
|
||||
def __init__(self):
|
||||
"""
|
||||
Initialize the VideoSaverNode.
|
||||
|
||||
Sets up logging with node-specific context and initializes output directories.
|
||||
Ensures the output directory structure exists and is writable.
|
||||
"""
|
||||
self.logger = logger.bind(node="VideoSaver")
|
||||
self.output_dir = os.path.join("ComfyUI", "output")
|
||||
os.makedirs(self.output_dir, exist_ok=True)
|
||||
self.logger.info(f"Initialized VideoSaver node with output dir: {self.output_dir}")
|
||||
|
||||
def save_video(
|
||||
self,
|
||||
video: str,
|
||||
filename: str,
|
||||
format: str = "mp4",
|
||||
generate_preview: bool = True,
|
||||
preview_frames: int = 4
|
||||
) -> Tuple[str, Optional[Image.Image]]:
|
||||
"""
|
||||
Save enhanced video and generate preview.
|
||||
|
||||
This method handles the complete video saving process, including format
|
||||
conversion, preview generation, and resource tracking. It provides detailed
|
||||
logging of the save operation and system resource usage.
|
||||
|
||||
Args:
|
||||
video: Input video data or path
|
||||
filename: Desired output filename
|
||||
format: Output video format (default: "mp4")
|
||||
generate_preview: Whether to create preview (default: True)
|
||||
preview_frames: Number of frames for preview (default: 4)
|
||||
|
||||
Returns:
|
||||
Tuple containing:
|
||||
- str: Path to saved video file
|
||||
- Optional[Image]: Preview image if generated, None otherwise
|
||||
|
||||
Raises:
|
||||
ValueError: If video saving fails
|
||||
OSError: If output directory is not writable
|
||||
Exception: For other video processing errors
|
||||
|
||||
Example:
|
||||
saver = VideoSaverNode()
|
||||
path, preview = saver.save_video(
|
||||
video="input.mp4",
|
||||
filename="enhanced.mp4",
|
||||
generate_preview=True
|
||||
)
|
||||
"""
|
||||
try:
|
||||
start_time = time.time()
|
||||
output_path = os.path.join(self.output_dir, filename)
|
||||
preview = None
|
||||
|
||||
if not os.path.exists(video):
|
||||
raise FileNotFoundError(f"Input video not found: {video}")
|
||||
|
||||
self.logger.info(f"Starting video save operation: {filename}")
|
||||
|
||||
cap = cv2.VideoCapture(video)
|
||||
total_frames = int(cap.get(cv2.CAP_PROP_FRAME_COUNT))
|
||||
fps = cap.get(cv2.CAP_PROP_FPS)
|
||||
width = int(cap.get(cv2.CAP_PROP_FRAME_WIDTH))
|
||||
height = int(cap.get(cv2.CAP_PROP_FRAME_HEIGHT))
|
||||
|
||||
if generate_preview:
|
||||
frames = []
|
||||
frame_indices = np.linspace(0, total_frames-1, preview_frames, dtype=int)
|
||||
|
||||
for idx in frame_indices:
|
||||
cap.set(cv2.CAP_PROP_POS_FRAMES, idx)
|
||||
ret, frame = cap.read()
|
||||
if ret:
|
||||
frame = cv2.cvtColor(frame, cv2.COLOR_BGR2RGB)
|
||||
frames.append(Image.fromarray(frame))
|
||||
|
||||
grid_size = int(np.ceil(np.sqrt(preview_frames)))
|
||||
grid_width = grid_size * width
|
||||
grid_height = grid_size * height
|
||||
preview = Image.new('RGB', (grid_width, grid_height))
|
||||
|
||||
for idx, frame in enumerate(frames):
|
||||
x = (idx % grid_size) * width
|
||||
y = (idx // grid_size) * height
|
||||
preview.paste(frame, (x, y))
|
||||
|
||||
cap.release()
|
||||
|
||||
import shutil
|
||||
shutil.copy2(video, output_path)
|
||||
|
||||
save_time = time.time() - start_time
|
||||
gpu_mem = torch.cuda.memory_allocated() if torch.cuda.is_available() else 0
|
||||
|
||||
self.logger.success(
|
||||
"Video saved successfully",
|
||||
extra={
|
||||
"metrics": {
|
||||
"save_time": f"{save_time:.2f}s",
|
||||
"gpu_memory": f"{gpu_mem/1024/1024:.1f}MB",
|
||||
"video_info": {
|
||||
"path": output_path,
|
||||
"frames": total_frames,
|
||||
"fps": fps,
|
||||
"resolution": f"{width}x{height}",
|
||||
"format": format,
|
||||
"size": os.path.getsize(output_path)
|
||||
}
|
||||
}
|
||||
}
|
||||
)
|
||||
|
||||
return output_path, preview
|
||||
|
||||
except Exception as e:
|
||||
self.logger.exception(
|
||||
"Failed to save video",
|
||||
extra={
|
||||
"error": str(e),
|
||||
"filename": filename,
|
||||
"output_path": output_path,
|
||||
"system_info": {
|
||||
"disk_space": psutil.disk_usage(self.output_dir).free,
|
||||
"gpu_memory": torch.cuda.memory_summary() if torch.cuda.is_available() else "N/A"
|
||||
}
|
||||
}
|
||||
)
|
||||
raise
|
||||
|
||||
@classmethod
|
||||
def IS_CHANGED(cls, filename: str) -> bool:
|
||||
"""
|
||||
Check if the output file already exists.
|
||||
|
||||
This method checks for existing output files to prevent overwriting
|
||||
without explicit confirmation.
|
||||
|
||||
Args:
|
||||
filename: Name of output file to check
|
||||
|
||||
Returns:
|
||||
bool: True if file doesn't exist, False otherwise
|
||||
|
||||
Example:
|
||||
changed = VideoSaverNode.IS_CHANGED("enhanced.mp4")
|
||||
"""
|
||||
return not os.path.exists(os.path.join("ComfyUI", "output", filename))
|
||||
Reference in New Issue
Block a user