File size: 5,316 Bytes
872b0a0 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 | # from skimage.color import rgb2yuv
# import cv2
import numpy as np
import torch
# from fast_slic.avx2 import SlicAvx2 as Slic
# from skimage.segmentation import slic as sk_slic
from models.UnSAMFlow.utils.warp_utils import flow_warp
# def run_slic_pt(img_batch, n_seg=200, compact=10, rd_select=(8, 16), fast=True): # Nx1xHxW
# """
# :param img: Nx3xHxW 0~1 float32
# :param n_seg:
# :param compact:
# :return: Nx1xHxW float32
# """
# B = img_batch.size(0)
# dtype = img_batch.type()
# img_batch = np.split(
# img_batch.detach().cpu().numpy().transpose([0, 2, 3, 1]), B, axis=0)
# out = []
# if fast:
# fast_slic = Slic(num_components=n_seg, compactness=compact, min_size_factor=0.8)
# for img in img_batch:
# img = np.copy((img * 255).squeeze(0).astype(np.uint8), order='C')
# if fast:
# img = cv2.cvtColor(img, cv2.COLOR_RGB2LAB)
# seg = fast_slic.iterate(img)
# else:
# seg = sk_slic(img, n_segments=200, compactness=10)
# if rd_select is not None:
# n_select = np.random.randint(rd_select[0], rd_select[1])
# select_list = np.random.choice(range(0, np.max(seg) + 1), n_select,
# replace=False)
# seg = np.bitwise_or.reduce([seg == seg_id for seg_id in select_list])
# out.append(seg)
# x_out = torch.tensor(np.stack(out)).type(dtype).unsqueeze(1)
# return x_out
def random_crop(img, full_segs, flow, occ_mask, crop_sz):
"""
:param img: Nx6xHxW
:param flows: n * [Nx2xHxW]
:param occ_masks: n * [Nx1xHxW]
:param crop_sz:
:return:
"""
_, _, h, w = img.size()
c_h, c_w = crop_sz
if c_h == h and c_w == w:
return img, flow, occ_mask
x1 = np.random.randint(0, w - c_w)
y1 = np.random.randint(0, h - c_h)
img = img[:, :, y1 : y1 + c_h, x1 : x1 + c_w]
full_segs = full_segs[:, :, y1 : y1 + c_h, x1 : x1 + c_w]
flow = flow[:, :, y1 : y1 + c_h, x1 : x1 + c_w]
occ_mask = occ_mask[:, :, y1 : y1 + c_h, x1 : x1 + c_w]
return img, full_segs, flow, occ_mask
# def semantic_connected_components(
# semseg, class_indices, width_range=(30, 200), height_range=(30, 100)
# ):
# """
# Input:
# semsegs: Onehot semantic segmentations of size [c, H, W]
# class_indices: A list of the indices of the classes of interest.
# For example, [car_idx] or [sign_idx, pole_idx, traffic_light_idx]
# width_range, height_range: The width and height ranges for the objects of interest.
# Output:
# list of masks for cars in the size range
# """
# curr_sem = semseg[class_indices].sum(dim=0)
# curr_sem = (curr_sem[:, :, None] * 255).numpy().astype(np.uint8)
# num_labels, labels = cv2.connectedComponents(curr_sem)
# sem_list = []
# # 0 is background, so ignore them.
# for i in range(1, num_labels):
# curr_obj = labels == i
# hs, ws = np.where(curr_obj)
# h_len, w_len = np.max(hs) - np.min(hs), np.max(ws) - np.min(ws)
# if (height_range[0] < h_len < height_range[1]) and (
# width_range[0] < w_len < width_range[1]
# ):
# if (
# curr_obj.sum() / ((h_len + 1) * (w_len + 1)) > 0.6
# ): # filter some wrong car estimates or largely occluded cars
# sem_list.append(curr_obj)
# return sem_list
# def find_semantic_group(semseg, class_indices, win_width=200):
# curr_sem = semseg[class_indices].sum(dim=0).numpy()
# freq = curr_sem.mean(axis=0) # 1d frequency
# freq_win = np.convolve(
# freq, np.ones(win_width) / win_width, mode="valid"
# ) # find the most frequent window
# if max(freq_win) > 0.1:
# # optimal window: [left, right)
# left = np.argmax(freq_win)
# right = left + win_width
# curr_sem[:, :left] = 0
# curr_sem[:, right:] = 0
# return curr_sem
# else:
# return None
def add_fake_object(input_dict):
# prepare input
img1_ot = input_dict["img1_tgt"]
img2_ot = input_dict["img2_tgt"]
full_seg1_ot = input_dict["full_seg1_st"]
full_seg2_ot = input_dict["full_seg2_st"]
flow_ot = input_dict["flow_tgt"]
noc_ot = input_dict["noc_tgt"]
img = input_dict["img_src"]
obj_mask = input_dict["obj_mask"]
motion = input_dict["motion"][:, :, None, None]
b, _, h, w = img1_ot.shape
N1 = full_seg1_ot.max()
N2 = full_seg2_ot.max()
# add object to frame 1
img1_ot = obj_mask * img + (1 - obj_mask) * img1_ot
full_seg1_ot = obj_mask * (N1 + 1) + (1 - obj_mask) * full_seg1_ot
# add object to frame 2
new_obj_mask = flow_warp(obj_mask, -motion.repeat(1, 1, h, w), pad="zeros")
new_img = flow_warp(img, -motion.repeat(1, 1, h, w), pad="border")
img2_ot = new_obj_mask * new_img + (1 - new_obj_mask) * img2_ot
full_seg2_ot = new_obj_mask * (N2 + 1) + (1 - new_obj_mask) * full_seg2_ot
# change flow
flow_ot = obj_mask * motion + (1 - obj_mask) * flow_ot
noc_ot = torch.max(noc_ot, obj_mask) # where we are confident about flow_ot
return img1_ot, img2_ot, full_seg1_ot, full_seg2_ot, flow_ot, noc_ot, new_obj_mask
|