georgefen/Face-Landmark-ControlNet
116
1'''2modified by lihaoweicv3pytorch version4'''5 6'''7M-LSD8Copyright 2021-present NAVER Corp.9Apache License v2.010'''11 12import os13import numpy as np14import cv215import torch16from torch.nn import functional as F17 18 19def deccode_output_score_and_ptss(tpMap, topk_n = 200, ksize = 5):20 '''21 tpMap:22 center: tpMap[1, 0, :, :]23 displacement: tpMap[1, 1:5, :, :]24 '''25 b, c, h, w = tpMap.shape26 assert b==1, 'only support bsize==1'27 displacement = tpMap[:, 1:5, :, :][0]28 center = tpMap[:, 0, :, :]29 heat = torch.sigmoid(center)30 hmax = F.max_pool2d( heat, (ksize, ksize), stride=1, padding=(ksize-1)//2)31 keep = (hmax == heat).float()32 heat = heat * keep33 heat = heat.reshape(-1, )34 35 scores, indices = torch.topk(heat, topk_n, dim=-1, largest=True)36 yy = torch.floor_divide(indices, w).unsqueeze(-1)37 xx = torch.fmod(indices, w).unsqueeze(-1)38 ptss = torch.cat((yy, xx),dim=-1)39 40 ptss = ptss.detach().cpu().numpy()41 scores = scores.detach().cpu().numpy()42 displacement = displacement.detach().cpu().numpy()43 displacement = displacement.transpose((1,2,0))44 return ptss, scores, displacement45 46 47def pred_lines(image, model,48 input_shape=[512, 512],49 score_thr=0.10,50 dist_thr=20.0):51 h, w, _ = image.shape52 h_ratio, w_ratio = [h / input_shape[0], w / input_shape[1]]53 54 resized_image = np.concatenate([cv2.resize(image, (input_shape[1], input_shape[0]), interpolation=cv2.INTER_AREA),55 np.ones([input_shape[0], input_shape[1], 1])], axis=-1)56 57 resized_image = resized_image.transpose((2,0,1))58 batch_image = np.expand_dims(resized_image, axis=0).astype('float32')59 batch_image = (batch_image / 127.5) - 1.060 61 batch_image = torch.from_numpy(batch_image).float().cuda()62 outputs = model(batch_image)63 pts, pts_score, vmap = deccode_output_score_and_ptss(outputs, 200, 3)64 start = vmap[:, :, :2]65 end = vmap[:, :, 2:]66 dist_map = np.sqrt(np.sum((start - end) ** 2, axis=-1))67 68 segments_list = []69 for center, score in zip(pts, pts_score):70 y, x = center71 distance = dist_map[y, x]72 if score > score_thr and distance > dist_thr:73 disp_x_start, disp_y_start, disp_x_end, disp_y_end = vmap[y, x, :]74 x_start = x + disp_x_start75 y_start = y + disp_y_start76 x_end = x + disp_x_end77 y_end = y + disp_y_end78 segments_list.append([x_start, y_start, x_end, y_end])79 80 lines = 2 * np.array(segments_list) # 256 > 51281 lines[:, 0] = lines[:, 0] * w_ratio82 lines[:, 1] = lines[:, 1] * h_ratio83 lines[:, 2] = lines[:, 2] * w_ratio84 lines[:, 3] = lines[:, 3] * h_ratio85 86 return lines87 88 89def pred_squares(image,90 model,91 input_shape=[512, 512],92 params={'score': 0.06,93 'outside_ratio': 0.28,94 'inside_ratio': 0.45,95 'w_overlap': 0.0,96 'w_degree': 1.95,97 'w_length': 0.0,98 'w_area': 1.86,99 'w_center': 0.14}):100 '''101 shape = [height, width]102 '''103 h, w, _ = image.shape104 original_shape = [h, w]105 106 resized_image = np.concatenate([cv2.resize(image, (input_shape[0], input_shape[1]), interpolation=cv2.INTER_AREA),107 np.ones([input_shape[0], input_shape[1], 1])], axis=-1)108 resized_image = resized_image.transpose((2, 0, 1))109 batch_image = np.expand_dims(resized_image, axis=0).astype('float32')110 batch_image = (batch_image / 127.5) - 1.0111 112 batch_image = torch.from_numpy(batch_image).float().cuda()113 outputs = model(batch_image)114 115 pts, pts_score, vmap = deccode_output_score_and_ptss(outputs, 200, 3)116 start = vmap[:, :, :2] # (x, y)117 end = vmap[:, :, 2:] # (x, y)118 dist_map = np.sqrt(np.sum((start - end) ** 2, axis=-1))119 120 junc_list = []121 segments_list = []122 for junc, score in zip(pts, pts_score):123 y, x = junc124 distance = dist_map[y, x]125 if score > params['score'] and distance > 20.0:126 junc_list.append([x, y])127 disp_x_start, disp_y_start, disp_x_end, disp_y_end = vmap[y, x, :]128 d_arrow = 1.0129 x_start = x + d_arrow * disp_x_start130 y_start = y + d_arrow * disp_y_start131 x_end = x + d_arrow * disp_x_end132 y_end = y + d_arrow * disp_y_end133 segments_list.append([x_start, y_start, x_end, y_end])134 135 segments = np.array(segments_list)136 137 ####### post processing for squares138 # 1. get unique lines139 point = np.array([[0, 0]])140 point = point[0]141 start = segments[:, :2]142 end = segments[:, 2:]143 diff = start - end144 a = diff[:, 1]145 b = -diff[:, 0]146 c = a * start[:, 0] + b * start[:, 1]147 148 d = np.abs(a * point[0] + b * point[1] - c) / np.sqrt(a ** 2 + b ** 2 + 1e-10)149 theta = np.arctan2(diff[:, 0], diff[:, 1]) * 180 / np.pi150 theta[theta < 0.0] += 180151 hough = np.concatenate([d[:, None], theta[:, None]], axis=-1)152 153 d_quant = 1154 theta_quant = 2155 hough[:, 0] //= d_quant156 hough[:, 1] //= theta_quant157 _, indices, counts = np.unique(hough, axis=0, return_index=True, return_counts=True)158 159 acc_map = np.zeros([512 // d_quant + 1, 360 // theta_quant + 1], dtype='float32')160 idx_map = np.zeros([512 // d_quant + 1, 360 // theta_quant + 1], dtype='int32') - 1161 yx_indices = hough[indices, :].astype('int32')162 acc_map[yx_indices[:, 0], yx_indices[:, 1]] = counts163 idx_map[yx_indices[:, 0], yx_indices[:, 1]] = indices164 165 acc_map_np = acc_map166 # acc_map = acc_map[None, :, :, None]167 #168 # ### fast suppression using tensorflow op169 # acc_map = tf.constant(acc_map, dtype=tf.float32)170 # max_acc_map = tf.keras.layers.MaxPool2D(pool_size=(5, 5), strides=1, padding='same')(acc_map)171 # acc_map = acc_map * tf.cast(tf.math.equal(acc_map, max_acc_map), tf.float32)172 # flatten_acc_map = tf.reshape(acc_map, [1, -1])173 # topk_values, topk_indices = tf.math.top_k(flatten_acc_map, k=len(pts))174 # _, h, w, _ = acc_map.shape175 # y = tf.expand_dims(topk_indices // w, axis=-1)176 # x = tf.expand_dims(topk_indices % w, axis=-1)177 # yx = tf.concat([y, x], axis=-1)178 179 ### fast suppression using pytorch op180 acc_map = torch.from_numpy(acc_map_np).unsqueeze(0).unsqueeze(0)181 _,_, h, w = acc_map.shape182 max_acc_map = F.max_pool2d(acc_map,kernel_size=5, stride=1, padding=2)183 acc_map = acc_map * ( (acc_map == max_acc_map).float() )184 flatten_acc_map = acc_map.reshape([-1, ])185 186 scores, indices = torch.topk(flatten_acc_map, len(pts), dim=-1, largest=True)187 yy = torch.div(indices, w, rounding_mode='floor').unsqueeze(-1)188 xx = torch.fmod(indices, w).unsqueeze(-1)189 yx = torch.cat((yy, xx), dim=-1)190 191 yx = yx.detach().cpu().numpy()192 193 topk_values = scores.detach().cpu().numpy()194 indices = idx_map[yx[:, 0], yx[:, 1]]195 basis = 5 // 2196 197 merged_segments = []198 for yx_pt, max_indice, value in zip(yx, indices, topk_values):199 y, x = yx_pt200 if max_indice == -1 or value == 0:201 continue202 segment_list = []203 for y_offset in range(-basis, basis + 1):204 for x_offset in range(-basis, basis + 1):205 indice = idx_map[y + y_offset, x + x_offset]206 cnt = int(acc_map_np[y + y_offset, x + x_offset])207 if indice != -1:208 segment_list.append(segments[indice])209 if cnt > 1:210 check_cnt = 1211 current_hough = hough[indice]212 for new_indice, new_hough in enumerate(hough):213 if (current_hough == new_hough).all() and indice != new_indice:214 segment_list.append(segments[new_indice])215 check_cnt += 1216 if check_cnt == cnt:217 break218 group_segments = np.array(segment_list).reshape([-1, 2])219 sorted_group_segments = np.sort(group_segments, axis=0)220 x_min, y_min = sorted_group_segments[0, :]221 x_max, y_max = sorted_group_segments[-1, :]222 223 deg = theta[max_indice]224 if deg >= 90:225 merged_segments.append([x_min, y_max, x_max, y_min])226 else:227 merged_segments.append([x_min, y_min, x_max, y_max])228 229 # 2. get intersections230 new_segments = np.array(merged_segments) # (x1, y1, x2, y2)231 start = new_segments[:, :2] # (x1, y1)232 end = new_segments[:, 2:] # (x2, y2)233 new_centers = (start + end) / 2.0234 diff = start - end235 dist_segments = np.sqrt(np.sum(diff ** 2, axis=-1))236 237 # ax + by = c238 a = diff[:, 1]239 b = -diff[:, 0]240 c = a * start[:, 0] + b * start[:, 1]241 pre_det = a[:, None] * b[None, :]242 det = pre_det - np.transpose(pre_det)243 244 pre_inter_y = a[:, None] * c[None, :]245 inter_y = (pre_inter_y - np.transpose(pre_inter_y)) / (det + 1e-10)246 pre_inter_x = c[:, None] * b[None, :]247 inter_x = (pre_inter_x - np.transpose(pre_inter_x)) / (det + 1e-10)248 inter_pts = np.concatenate([inter_x[:, :, None], inter_y[:, :, None]], axis=-1).astype('int32')249 250 # 3. get corner information251 # 3.1 get distance252 '''253 dist_segments:254 | dist(0), dist(1), dist(2), ...|255 dist_inter_to_segment1:256 | dist(inter,0), dist(inter,0), dist(inter,0), ... |257 | dist(inter,1), dist(inter,1), dist(inter,1), ... |258 ...259 dist_inter_to_semgnet2:260 | dist(inter,0), dist(inter,1), dist(inter,2), ... |261 | dist(inter,0), dist(inter,1), dist(inter,2), ... |262 ...263 '''264 265 dist_inter_to_segment1_start = np.sqrt(266 np.sum(((inter_pts - start[:, None, :]) ** 2), axis=-1, keepdims=True)) # [n_batch, n_batch, 1]267 dist_inter_to_segment1_end = np.sqrt(268 np.sum(((inter_pts - end[:, None, :]) ** 2), axis=-1, keepdims=True)) # [n_batch, n_batch, 1]269 dist_inter_to_segment2_start = np.sqrt(270 np.sum(((inter_pts - start[None, :, :]) ** 2), axis=-1, keepdims=True)) # [n_batch, n_batch, 1]271 dist_inter_to_segment2_end = np.sqrt(272 np.sum(((inter_pts - end[None, :, :]) ** 2), axis=-1, keepdims=True)) # [n_batch, n_batch, 1]273 274 # sort ascending275 dist_inter_to_segment1 = np.sort(276 np.concatenate([dist_inter_to_segment1_start, dist_inter_to_segment1_end], axis=-1),277 axis=-1) # [n_batch, n_batch, 2]278 dist_inter_to_segment2 = np.sort(279 np.concatenate([dist_inter_to_segment2_start, dist_inter_to_segment2_end], axis=-1),280 axis=-1) # [n_batch, n_batch, 2]281 282 # 3.2 get degree283 inter_to_start = new_centers[:, None, :] - inter_pts284 deg_inter_to_start = np.arctan2(inter_to_start[:, :, 1], inter_to_start[:, :, 0]) * 180 / np.pi285 deg_inter_to_start[deg_inter_to_start < 0.0] += 360286 inter_to_end = new_centers[None, :, :] - inter_pts287 deg_inter_to_end = np.arctan2(inter_to_end[:, :, 1], inter_to_end[:, :, 0]) * 180 / np.pi288 deg_inter_to_end[deg_inter_to_end < 0.0] += 360289 290 '''291 B -- G292 | |293 C -- R294 B : blue / G: green / C: cyan / R: red295 296 0 -- 1297 | |298 3 -- 2299 '''300 # rename variables301 deg1_map, deg2_map = deg_inter_to_start, deg_inter_to_end302 # sort deg ascending303 deg_sort = np.sort(np.concatenate([deg1_map[:, :, None], deg2_map[:, :, None]], axis=-1), axis=-1)304 305 deg_diff_map = np.abs(deg1_map - deg2_map)306 # we only consider the smallest degree of intersect307 deg_diff_map[deg_diff_map > 180] = 360 - deg_diff_map[deg_diff_map > 180]308 309 # define available degree range310 deg_range = [60, 120]311 312 corner_dict = {corner_info: [] for corner_info in range(4)}313 inter_points = []314 for i in range(inter_pts.shape[0]):315 for j in range(i + 1, inter_pts.shape[1]):316 # i, j > line index, always i < j317 x, y = inter_pts[i, j, :]318 deg1, deg2 = deg_sort[i, j, :]319 deg_diff = deg_diff_map[i, j]320 321 check_degree = deg_diff > deg_range[0] and deg_diff < deg_range[1]322 323 outside_ratio = params['outside_ratio'] # over ratio >>> drop it!324 inside_ratio = params['inside_ratio'] # over ratio >>> drop it!325 check_distance = ((dist_inter_to_segment1[i, j, 1] >= dist_segments[i] and \326 dist_inter_to_segment1[i, j, 0] <= dist_segments[i] * outside_ratio) or \327 (dist_inter_to_segment1[i, j, 1] <= dist_segments[i] and \328 dist_inter_to_segment1[i, j, 0] <= dist_segments[i] * inside_ratio)) and \329 ((dist_inter_to_segment2[i, j, 1] >= dist_segments[j] and \330 dist_inter_to_segment2[i, j, 0] <= dist_segments[j] * outside_ratio) or \331 (dist_inter_to_segment2[i, j, 1] <= dist_segments[j] and \332 dist_inter_to_segment2[i, j, 0] <= dist_segments[j] * inside_ratio))333 334 if check_degree and check_distance:335 corner_info = None336 337 if (deg1 >= 0 and deg1 <= 45 and deg2 >= 45 and deg2 <= 120) or \338 (deg2 >= 315 and deg1 >= 45 and deg1 <= 120):339 corner_info, color_info = 0, 'blue'340 elif (deg1 >= 45 and deg1 <= 125 and deg2 >= 125 and deg2 <= 225):341 corner_info, color_info = 1, 'green'342 elif (deg1 >= 125 and deg1 <= 225 and deg2 >= 225 and deg2 <= 315):343 corner_info, color_info = 2, 'black'344 elif (deg1 >= 0 and deg1 <= 45 and deg2 >= 225 and deg2 <= 315) or \345 (deg2 >= 315 and deg1 >= 225 and deg1 <= 315):346 corner_info, color_info = 3, 'cyan'347 else:348 corner_info, color_info = 4, 'red' # we don't use it349 continue350 351 corner_dict[corner_info].append([x, y, i, j])352 inter_points.append([x, y])353 354 square_list = []355 connect_list = []356 segments_list = []357 for corner0 in corner_dict[0]:358 for corner1 in corner_dict[1]:359 connect01 = False360 for corner0_line in corner0[2:]:361 if corner0_line in corner1[2:]:362 connect01 = True363 break364 if connect01:365 for corner2 in corner_dict[2]:366 connect12 = False367 for corner1_line in corner1[2:]:368 if corner1_line in corner2[2:]:369 connect12 = True370 break371 if connect12:372 for corner3 in corner_dict[3]:373 connect23 = False374 for corner2_line in corner2[2:]:375 if corner2_line in corner3[2:]:376 connect23 = True377 break378 if connect23:379 for corner3_line in corner3[2:]:380 if corner3_line in corner0[2:]:381 # SQUARE!!!382 '''383 0 -- 1384 | |385 3 -- 2386 square_list:387 order: 0 > 1 > 2 > 3388 | x0, y0, x1, y1, x2, y2, x3, y3 |389 | x0, y0, x1, y1, x2, y2, x3, y3 |390 ...391 connect_list:392 order: 01 > 12 > 23 > 30393 | line_idx01, line_idx12, line_idx23, line_idx30 |394 | line_idx01, line_idx12, line_idx23, line_idx30 |395 ...396 segments_list:397 order: 0 > 1 > 2 > 3398 | line_idx0_i, line_idx0_j, line_idx1_i, line_idx1_j, line_idx2_i, line_idx2_j, line_idx3_i, line_idx3_j |399 | line_idx0_i, line_idx0_j, line_idx1_i, line_idx1_j, line_idx2_i, line_idx2_j, line_idx3_i, line_idx3_j |400 ...401 '''402 square_list.append(corner0[:2] + corner1[:2] + corner2[:2] + corner3[:2])403 connect_list.append([corner0_line, corner1_line, corner2_line, corner3_line])404 segments_list.append(corner0[2:] + corner1[2:] + corner2[2:] + corner3[2:])405 406 def check_outside_inside(segments_info, connect_idx):407 # return 'outside or inside', min distance, cover_param, peri_param408 if connect_idx == segments_info[0]:409 check_dist_mat = dist_inter_to_segment1410 else:411 check_dist_mat = dist_inter_to_segment2412 413 i, j = segments_info414 min_dist, max_dist = check_dist_mat[i, j, :]415 connect_dist = dist_segments[connect_idx]416 if max_dist > connect_dist:417 return 'outside', min_dist, 0, 1418 else:419 return 'inside', min_dist, -1, -1420 421 top_square = None422 423 try:424 map_size = input_shape[0] / 2425 squares = np.array(square_list).reshape([-1, 4, 2])426 score_array = []427 connect_array = np.array(connect_list)428 segments_array = np.array(segments_list).reshape([-1, 4, 2])429 430 # get degree of corners:431 squares_rollup = np.roll(squares, 1, axis=1)432 squares_rolldown = np.roll(squares, -1, axis=1)433 vec1 = squares_rollup - squares434 normalized_vec1 = vec1 / (np.linalg.norm(vec1, axis=-1, keepdims=True) + 1e-10)435 vec2 = squares_rolldown - squares436 normalized_vec2 = vec2 / (np.linalg.norm(vec2, axis=-1, keepdims=True) + 1e-10)437 inner_products = np.sum(normalized_vec1 * normalized_vec2, axis=-1) # [n_squares, 4]438 squares_degree = np.arccos(inner_products) * 180 / np.pi # [n_squares, 4]439 440 # get square score441 overlap_scores = []442 degree_scores = []443 length_scores = []444 445 for connects, segments, square, degree in zip(connect_array, segments_array, squares, squares_degree):446 '''447 0 -- 1448 | |449 3 -- 2450 451 # segments: [4, 2]452 # connects: [4]453 '''454 455 ###################################### OVERLAP SCORES456 cover = 0457 perimeter = 0458 # check 0 > 1 > 2 > 3459 square_length = []460 461 for start_idx in range(4):462 end_idx = (start_idx + 1) % 4463 464 connect_idx = connects[start_idx] # segment idx of segment01465 start_segments = segments[start_idx]466 end_segments = segments[end_idx]467 468 start_point = square[start_idx]469 end_point = square[end_idx]470 471 # check whether outside or inside472 start_position, start_min, start_cover_param, start_peri_param = check_outside_inside(start_segments,473 connect_idx)474 end_position, end_min, end_cover_param, end_peri_param = check_outside_inside(end_segments, connect_idx)475 476 cover += dist_segments[connect_idx] + start_cover_param * start_min + end_cover_param * end_min477 perimeter += dist_segments[connect_idx] + start_peri_param * start_min + end_peri_param * end_min478 479 square_length.append(480 dist_segments[connect_idx] + start_peri_param * start_min + end_peri_param * end_min)481 482 overlap_scores.append(cover / perimeter)483 ######################################484 ###################################### DEGREE SCORES485 '''486 deg0 vs deg2487 deg1 vs deg3488 '''489 deg0, deg1, deg2, deg3 = degree490 deg_ratio1 = deg0 / deg2491 if deg_ratio1 > 1.0:492 deg_ratio1 = 1 / deg_ratio1493 deg_ratio2 = deg1 / deg3494 if deg_ratio2 > 1.0:495 deg_ratio2 = 1 / deg_ratio2496 degree_scores.append((deg_ratio1 + deg_ratio2) / 2)497 ######################################498 ###################################### LENGTH SCORES499 '''500 len0 vs len2501 len1 vs len3502 '''503 len0, len1, len2, len3 = square_length504 len_ratio1 = len0 / len2 if len2 > len0 else len2 / len0505 len_ratio2 = len1 / len3 if len3 > len1 else len3 / len1506 length_scores.append((len_ratio1 + len_ratio2) / 2)507 508 ######################################509 510 overlap_scores = np.array(overlap_scores)511 overlap_scores /= np.max(overlap_scores)512 513 degree_scores = np.array(degree_scores)514 # degree_scores /= np.max(degree_scores)515 516 length_scores = np.array(length_scores)517 518 ###################################### AREA SCORES519 area_scores = np.reshape(squares, [-1, 4, 2])520 area_x = area_scores[:, :, 0]521 area_y = area_scores[:, :, 1]522 correction = area_x[:, -1] * area_y[:, 0] - area_y[:, -1] * area_x[:, 0]523 area_scores = np.sum(area_x[:, :-1] * area_y[:, 1:], axis=-1) - np.sum(area_y[:, :-1] * area_x[:, 1:], axis=-1)524 area_scores = 0.5 * np.abs(area_scores + correction)525 area_scores /= (map_size * map_size) # np.max(area_scores)526 ######################################527 528 ###################################### CENTER SCORES529 centers = np.array([[256 // 2, 256 // 2]], dtype='float32') # [1, 2]530 # squares: [n, 4, 2]531 square_centers = np.mean(squares, axis=1) # [n, 2]532 center2center = np.sqrt(np.sum((centers - square_centers) ** 2))533 center_scores = center2center / (map_size / np.sqrt(2.0))534 535 '''536 score_w = [overlap, degree, area, center, length]537 '''538 score_w = [0.0, 1.0, 10.0, 0.5, 1.0]539 score_array = params['w_overlap'] * overlap_scores \540 + params['w_degree'] * degree_scores \541 + params['w_area'] * area_scores \542 - params['w_center'] * center_scores \543 + params['w_length'] * length_scores544 545 best_square = []546 547 sorted_idx = np.argsort(score_array)[::-1]548 score_array = score_array[sorted_idx]549 squares = squares[sorted_idx]550 551 except Exception as e:552 pass553 554 '''return list555 merged_lines, squares, scores556 '''557 558 try:559 new_segments[:, 0] = new_segments[:, 0] * 2 / input_shape[1] * original_shape[1]560 new_segments[:, 1] = new_segments[:, 1] * 2 / input_shape[0] * original_shape[0]561 new_segments[:, 2] = new_segments[:, 2] * 2 / input_shape[1] * original_shape[1]562 new_segments[:, 3] = new_segments[:, 3] * 2 / input_shape[0] * original_shape[0]563 except:564 new_segments = []565 566 try:567 squares[:, :, 0] = squares[:, :, 0] * 2 / input_shape[1] * original_shape[1]568 squares[:, :, 1] = squares[:, :, 1] * 2 / input_shape[0] * original_shape[0]569 except:570 squares = []571 score_array = []572 573 try:574 inter_points = np.array(inter_points)575 inter_points[:, 0] = inter_points[:, 0] * 2 / input_shape[1] * original_shape[1]576 inter_points[:, 1] = inter_points[:, 1] * 2 / input_shape[0] * original_shape[0]577 except:578 inter_points = []579 580 return new_segments, squares, score_array, inter_points581 