Refactoring inference function

This commit is contained in:
daniabib
2024-05-14 21:09:25 +00:00
parent af274f57f9
commit f5715db29e
+131 -113
View File
@@ -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())