Refactor import statements in inference_utils.py

This commit is contained in:
vikramxd
2024-11-13 09:04:42 +00:00
parent bf33d8fd97
commit 58661da3bd
5 changed files with 432 additions and 295 deletions
@@ -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_"
+38
View File
@@ -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_"
+156 -138
View File
@@ -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
+177 -154
View File
@@ -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
+1 -3
View File
@@ -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()