Refactor import statements in inference_utils.py
This commit is contained in:
@@ -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_"
|
||||
@@ -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
@@ -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
|
||||
|
||||
@@ -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
|
||||
@@ -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()
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user