diff --git a/propainter.py b/propainter.py index 84455a7..b2312dc 100644 --- a/propainter.py +++ b/propainter.py @@ -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())