WompUniversity/Inpaint-Anything-no-errors
0
1import cv22import numpy as np3 4def paste_object(source, source_mask, target, target_coords, resize_scale=1):5 assert target_coords[0] < target.shape[1] and target_coords[1] < target.shape[0]6 # Find the bounding box of the source_mask7 x, y, w, h = cv2.boundingRect(source_mask)8 assert h < source.shape[0] and w < source.shape[1]9 obj = source[y:y+h, x:x+w]10 obj_msk = source_mask[y:y+h, x:x+w]11 if resize_scale != 1:12 obj = cv2.resize(obj, (0,0), fx=resize_scale, fy=resize_scale)13 obj_msk = cv2.resize(obj_msk, (0,0), fx=resize_scale, fy=resize_scale)14 _, _, w, h = cv2.boundingRect(obj_msk)15 16 xt = max(0, target_coords[0]-w//2)17 yt = max(0, target_coords[1]-h//2)18 if target_coords[0]-w//2 < 0:19 obj = obj[:, w//2-target_coords[0]:]20 obj_msk = obj_msk[:, w//2-target_coords[0]:]21 if target_coords[0]+w//2 > target.shape[1]:22 obj = obj[:, :target.shape[1]-target_coords[0]+w//2]23 obj_msk = obj_msk[:, :target.shape[1]-target_coords[0]+w//2]24 if target_coords[1]-h//2 < 0:25 obj = obj[h//2-target_coords[1]:, :]26 obj_msk = obj_msk[h//2-target_coords[1]:, :]27 if target_coords[1]+h//2 > target.shape[0]:28 obj = obj[:target.shape[0]-target_coords[1]+h//2, :]29 obj_msk = obj_msk[:target.shape[0]-target_coords[1]+h//2, :]30 _, _, w, h = cv2.boundingRect(obj_msk)31 32 target[yt:yt+h, xt:xt+w][obj_msk==255] = obj[obj_msk==255]33 target_mask = np.zeros_like(target)34 target_mask = cv2.cvtColor(target_mask, cv2.COLOR_BGR2GRAY)35 target_mask[yt:yt+h, xt:xt+w][obj_msk==255] = 25536 37 return target, target_mask38 39if __name__ == '__main__':40 source = cv2.imread('example/boat.jpg')41 source_mask = cv2.imread('example/boat_mask_1.png', 0)42 target = cv2.imread('example/hippopotamus.jpg')43 print(source.shape, source_mask.shape, target.shape)44 45 target_coords = (700, 400) # (x, y)46 resize_scale = 147 target, target_mask = paste_object(source, source_mask, target, target_coords, resize_scale)48 cv2.imwrite('target_pasted.png', target)49 cv2.imwrite('target_mask.png', target_mask)50 print(target.shape, target_mask.shape)