diff --git a/propainter_inference.py b/propainter_inference.py index 7041fea..1a15666 100644 --- a/propainter_inference.py +++ b/propainter_inference.py @@ -7,6 +7,7 @@ from tqdm import tqdm from .model.modules.flow_comp_raft import RAFT_bi from .model.recurrent_flow_completion import RecurrentFlowCompleteNet from .model.propainter import InpaintGenerator +from .utils.model_utils import Models from numpy.typing import NDArray @@ -24,15 +25,21 @@ class ProPainterConfig: fp16: str video_length: int input_size: int + device: torch.device ouput_size: int = field(init=False) process_size: tuple[int, int] = field(init=False) + use_half: bool = field(init=False) - def __post_init__(self): + def __post_init__(self) -> None: + """Initialize output size, process_size and use-half.""" self.output_size = (self.width, self.height) self.process_size = ( self.output_size[0] - self.output_size[0] % 8, self.output_size[1] - self.output_size[1] % 8, ) + self.use_half = self.fp16 == "enable" + if self.device == torch.device("cpu"): + self.use_half = False def get_ref_index( @@ -261,12 +268,13 @@ def feature_propagation( ref_ids = get_ref_index(f, neighbor_ids, config, ref_num) selected_imgs = updated_frames[:, neighbor_ids + ref_ids, :, :, :] selected_masks = masks_dilated[:, neighbor_ids + ref_ids, :, :, :] + if config.use_half: + selected_masks = selected_masks.half() selected_update_masks = updated_masks[:, neighbor_ids + ref_ids, :, :, :] selected_pred_flows_bi = ( prediction_flows[0][:, neighbor_ids[:-1], :, :, :], prediction_flows[1][:, neighbor_ids[:-1], :, :, :], ) - with torch.no_grad(): # 1.0 indicates mask l_t = len(neighbor_ids) @@ -311,37 +319,30 @@ def feature_propagation( def process_inpainting( - frames, - flow_masks, - masks_dilated, - node_config, - raft_model, - flow_model, - inpaint_model, - device, -): - use_half = node_config.fp16 == "enable" - if device == torch.device("cpu"): - use_half = False - + models: Models, + frames: torch.Tensor, + flow_masks: torch.Tensor, + masks_dilated: torch.Tensor, + config: ProPainterConfig, +) -> tuple[torch.Tensor, torch.Tensor, tuple[torch.Tensor, torch.Tensor]]: + """Apply inpainting on video using recurrent flow and ProPainter model.""" with torch.no_grad(): - gt_flows_bi = compute_flow(raft_model, frames, node_config) + gt_flows_bi = compute_flow(models.raft_model, frames, config) - if use_half: + if config.use_half: frames, flow_masks, masks_dilated = ( frames.half(), flow_masks.half(), masks_dilated.half(), ) gt_flows_bi = (gt_flows_bi[0].half(), gt_flows_bi[1].half()) - flow_model = flow_model.half() - inpaint_model = inpaint_model.half() pred_flows_bi = complete_flow( - flow_model, gt_flows_bi, flow_masks, node_config.subvideo_length + models.flow_model, gt_flows_bi, flow_masks, config.subvideo_length ) updated_frames, updated_masks = image_propagation( - inpaint_model, frames, masks_dilated, pred_flows_bi, node_config + models.inpaint_model, frames, masks_dilated, pred_flows_bi, config ) + return updated_frames, updated_masks, pred_flows_bi diff --git a/propainter_nodes.py b/propainter_nodes.py index df923fb..efc77e7 100644 --- a/propainter_nodes.py +++ b/propainter_nodes.py @@ -103,28 +103,26 @@ class ProPainterInpaint: fp16, video_length, input_size, + device ) frames, flow_masks, masks_dilated, original_frames = prepare_frames_and_masks( frames, mask, node_config, device ) - raft_model, flow_model, inpaint_model = initialize_models(device) + models = initialize_models(device, node_config.fp16) print(f"\nProcessing {node_config.video_length} frames...") updated_frames, updated_masks, pred_flows_bi = process_inpainting( + models, frames, flow_masks, masks_dilated, node_config, - raft_model, - flow_model, - inpaint_model, - device, ) composed_frames = feature_propagation( - inpaint_model, + models.inpaint_model, updated_frames, updated_masks, masks_dilated, diff --git a/requirements.txt.example b/requirements.txt.example deleted file mode 100644 index 0fb5350..0000000 --- a/requirements.txt.example +++ /dev/null @@ -1,16 +0,0 @@ -av -addict -einops -future -numpy -scipy -opencv-python -matplotlib -scikit-image -torch>=1.7.1 -torchvision>=0.8.2 -imageio-ffmpeg -pyyaml -requests -timm -yapf \ No newline at end of file diff --git a/utils/model_utils.py b/utils/model_utils.py index 719d999..2676d2c 100644 --- a/utils/model_utils.py +++ b/utils/model_utils.py @@ -1,3 +1,4 @@ +from dataclasses import dataclass import os from torch import device @@ -8,6 +9,13 @@ from ..model.recurrent_flow_completion import RecurrentFlowCompleteNet from ..model.propainter import InpaintGenerator +@dataclass +class Models: + raft_model: RAFT_bi + flow_model: RecurrentFlowCompleteNet + inpaint_model: InpaintGenerator + + pretrain_model_url = "https://github.com/sczhou/ProPainter/releases/download/v0.1.0/" @@ -47,11 +55,14 @@ def load_inpaint_model(device: device) -> InpaintGenerator: return inpaint_model -def initialize_models( - device: device, -) -> tuple[RAFT_bi, RecurrentFlowCompleteNet, InpaintGenerator]: +def initialize_models(device: device, use_half: str) -> Models: "Return initialized inference models." raft_model = load_raft_model(device) flow_model = load_recurrent_flow_model(device) inpaint_model = load_inpaint_model(device) - return raft_model, flow_model, inpaint_model + + if use_half == "enable": + # raft_model = raft_model.half() + flow_model = flow_model.half() + inpaint_model = inpaint_model.half() + return Models(raft_model, flow_model, inpaint_model)