Refactoring Models class
This commit is contained in:
+22
-21
@@ -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
|
||||
|
||||
+4
-6
@@ -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,
|
||||
|
||||
@@ -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
|
||||
+15
-4
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user