Refactoring inference function
This commit is contained in:
+131
-113
@@ -116,6 +116,133 @@ def complete_flow(recurrent_flow_model, flows_tuple, flow_masks, subvideo_length
|
||||
|
||||
return pred_flows_bi
|
||||
|
||||
|
||||
def image_propagation(inpaint_model,
|
||||
frames,
|
||||
masks_dilated,
|
||||
prediction_flows,
|
||||
video_length,
|
||||
subvideo_length,
|
||||
process_size):
|
||||
"""
|
||||
The masked frames are computed by blending original frames and propagated images based on the masks. The process is again segmented if the video is longer than a defined threshold (subvideo_length_img_prop).
|
||||
"""
|
||||
process_width, process_height = process_size
|
||||
masked_frames = frames * (1 - masks_dilated)
|
||||
ic(masked_frames.size())
|
||||
subvideo_length_img_prop = min(100, subvideo_length) # ensure a minimum of 100 frames for image propagation
|
||||
if video_length > subvideo_length_img_prop:
|
||||
updated_frames, updated_masks = [], []
|
||||
pad_len = 10
|
||||
for f in range(0, video_length, subvideo_length_img_prop):
|
||||
s_f = max(0, f - pad_len)
|
||||
e_f = min(video_length, f + subvideo_length_img_prop + pad_len)
|
||||
pad_len_s = max(0, f) - s_f
|
||||
pad_len_e = e_f - min(video_length, f + subvideo_length_img_prop)
|
||||
b, t, _, _, _ = masks_dilated[:, s_f:e_f].size()
|
||||
pred_flows_bi_sub = (prediction_flows[0][:, s_f:e_f-1], prediction_flows[1][:, s_f:e_f-1])
|
||||
prop_imgs_sub, updated_local_masks_sub = inpaint_model.img_propagation(masked_frames[:, s_f:e_f],
|
||||
pred_flows_bi_sub,
|
||||
masks_dilated[:, s_f:e_f],
|
||||
'nearest')
|
||||
updated_frames_sub = frames[:, s_f:e_f] * (1 - masks_dilated[:, s_f:e_f]) + \
|
||||
prop_imgs_sub.view(b, t, 3, process_height, process_width) * masks_dilated[:, s_f:e_f]
|
||||
updated_masks_sub = updated_local_masks_sub.view(b, t, 1, process_height, process_width)
|
||||
|
||||
updated_frames.append(updated_frames_sub[:, pad_len_s:e_f-s_f-pad_len_e])
|
||||
updated_masks.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)
|
||||
else:
|
||||
b, t, _, _, _ = masks_dilated.size()
|
||||
prop_imgs, updated_local_masks = inpaint_model.img_propagation(masked_frames, prediction_flows, masks_dilated, 'nearest')
|
||||
updated_frames = frames * (1 - masks_dilated) + prop_imgs.view(b, t, 3, process_height, process_width) * masks_dilated
|
||||
updated_masks = updated_local_masks.view(b, t, 1, process_height, process_width)
|
||||
torch.cuda.empty_cache()
|
||||
ic(updated_frames.size())
|
||||
ic(updated_masks.size())
|
||||
|
||||
return updated_frames, updated_masks
|
||||
|
||||
|
||||
def feature_propagation(inpaint_model,
|
||||
updated_frames,
|
||||
updated_masks,
|
||||
masks_dilated,
|
||||
prediction_flows,
|
||||
original_frames,
|
||||
video_length,
|
||||
subvideo_length,
|
||||
neighbor_length,
|
||||
ref_stride,
|
||||
process_size):
|
||||
"""
|
||||
Feature Propagation and Transformation: This is done in a loop where features from neighboring frames are propagated using a model. The result is adjusted for color normalization and combined with original frames to produce the final composited frames.
|
||||
"""
|
||||
process_width, process_height = process_size
|
||||
|
||||
comp_frames = [None] * video_length
|
||||
|
||||
neighbor_stride = neighbor_length // 2
|
||||
if video_length > subvideo_length:
|
||||
ref_num = subvideo_length // ref_stride
|
||||
else:
|
||||
ref_num = -1
|
||||
|
||||
for f in tqdm(range(0, video_length, neighbor_stride)):
|
||||
neighbor_ids = [
|
||||
i for i in range(max(0, f - neighbor_stride),
|
||||
min(video_length, f + neighbor_stride + 1))
|
||||
]
|
||||
ref_ids = get_ref_index(f, neighbor_ids, video_length, ref_stride, ref_num)
|
||||
selected_imgs = updated_frames[:, neighbor_ids + ref_ids, :, :, :]
|
||||
selected_masks = masks_dilated[:, neighbor_ids + ref_ids, :, :, :]
|
||||
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)
|
||||
|
||||
# pred_img = selected_imgs # results of image propagation
|
||||
pred_img = inpaint_model(selected_imgs, selected_pred_flows_bi, selected_masks, selected_update_masks, l_t)
|
||||
|
||||
pred_img = pred_img.view(-1, 3, process_height, process_width)
|
||||
|
||||
pred_img = (pred_img + 1) / 2
|
||||
pred_img = pred_img.cpu().permute(0, 2, 3, 1).numpy() * 255
|
||||
binary_masks = masks_dilated[0, neighbor_ids, :, :, :].cpu().permute(
|
||||
0, 2, 3, 1).numpy().astype(np.uint8)
|
||||
for i in range(len(neighbor_ids)):
|
||||
idx = neighbor_ids[i]
|
||||
img = np.array(pred_img[i]).astype(np.uint8) * binary_masks[i] \
|
||||
+ original_frames[idx] * (1 - binary_masks[i])
|
||||
if comp_frames[idx] is None:
|
||||
comp_frames[idx] = img
|
||||
else:
|
||||
comp_frames[idx] = comp_frames[idx].astype(np.float32) * 0.5 + img.astype(np.float32) * 0.5
|
||||
|
||||
comp_frames[idx] = comp_frames[idx].astype(np.uint8)
|
||||
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
ic(type(comp_frames[0]))
|
||||
ic(comp_frames[0].shape)
|
||||
ic(comp_frames[0].dtype)
|
||||
|
||||
|
||||
# For Debugging
|
||||
for idx in range(video_length):
|
||||
f = comp_frames[idx]
|
||||
f = cv2.resize(f, process_size, interpolation = cv2.INTER_CUBIC)
|
||||
f = cv2.cvtColor(f, cv2.COLOR_BGR2RGB)
|
||||
img_save_root = os.path.join("custom_nodes/ComfyUI-ProPainter-Nodes/results", "frames", str(idx).zfill(4)+'.png')
|
||||
imwrite(f, img_save_root)
|
||||
|
||||
return comp_frames
|
||||
|
||||
class ProPainter:
|
||||
"""
|
||||
ProPainter Inpainter
|
||||
@@ -218,9 +345,7 @@ class ProPainter:
|
||||
|
||||
frames, process_size = resize_images(frames, input_size, output_size)
|
||||
ic(frames[0].size)
|
||||
|
||||
process_width, process_height = process_size
|
||||
|
||||
|
||||
flow_masks, masks_dilated = read_masks(mask, input_size, output_size, mask.size(dim=0), flow_mask_dilates, mask_dilates)
|
||||
|
||||
ic(type(flow_masks[0]))
|
||||
@@ -248,10 +373,6 @@ class ProPainter:
|
||||
fix_flow_complete = load_recurrent_flow_model(device)
|
||||
model = load_inpaint_model(device)
|
||||
|
||||
|
||||
##############################################
|
||||
# ProPainter inference
|
||||
##############################################
|
||||
video_length = frames.size(dim=1)
|
||||
print(f'\nProcessing {video_length} frames...')
|
||||
|
||||
@@ -274,112 +395,9 @@ class ProPainter:
|
||||
ic(type(masks_dilated))
|
||||
ic(masks_dilated.size())
|
||||
|
||||
# TODO: Finish refactoring inference function
|
||||
# ---- image propagation ----
|
||||
"""
|
||||
The masked frames are computed by blending original frames and propagated images based on the masks. The process is again segmented if the video is longer than a defined threshold (subvideo_length_img_prop).
|
||||
"""
|
||||
masked_frames = frames * (1 - masks_dilated)
|
||||
ic(masked_frames.size())
|
||||
subvideo_length_img_prop = min(100, subvideo_length) # ensure a minimum of 100 frames for image propagation
|
||||
if video_length > subvideo_length_img_prop:
|
||||
updated_frames, updated_masks = [], []
|
||||
pad_len = 10
|
||||
for f in range(0, video_length, subvideo_length_img_prop):
|
||||
s_f = max(0, f - pad_len)
|
||||
e_f = min(video_length, f + subvideo_length_img_prop + pad_len)
|
||||
pad_len_s = max(0, f) - s_f
|
||||
pad_len_e = e_f - min(video_length, f + subvideo_length_img_prop)
|
||||
# TODO: Check error in prop_imgs_sub.view() for some masks
|
||||
b, t, _, _, _ = masks_dilated[:, s_f:e_f].size()
|
||||
pred_flows_bi_sub = (pred_flows_bi[0][:, s_f:e_f-1], pred_flows_bi[1][:, s_f:e_f-1])
|
||||
prop_imgs_sub, updated_local_masks_sub = model.img_propagation(masked_frames[:, s_f:e_f],
|
||||
pred_flows_bi_sub,
|
||||
masks_dilated[:, s_f:e_f],
|
||||
'nearest')
|
||||
updated_frames_sub = frames[:, s_f:e_f] * (1 - masks_dilated[:, s_f:e_f]) + \
|
||||
prop_imgs_sub.view(b, t, 3, process_height, process_width) * masks_dilated[:, s_f:e_f]
|
||||
updated_masks_sub = updated_local_masks_sub.view(b, t, 1, process_height, process_width)
|
||||
|
||||
updated_frames.append(updated_frames_sub[:, pad_len_s:e_f-s_f-pad_len_e])
|
||||
updated_masks.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)
|
||||
else:
|
||||
b, t, _, _, _ = masks_dilated.size()
|
||||
prop_imgs, updated_local_masks = model.img_propagation(masked_frames, pred_flows_bi, masks_dilated, 'nearest')
|
||||
updated_frames = frames * (1 - masks_dilated) + prop_imgs.view(b, t, 3, height, width) * masks_dilated
|
||||
updated_masks = updated_local_masks.view(b, t, 1, height, width)
|
||||
torch.cuda.empty_cache()
|
||||
ic(updated_frames.size())
|
||||
ic(updated_masks.size())
|
||||
|
||||
|
||||
|
||||
# ---- feature propagation + transformer ----
|
||||
"""
|
||||
Feature Propagation and Transformation: This is done in a loop where features from neighboring frames are propagated using a model. The result is adjusted for color normalization and combined with original frames to produce the final composited frames.
|
||||
"""
|
||||
comp_frames = [None] * video_length
|
||||
|
||||
neighbor_stride = neighbor_length // 2
|
||||
if video_length > subvideo_length:
|
||||
ref_num = subvideo_length // ref_stride
|
||||
else:
|
||||
ref_num = -1
|
||||
|
||||
for f in tqdm(range(0, video_length, neighbor_stride)):
|
||||
neighbor_ids = [
|
||||
i for i in range(max(0, f - neighbor_stride),
|
||||
min(video_length, f + neighbor_stride + 1))
|
||||
]
|
||||
ref_ids = get_ref_index(f, neighbor_ids, video_length, ref_stride, ref_num)
|
||||
selected_imgs = updated_frames[:, neighbor_ids + ref_ids, :, :, :]
|
||||
selected_masks = masks_dilated[:, neighbor_ids + ref_ids, :, :, :]
|
||||
selected_update_masks = updated_masks[:, neighbor_ids + ref_ids, :, :, :]
|
||||
selected_pred_flows_bi = (pred_flows_bi[0][:, neighbor_ids[:-1], :, :, :], pred_flows_bi[1][:, neighbor_ids[:-1], :, :, :])
|
||||
|
||||
with torch.no_grad():
|
||||
# 1.0 indicates mask
|
||||
l_t = len(neighbor_ids)
|
||||
|
||||
# pred_img = selected_imgs # results of image propagation
|
||||
pred_img = model(selected_imgs, selected_pred_flows_bi, selected_masks, selected_update_masks, l_t)
|
||||
|
||||
pred_img = pred_img.view(-1, 3, process_height, process_width)
|
||||
|
||||
pred_img = (pred_img + 1) / 2
|
||||
pred_img = pred_img.cpu().permute(0, 2, 3, 1).numpy() * 255
|
||||
binary_masks = masks_dilated[0, neighbor_ids, :, :, :].cpu().permute(
|
||||
0, 2, 3, 1).numpy().astype(np.uint8)
|
||||
for i in range(len(neighbor_ids)):
|
||||
idx = neighbor_ids[i]
|
||||
img = np.array(pred_img[i]).astype(np.uint8) * binary_masks[i] \
|
||||
+ ori_frames[idx] * (1 - binary_masks[i])
|
||||
if comp_frames[idx] is None:
|
||||
comp_frames[idx] = img
|
||||
else:
|
||||
comp_frames[idx] = comp_frames[idx].astype(np.float32) * 0.5 + img.astype(np.float32) * 0.5
|
||||
|
||||
comp_frames[idx] = comp_frames[idx].astype(np.uint8)
|
||||
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
ic(type(comp_frames[0]))
|
||||
ic(comp_frames[0].shape)
|
||||
ic(comp_frames[0].dtype)
|
||||
|
||||
|
||||
# For Debugging
|
||||
for idx in range(video_length):
|
||||
f = comp_frames[idx]
|
||||
f = cv2.resize(f, process_size, interpolation = cv2.INTER_CUBIC)
|
||||
f = cv2.cvtColor(f, cv2.COLOR_BGR2RGB)
|
||||
img_save_root = os.path.join("custom_nodes/ComfyUI-ProPainter-Nodes/results", "frames", str(idx).zfill(4)+'.png')
|
||||
imwrite(f, img_save_root)
|
||||
|
||||
updated_frames, updated_masks = image_propagation(model, frames, masks_dilated, pred_flows_bi, video_length, subvideo_length, process_size)
|
||||
|
||||
comp_frames = feature_propagation(model, updated_frames, updated_masks, masks_dilated, pred_flows_bi, ori_frames, video_length, subvideo_length, neighbor_length, ref_stride, process_size)
|
||||
|
||||
output_frames = [torch.from_numpy(frame.astype(np.float32) / 255.0) for frame in comp_frames]
|
||||
ic(output_frames[0].size())
|
||||
|
||||
Reference in New Issue
Block a user