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