Refactoring with mypy

This commit is contained in:
daniabib
2024-05-16 23:56:07 +00:00
parent 6efcd95df1
commit cf8d0e45df
3 changed files with 13 additions and 13 deletions
+10 -10
View File
@@ -115,7 +115,7 @@ def complete_flow(
"""
flow_length = flows_tuple[0].size(dim=1)
if flow_length > subvideo_length:
pred_flows_f, pred_flows_b = [], []
pred_flows_f_list, pred_flows_b_list = [], []
pad_len = 5
for f in range(0, flow_length, subvideo_length):
s_f = max(0, f - pad_len)
@@ -132,16 +132,16 @@ def complete_flow(
flow_masks[:, s_f : e_f + 1],
)
pred_flows_f.append(
pred_flows_f_list.append(
pred_flows_bi_sub[0][:, pad_len_s : e_f - s_f - pad_len_e]
)
pred_flows_b.append(
pred_flows_b_list.append(
pred_flows_bi_sub[1][:, pad_len_s : e_f - s_f - pad_len_e]
)
torch.cuda.empty_cache()
pred_flows_f = torch.cat(pred_flows_f, dim=1)
pred_flows_b = torch.cat(pred_flows_b, dim=1)
pred_flows_f = torch.cat(pred_flows_f_list, dim=1)
pred_flows_b = torch.cat(pred_flows_b_list, dim=1)
pred_flows_bi = (pred_flows_f, pred_flows_b)
@@ -175,7 +175,7 @@ def image_propagation(
100, config.subvideo_length
) # ensure a minimum of 100 frames for image propagation
if config.video_length > subvideo_length_img_prop:
updated_frames, updated_masks = [], []
updated_frames_list, updated_masks_list = [], []
pad_len = 10
for f in range(0, config.video_length, subvideo_length_img_prop):
s_f = max(0, f - pad_len)
@@ -202,16 +202,16 @@ def image_propagation(
b, t, 1, process_height, process_width
)
updated_frames.append(
updated_frames_list.append(
updated_frames_sub[:, pad_len_s : e_f - s_f - pad_len_e]
)
updated_masks.append(
updated_masks_list.append(
updated_masks_sub[:, pad_len_s : e_f - s_f - pad_len_e]
)
torch.cuda.empty_cache()
updated_frames = torch.cat(updated_frames, dim=1)
updated_masks = torch.cat(updated_masks, dim=1)
updated_frames = torch.cat(updated_frames_list, dim=1)
updated_masks = torch.cat(updated_masks_list, dim=1)
else:
b, t, _, _, _ = masks_dilated.size()
prop_imgs, updated_local_masks = inpaint_model.img_propagation(
+1 -1
View File
@@ -83,7 +83,7 @@ class ProPainterInpaint:
subvideo_length: int,
raft_iter: int,
fp16: str,
) -> tuple[torch.Tensor]:
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
"""Perform inpainting on images input using the ProPainter model inference."""
device = model_management.get_torch_device()
+2 -2
View File
@@ -128,8 +128,8 @@ def read_masks(
masks_dilated.append(Image.fromarray(mask_img * 255))
if len(mask_imgs) == 1:
flow_masks = flow_masks * config.length
masks_dilated = masks_dilated * config.length
flow_masks = flow_masks * config.video_length
masks_dilated = masks_dilated * config.video_length
return flow_masks, masks_dilated