diff --git a/src/VEnhancer/configs/distributred_venhancer_config.py b/src/VEnhancer/configs/distributred_venhancer_config.py new file mode 100644 index 0000000..812f992 --- /dev/null +++ b/src/VEnhancer/configs/distributred_venhancer_config.py @@ -0,0 +1,60 @@ +"""Configuration settings for VEnhancer parallel inference.""" + +from typing import Literal +from pydantic import BaseSettings, Field + + +class VEnhancerConfig(BaseSettings): + """Configuration settings for parallel VEnhancer model. + + Attributes: + result_dir: Directory to save enhanced videos + version: Model version to use (v1 or v2) + model_path: Custom path to model checkpoint + solver_mode: Solver mode for inference (fast or normal) + steps: Number of inference steps + guide_scale: Guidance scale for text conditioning + s_cond: Conditioning strength + noise_aug: Noise augmentation level + target_fps: Target frames per second + up_scale: Upscaling factor + repo_id: Hugging Face model repository ID + seed: Random seed for reproducibility + """ + + result_dir: str = Field(default="./results/", description="Directory to save enhanced videos") + version: Literal["v1", "v2"] = Field(default="v1", description="Model version") + model_path: str = Field(default="", description="Path to model checkpoint") + solver_mode: Literal["fast", "normal"] = Field(default="fast", description="Solver mode") + steps: int = Field(default=15, description="Number of inference steps") + guide_scale: float = Field(default=7.5, description="Guidance scale") + s_cond: float = Field(default=8.0, description="Conditioning strength") + noise_aug: int = Field(default=200, ge=0, le=300, description="Noise augmentation level") + target_fps: int = Field(default=24, ge=8, le=60, description="Target FPS") + up_scale: float = Field(default=4.0, ge=1.0, le=8.0, description="Upscaling factor") + repo_id: str = Field(default="jwhejwhe/VEnhancer", description="HuggingFace model repository") + seed: int = Field(default=666, description="Random seed") + + class Config: + env_prefix = "VENHANCER_" + + +class DistributedConfig(BaseSettings): + """Configuration for distributed inference setup. + + Attributes: + world_size: Total number of GPUs + rank: Global rank of current process + local_rank: Local GPU ID + backend: PyTorch distributed backend + init_method: Distributed initialization method + """ + + world_size: int = Field(default=1, ge=1, description="Total number of GPUs") + rank: int = Field(default=0, description="Global rank of current process") + local_rank: int = Field(default=0, description="Local GPU ID") + backend: Literal["nccl", "gloo"] = Field(default="nccl", description="Distributed backend") + init_method: str = Field(default="env://", description="Distributed init method") + + class Config: + env_prefix = "DIST_" diff --git a/src/VEnhancer/configs/venhnacer_config.py b/src/VEnhancer/configs/venhnacer_config.py new file mode 100644 index 0000000..54f3a08 --- /dev/null +++ b/src/VEnhancer/configs/venhnacer_config.py @@ -0,0 +1,38 @@ +"""Configuration settings for VEnhancer model.""" + +from pydantic import BaseSettings, Field + + +class VEnhancerConfig(BaseSettings): + """Configuration settings for VEnhancer model. + + Attributes: + result_dir: Directory to save enhanced videos + version: Model version to use (v1 or v2) + model_path: Custom path to model checkpoint + solver_mode: Solver mode for inference (fast or normal) + steps: Number of inference steps + guide_scale: Guidance scale for text conditioning + s_cond: Conditioning strength + noise_aug: Noise augmentation level + target_fps: Target frames per second + up_scale: Upscaling factor + repo_id: Hugging Face model repository ID + seed: Random seed for reproducibility + """ + + result_dir: str = Field(default="./results/", description="Directory to save enhanced videos") + version: str = Field(default="v1", description="Model version") + model_path: str = Field(default="", description="Path to model checkpoint") + solver_mode: str = Field(default="fast", description="Solver mode (fast or normal)") + steps: int = Field(default=15, description="Number of inference steps") + guide_scale: float = Field(default=7.5, description="Guidance scale") + s_cond: float = Field(default=8.0, description="Conditioning strength") + noise_aug: int = Field(default=200, ge=0, le=300, description="Noise augmentation level") + target_fps: int = Field(default=24, ge=8, le=60, description="Target FPS") + up_scale: float = Field(default=4.0, ge=1.0, le=8.0, description="Upscaling factor") + repo_id: str = Field(default="jwhejwhe/VEnhancer", description="HuggingFace model repository") + seed: int = Field(default=666, description="Random seed") + + class Config: + env_prefix = "VENHANCER_" diff --git a/src/VEnhancer/enhance_a_video.py b/src/VEnhancer/enhance_a_video.py index 2811e2b..a25dcd7 100644 --- a/src/VEnhancer/enhance_a_video.py +++ b/src/VEnhancer/enhance_a_video.py @@ -1,195 +1,213 @@ -from argparse import ArgumentParser, Namespace -import glob +""" +VEnhancer: A text-guided video enhancement model that can upscale resolution, +adjust frame rates, and enhance video quality based on text prompts. +""" +import os +import glob +from typing import List, Optional + +import torch from easydict import EasyDict from huggingface_hub import hf_hub_download -from inference_utils import * -from video_to_video.utils.seed import setup_seed -from video_to_video.video_to_video_model import VideoToVideo +from VEnhancer.inference_utils import ( + get_logger, + load_video, + preprocess, + adjust_resolution, + make_mask_cond, + collate_fn, + tensor2vid, + save_video, + load_prompt_list, +) +from VEnhancer.video_to_video.utils.seed import setup_seed +from VEnhancer.video_to_video.video_to_video_model import VideoToVideo +from VEnhancer.configs.venhnacer_config import VEnhancerConfig logger = get_logger() class VEnhancer: - def __init__( - self, - result_dir="./results/", - version="v1", - model_path="", - solver_mode="fast", - steps=15, - guide_scale=7.5, - s_cond=8, - ): - if not model_path: - self.download_model(version=version) + """Video enhancement model with text guidance. + + This class implements a video enhancement model that can upscale resolution, + adjust frame rates, and enhance video quality based on text prompts. + """ + + def __init__(self, config: Optional[VEnhancerConfig] = None): + """Initialize VEnhancer model. + + Args: + config: VEnhancerConfig object containing model settings + """ + self.config = config or VEnhancerConfig() + + if not self.config.model_path: + self.download_model() else: - self.model_path = model_path + self.model_path = self.config.model_path + assert os.path.exists(self.model_path), "Error: checkpoint Not Found!" logger.info(f"checkpoint_path: {self.model_path}") - self.result_dir = result_dir - os.makedirs(self.result_dir, exist_ok=True) + os.makedirs(self.config.result_dir, exist_ok=True) model_cfg = EasyDict(__name__="model_cfg") model_cfg.model_path = self.model_path self.model = VideoToVideo(model_cfg) - steps = 15 if solver_mode == "fast" else steps - self.solver_mode = solver_mode - self.steps = steps - self.guide_scale = guide_scale - self.s_cond = s_cond + def enhance_a_video( + self, + video_path: str, + prompt: str, + up_scale: Optional[float] = None, + target_fps: Optional[int] = None, + noise_aug: Optional[int] = None, + ) -> str: + """Enhance a video using text guidance. - def enhance_a_video(self, video_path, prompt, up_scale=4, target_fps=24, noise_aug=300): + Args: + video_path: Path to input video file + prompt: Text prompt for enhancement guidance + up_scale: Optional upscaling factor (overrides config) + target_fps: Optional target FPS (overrides config) + noise_aug: Optional noise augmentation level (overrides config) + + Returns: + str: Path to enhanced video file + """ + up_scale = up_scale or self.config.up_scale + target_fps = target_fps or self.config.target_fps + noise_aug = noise_aug or self.config.noise_aug save_name = os.path.splitext(os.path.basename(video_path))[0] - text = prompt - logger.info(f"text: {text}") - caption = text + self.model.positive_prompt + caption = prompt + self.model.positive_prompt + logger.info(f"Processing with prompt: {prompt}") + # Load and preprocess video input_frames, input_fps = load_video(video_path) in_f_num = len(input_frames) - logger.info(f"input frames length: {in_f_num}") - logger.info(f"input fps: {input_fps}") + logger.info(f"Input frames: {in_f_num}, FPS: {input_fps}") + + # Calculate frame interpolation interp_f_num = max(round(target_fps / input_fps) - 1, 0) interp_f_num = min(interp_f_num, 8) target_fps = input_fps * (interp_f_num + 1) - logger.info(f"target_fps: {target_fps}") + logger.info(f"Target FPS: {target_fps}") + # Process video data video_data = preprocess(input_frames) _, _, h, w = video_data.shape - logger.info(f"input resolution: {(h, w)}") target_h, target_w = adjust_resolution(h, w, up_scale) - logger.info(f"target resolution: {(target_h, target_w)}") - - mask_cond = make_mask_cond(in_f_num, interp_f_num) - mask_cond = torch.Tensor(mask_cond).long() + logger.info(f"Resolution: {h}x{w} → {target_h}x{target_w}") + # Prepare conditioning + mask_cond = torch.Tensor(make_mask_cond(in_f_num, interp_f_num)).long() noise_aug = min(max(noise_aug, 0), 300) - logger.info(f"noise augmentation: {noise_aug}") - logger.info(f"scale s is set to: {self.s_cond}") + logger.info(f"Noise augmentation: {noise_aug}") - pre_data = {"video_data": video_data, "y": caption} - pre_data["mask_cond"] = mask_cond - pre_data["s_cond"] = self.s_cond - pre_data["interp_f_num"] = interp_f_num - pre_data["target_res"] = (target_h, target_w) - pre_data["t_hint"] = noise_aug + # Prepare inference data + pre_data = { + "video_data": video_data, + "y": caption, + "mask_cond": mask_cond, + "s_cond": self.config.s_cond, + "interp_f_num": interp_f_num, + "target_res": (target_h, target_w), + "t_hint": noise_aug, + } - total_noise_levels = 900 - setup_seed(666) + setup_seed(self.config.seed) + # Run inference with torch.no_grad(): data_tensor = collate_fn(pre_data, "cuda:0") output = self.model.test( data_tensor, - total_noise_levels, - steps=self.steps, - solver_mode=self.solver_mode, - guide_scale=self.guide_scale, + total_noise_levels=900, + steps=self.config.steps, + solver_mode=self.config.solver_mode, + guide_scale=self.config.guide_scale, noise_aug=noise_aug, ) + # Save results output = tensor2vid(output) - save_video(output, self.result_dir, f"{save_name}.mp4", fps=target_fps) - return os.path.join(self.result_dir, save_name) + save_video(output, self.config.result_dir, f"{save_name}.mp4", fps=target_fps) + return os.path.join(self.config.result_dir, save_name) - def download_model(self, version="v1"): - REPO_ID = "jwhejwhe/VEnhancer" - filename = "venhancer_paper.pt" - if version == "v2": - filename = "venhancer_v2.pt" + def download_model(self) -> None: + """Download model checkpoint from Hugging Face.""" + filename = "venhancer_v2.pt" if self.config.version == "v2" else "venhancer_paper.pt" ckpt_dir = "./ckpts/" os.makedirs(ckpt_dir, exist_ok=True) + local_file = os.path.join(ckpt_dir, filename) if not os.path.exists(local_file): - logger.info(f"Downloading the VEnhancer checkpoint...") - hf_hub_download(repo_id=REPO_ID, filename=filename, local_dir=ckpt_dir) + logger.info("Downloading the VEnhancer checkpoint...") + hf_hub_download( + repo_id=self.config.repo_id, + filename=filename, + local_dir=ckpt_dir + ) self.model_path = local_file + def process_batch( + self, + input_path: str, + prompt: Optional[str] = None, + prompt_path: Optional[str] = None, + filename_as_prompt: bool = False, + ) -> List[str]: + """Process multiple videos in batch. -def parse_args() -> Namespace: - parser = ArgumentParser() + Args: + input_path: Path to input video or directory + prompt: Optional text prompt for all videos + prompt_path: Optional path to file containing prompts + filename_as_prompt: Use filename as prompt - parser.add_argument("--input_path", required=True, type=str, help="input video path") - parser.add_argument("--save_dir", type=str, default="results", help="save directory") - parser.add_argument("--version", type=str, default="v1", choices=["v1", "v2"], help="model version") - parser.add_argument("--model_path", type=str, default="", help="model path") + Returns: + List[str]: Paths to enhanced video files - parser.add_argument("--prompt", type=str, default="a good video", help="prompt") - parser.add_argument("--prompt_path", type=str, default="", help="prompt path") - parser.add_argument("--filename_as_prompt", action="store_true") - - parser.add_argument("--cfg", type=float, default=7.5) - parser.add_argument("--solver_mode", type=str, default="fast", choices=["fast", "normal"], help="fast | normal") - parser.add_argument("--steps", type=int, default=15) - - parser.add_argument("--noise_aug", type=int, default=200, help="noise augmentation") - parser.add_argument("--target_fps", type=int, default=24) - parser.add_argument("--up_scale", type=float, default=4) - parser.add_argument("--s_cond", type=float, default=8) - - return parser.parse_args() - - -def main(): - - args = parse_args() - - input_path = args.input_path - prompt = args.prompt - prompt_path = args.prompt_path - filename_as_prompt = args.filename_as_prompt - model_path = args.model_path - version = args.version - save_dir = args.save_dir - - noise_aug = args.noise_aug - up_scale = args.up_scale - target_fps = args.target_fps - s_cond = args.s_cond - - steps = args.steps - solver_mode = args.solver_mode - guide_scale = args.cfg - - venhancer = VEnhancer( - result_dir=save_dir, - version=version, - model_path=model_path, - solver_mode=solver_mode, - steps=steps, - guide_scale=guide_scale, - s_cond=s_cond, - ) - - if os.path.isdir(input_path): - file_path_list = sorted(glob.glob(os.path.join(input_path, "*.mp4"))) - elif os.path.isfile(input_path): - file_path_list = [input_path] - else: - raise TypeError("input must be a directory or video file!") - - prompt_list = None - if os.path.isfile(prompt_path): - prompt_list = load_prompt_list(prompt_path) - assert len(prompt_list) == len(file_path_list) - - for ind, file_path in enumerate(file_path_list): - logger.info(f"processing video {ind}, file_path: {file_path}") - if filename_as_prompt: - prompt = os.path.splitext(os.path.basename(file_path))[0] - elif prompt_list is not None: - prompt = prompt_list[ind] + Raises: + TypeError: If input_path is neither a file nor directory + """ + # Get list of video files + if os.path.isdir(input_path): + file_path_list = sorted(glob.glob(os.path.join(input_path, "*.mp4"))) + elif os.path.isfile(input_path): + file_path_list = [input_path] else: - prompt_path = os.path.splitext(file_path)[0] + ".txt" - if os.path.isfile(prompt_path): - logger.info(f"prompt_path: {prompt_path}") - prompt = load_prompt_list(prompt_path)[0] - venhancer.enhance_a_video(file_path, prompt, up_scale, target_fps, noise_aug) + raise TypeError("input must be a directory or video file!") + # Handle prompts + prompt_list = None + if os.path.isfile(prompt_path or ""): + prompt_list = load_prompt_list(prompt_path) + assert len(prompt_list) == len(file_path_list) -if __name__ == "__main__": - main() + enhanced_paths = [] + for idx, file_path in enumerate(file_path_list): + logger.info(f"Processing video {idx + 1}/{len(file_path_list)}") + + # Determine prompt for current video + current_prompt = prompt + if filename_as_prompt: + current_prompt = os.path.splitext(os.path.basename(file_path))[0] + elif prompt_list is not None: + current_prompt = prompt_list[idx] + elif not current_prompt: + prompt_file = os.path.splitext(file_path)[0] + ".txt" + if os.path.isfile(prompt_file): + current_prompt = load_prompt_list(prompt_file)[0] + else: + current_prompt = "a good video" + + # Process video + output_path = self.enhance_a_video(file_path, current_prompt) + enhanced_paths.append(output_path) + + return enhanced_paths diff --git a/src/VEnhancer/enhance_a_video_MultiGPU.py b/src/VEnhancer/enhance_a_video_MultiGPU.py index 71745fa..5dd9789 100644 --- a/src/VEnhancer/enhance_a_video_MultiGPU.py +++ b/src/VEnhancer/enhance_a_video_MultiGPU.py @@ -1,215 +1,238 @@ -from argparse import ArgumentParser, Namespace -import glob +""" +VEnhancerMULTIGPU: A distributed text-guided video enhancement model that can upscale resolution, +adjust frame rates, and enhance video quality based on text prompts using multiple GPUs. +""" +import os +import glob +from typing import List, Optional + +import torch +import torch.distributed as dist from easydict import EasyDict from huggingface_hub import hf_hub_download -import torch.cuda -import torch.distributed as dist -from inference_utils import * -from video_to_video.context_parallel import get_context_parallel_rank, initialize_context_parallel -from video_to_video.utils.seed import setup_seed -from video_to_video.video_to_video_model_parallel import VideoToVideoParallel +from VEnhancer.inference_utils import ( + get_logger, + load_video, + preprocess, + adjust_resolution, + make_mask_cond, + collate_fn, + tensor2vid, + save_video, + load_prompt_list, +) +from VEnhancer.video_to_video.context_parallel import ( + get_context_parallel_rank, + initialize_context_parallel, +) +from VEnhancer.video_to_video.utils.seed import setup_seed +from VEnhancer.video_to_video.video_to_video_model_parallel import VideoToVideoParallel +from VEnhancer.configs.distributred_venhancer_config import DistributedConfig logger = get_logger() -class VEnhancer: +class DistributedVEnhancer: + """Distributed video enhancement model with text guidance. + + This class implements a multi-GPU video enhancement model that can upscale resolution, + adjust frame rates, and enhance video quality based on text prompts. + """ + def __init__( self, - result_dir="./results/", - version="v1", - model_path="", - solver_mode="fast", - steps=15, - guide_scale=7.5, - s_cond=8, + dist_config: Optional[DistributedConfig] = None ): - if not model_path: - self.download_model(version=version) + """Initialize distributed VEnhancer model. + + Args: + model_config: VEnhancerConfig object containing model settings + dist_config: DistributedConfig object containing distributed setup + """ + self.dist_config = dist_config or DistributedConfig() + + self._setup_distributed() + + if not self.model_config.model_path: + self.download_model() else: - self.model_path = model_path + self.model_path = self.model_config.model_path + assert os.path.exists(self.model_path), "Error: checkpoint Not Found!" logger.info(f"checkpoint_path: {self.model_path}") - self.result_dir = result_dir - os.makedirs(self.result_dir, exist_ok=True) + os.makedirs(self.model_config.result_dir, exist_ok=True) model_cfg = EasyDict(__name__="model_cfg") model_cfg.model_path = self.model_path self.model = VideoToVideoParallel(model_cfg) - steps = 15 if solver_mode == "fast" else steps - self.solver_mode = solver_mode - self.steps = steps - self.guide_scale = guide_scale - self.s_cond = s_cond + def _setup_distributed(self) -> None: + """Initialize distributed training environment.""" + dist.init_process_group( + backend=self.dist_config.backend, + rank=self.dist_config.rank, + world_size=self.dist_config.world_size, + init_method=self.dist_config.init_method, + ) + torch.cuda.set_device(self.dist_config.local_rank) + initialize_context_parallel(self.dist_config.world_size) + logger.info(f"Initialized process group: rank={self.dist_config.rank}, " + f"world_size={self.dist_config.world_size}") - def enhance_a_video(self, video_path, prompt, up_scale=4, target_fps=24, noise_aug=300): + def enhance_a_video( + self, + video_path: str, + prompt: str, + up_scale: Optional[float] = None, + target_fps: Optional[int] = None, + noise_aug: Optional[int] = None, + ) -> str: + """Enhance a video using text guidance with distributed processing. + + Args: + video_path: Path to input video file + prompt: Text prompt for enhancement guidance + up_scale: Optional upscaling factor (overrides config) + target_fps: Optional target FPS (overrides config) + noise_aug: Optional noise augmentation level (overrides config) + + Returns: + str: Path to enhanced video file + """ + up_scale = up_scale or self.model_config.up_scale + target_fps = target_fps or self.model_config.target_fps + noise_aug = noise_aug or self.model_config.noise_aug save_name = os.path.splitext(os.path.basename(video_path))[0] - text = prompt - logger.info(f"text: {text}") - caption = text + self.model.positive_prompt + caption = prompt + self.model.positive_prompt + logger.info(f"Processing with prompt: {prompt}") + # Load and preprocess video input_frames, input_fps = load_video(video_path) in_f_num = len(input_frames) - logger.info(f"input frames length: {in_f_num}") - logger.info(f"input fps: {input_fps}") + logger.info(f"Input frames: {in_f_num}, FPS: {input_fps}") + + # Calculate frame interpolation interp_f_num = max(round(target_fps / input_fps) - 1, 0) interp_f_num = min(interp_f_num, 8) target_fps = input_fps * (interp_f_num + 1) - logger.info(f"target_fps: {target_fps}") + logger.info(f"Target FPS: {target_fps}") + # Process video data video_data = preprocess(input_frames) _, _, h, w = video_data.shape - logger.info(f"input resolution: {(h, w)}") target_h, target_w = adjust_resolution(h, w, up_scale) - logger.info(f"target resolution: {(target_h, target_w)}") - - mask_cond = make_mask_cond(in_f_num, interp_f_num) - mask_cond = torch.Tensor(mask_cond).long() + logger.info(f"Resolution: {h}x{w} → {target_h}x{target_w}") + # Prepare conditioning + mask_cond = torch.Tensor(make_mask_cond(in_f_num, interp_f_num)).long() noise_aug = min(max(noise_aug, 0), 300) - logger.info(f"noise augmentation: {noise_aug}") - logger.info(f"scale s is set to: {self.s_cond}") + logger.info(f"Noise augmentation: {noise_aug}") - pre_data = {"video_data": video_data, "y": caption} - pre_data["mask_cond"] = mask_cond - pre_data["s_cond"] = self.s_cond - pre_data["interp_f_num"] = interp_f_num - pre_data["target_res"] = (target_h, target_w) - pre_data["t_hint"] = noise_aug + # Prepare inference data + pre_data = { + "video_data": video_data, + "y": caption, + "mask_cond": mask_cond, + "s_cond": self.model_config.s_cond, + "interp_f_num": interp_f_num, + "target_res": (target_h, target_w), + "t_hint": noise_aug, + } - total_noise_levels = 900 - setup_seed(666) + setup_seed(self.model_config.seed) + # Run distributed inference with torch.no_grad(): - data_tensor = collate_fn(pre_data, "cuda") + data_tensor = collate_fn(pre_data, f"cuda:{self.dist_config.local_rank}") output = self.model.test( data_tensor, - total_noise_levels, - steps=self.steps, - solver_mode=self.solver_mode, - guide_scale=self.guide_scale, + total_noise_levels=900, + steps=self.model_config.steps, + solver_mode=self.model_config.solver_mode, + guide_scale=self.model_config.guide_scale, noise_aug=noise_aug, ) output = tensor2vid(output) + # Save results only on main process if get_context_parallel_rank() == 0: - save_video(output, self.result_dir, f"{save_name}.mp4", fps=target_fps) + save_video(output, self.model_config.result_dir, f"{save_name}.mp4", fps=target_fps) dist.barrier() - return os.path.join(self.result_dir, save_name) + return os.path.join(self.model_config.result_dir, save_name) - def download_model(self, version): - REPO_ID = "jwhejwhe/VEnhancer" - filename = "venhancer_paper.pt" - if version == "v2": - filename = "venhancer_v2.pt" + def download_model(self) -> None: + """Download model checkpoint from Hugging Face.""" + filename = "venhancer_v2.pt" if self.model_config.version == "v2" else "venhancer_paper.pt" ckpt_dir = "./ckpts/" os.makedirs(ckpt_dir, exist_ok=True) + local_file = os.path.join(ckpt_dir, filename) if not os.path.exists(local_file): - logger.info(f"Downloading the VEnhancer checkpoint...") - hf_hub_download(repo_id=REPO_ID, filename=filename, local_dir=ckpt_dir) + if get_context_parallel_rank() == 0: + logger.info("Downloading the VEnhancer checkpoint...") + hf_hub_download( + repo_id=self.model_config.repo_id, + filename=filename, + local_dir=ckpt_dir + ) + dist.barrier() self.model_path = local_file + def process_batch( + self, + input_path: str, + prompt: Optional[str] = None, + prompt_path: Optional[str] = None, + filename_as_prompt: bool = False, + ) -> List[str]: + """Process multiple videos in batch with distributed processing. -def parse_args() -> Namespace: - parser = ArgumentParser() + Args: + input_path: Path to input video or directory + prompt: Optional text prompt for all videos + prompt_path: Optional path to file containing prompts + filename_as_prompt: Use filename as prompt - parser.add_argument("--input_path", required=True, type=str, help="input video path") - parser.add_argument("--save_dir", type=str, default="results", help="save directory") - parser.add_argument("--version", type=str, default="v1", choices=["v1", "v2"], help="model version") - parser.add_argument("--model_path", type=str, default="", help="model path") + Returns: + List[str]: Paths to enhanced video files - parser.add_argument("--prompt", type=str, default="a good video", help="prompt") - parser.add_argument("--prompt_path", type=str, default="", help="prompt path") - parser.add_argument("--filename_as_prompt", action="store_true") - - parser.add_argument("--cfg", type=float, default=7.5) - parser.add_argument("--solver_mode", type=str, default="fast", choices=["fast", "normal"], help="fast | normal") - parser.add_argument("--steps", type=int, default=15) - - parser.add_argument("--noise_aug", type=int, default=200, help="noise augmentation") - parser.add_argument("--target_fps", type=int, default=24) - parser.add_argument("--up_scale", type=float, default=4) - parser.add_argument("--s_cond", type=float, default=8) - - return parser.parse_args() - - -def main(): - - args = parse_args() - - world_size = int(os.environ["WORLD_SIZE"]) - rank = int(os.environ["RANK"]) - gpu_id = int(os.environ["LOCAL_RANK"]) - - dist.init_process_group( - backend="nccl", - rank=rank, - world_size=world_size, - init_method="env://", - ) - torch.cuda.set_device(gpu_id) - - initialize_context_parallel(world_size) - - input_path = args.input_path - prompt = args.prompt - prompt_path = args.prompt_path - filename_as_prompt = args.filename_as_prompt - model_path = args.model_path - version = args.version - save_dir = args.save_dir - - noise_aug = args.noise_aug - up_scale = args.up_scale - target_fps = args.target_fps - s_cond = args.s_cond - - steps = args.steps - solver_mode = args.solver_mode - guide_scale = args.cfg - - venhancer = VEnhancer( - result_dir=save_dir, - version=version, - model_path=model_path, - solver_mode=solver_mode, - steps=steps, - guide_scale=guide_scale, - s_cond=s_cond, - ) - - if os.path.isdir(input_path): - file_path_list = sorted(glob.glob(os.path.join(input_path, "*.mp4"))) - elif os.path.isfile(input_path): - file_path_list = [input_path] - else: - raise TypeError("input must be a directory or video file!") - - prompt_list = None - if os.path.isfile(prompt_path): - prompt_list = load_prompt_list(prompt_path) - assert len(prompt_list) == len(file_path_list) - - for ind, file_path in enumerate(file_path_list): - logger.info(f"processing video {ind}, file_path: {file_path}") - if filename_as_prompt: - prompt = os.path.splitext(os.path.basename(file_path))[0] - elif prompt_list is not None: - prompt = prompt_list[ind] + Raises: + TypeError: If input_path is neither a file nor directory + """ + if os.path.isdir(input_path): + file_path_list = sorted(glob.glob(os.path.join(input_path, "*.mp4"))) + elif os.path.isfile(input_path): + file_path_list = [input_path] else: - prompt_path = os.path.splitext(file_path)[0] + ".txt" - if os.path.isfile(prompt_path): - logger.info(f"prompt_path: {prompt_path}") - prompt = load_prompt_list(prompt_path)[0] - venhancer.enhance_a_video(file_path, prompt, up_scale, target_fps, noise_aug) + raise TypeError("input must be a directory or video file!") + prompt_list = None + if os.path.isfile(prompt_path or ""): + prompt_list = load_prompt_list(prompt_path) + assert len(prompt_list) == len(file_path_list) -if __name__ == "__main__": - main() + enhanced_paths = [] + for idx, file_path in enumerate(file_path_list): + logger.info(f"Processing video {idx + 1}/{len(file_path_list)}") + + current_prompt = prompt + if filename_as_prompt: + current_prompt = os.path.splitext(os.path.basename(file_path))[0] + elif prompt_list is not None: + current_prompt = prompt_list[idx] + elif not current_prompt: + prompt_file = os.path.splitext(file_path)[0] + ".txt" + if os.path.isfile(prompt_file): + current_prompt = load_prompt_list(prompt_file)[0] + else: + current_prompt = "a good video" + + output_path = self.enhance_a_video(file_path, current_prompt) + enhanced_paths.append(output_path) + + return enhanced_paths \ No newline at end of file diff --git a/src/VEnhancer/inference_utils.py b/src/VEnhancer/inference_utils.py index 1e5733b..6789e19 100644 --- a/src/VEnhancer/inference_utils.py +++ b/src/VEnhancer/inference_utils.py @@ -10,9 +10,7 @@ import numpy as np import torch import torch.nn.functional as F import torchvision.transforms.functional as transforms_F - -from video_to_video.utils.logger import get_logger - +from VEnhancer.video_to_video.utils.logger import get_logger logger = get_logger()