File size: 4,212 Bytes
5013bb8 | 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 | import numpy as np
import torch
import torch.nn as nn
from .cluster import LaneNetPostProcessor
from .connect import sort_points_by_dist, connect_by_direction
def onehot_encoding(logits, dim=0):
max_idx = torch.argmax(logits, dim, keepdim=True)
one_hot = logits.new_full(logits.shape, 0)
one_hot.scatter_(dim, max_idx, 1)
return one_hot
def onehot_encoding_spread(logits, dim=1):
max_idx = torch.argmax(logits, dim, keepdim=True)
one_hot = logits.new_full(logits.shape, 0)
one_hot.scatter_(dim, max_idx, 1)
one_hot.scatter_(dim, torch.clamp(max_idx-1, min=0), 1)
one_hot.scatter_(dim, torch.clamp(max_idx-2, min=0), 1)
one_hot.scatter_(dim, torch.clamp(max_idx+1, max=logits.shape[dim]-1), 1)
one_hot.scatter_(dim, torch.clamp(max_idx+2, max=logits.shape[dim]-1), 1)
return one_hot
def get_pred_top2_direction(direction, dim=1):
direction = torch.softmax(direction, dim)
idx1 = torch.argmax(direction, dim)
idx1_onehot_spread = onehot_encoding_spread(direction, dim)
idx1_onehot_spread = idx1_onehot_spread.bool()
direction[idx1_onehot_spread] = 0
idx2 = torch.argmax(direction, dim)
direction = torch.stack([idx1, idx2], dim) - 1
return direction
def vectorize(segmentation, embedding, direction, angle_class):
segmentation = segmentation.softmax(0)
embedding = embedding.cpu()
direction = direction.permute(1, 2, 0).cpu()
direction = get_pred_top2_direction(direction, dim=-1)
max_pool_1 = nn.MaxPool2d((1, 5), padding=(0, 2), stride=1)
avg_pool_1 = nn.AvgPool2d((9, 5), padding=(4, 2), stride=1)
max_pool_2 = nn.MaxPool2d((5, 1), padding=(2, 0), stride=1)
avg_pool_2 = nn.AvgPool2d((5, 9), padding=(2, 4), stride=1)
post_processor = LaneNetPostProcessor(dbscan_eps=1.5, postprocess_min_samples=50)
oh_pred = onehot_encoding(segmentation).cpu().numpy()
confidences = []
line_types = []
simplified_coords = []
for i in range(1, oh_pred.shape[0]):
single_mask = oh_pred[i].astype('uint8')
single_embedding = embedding.permute(1, 2, 0)
single_class_inst_mask, single_class_inst_coords = post_processor.postprocess(single_mask, single_embedding)
if single_class_inst_mask is None:
continue
num_inst = len(single_class_inst_coords)
prob = segmentation[i]
prob[single_class_inst_mask == 0] = 0
nms_mask_1 = ((max_pool_1(prob.unsqueeze(0))[0] - prob) < 0.0001).cpu().numpy()
avg_mask_1 = avg_pool_1(prob.unsqueeze(0))[0].cpu().numpy()
nms_mask_2 = ((max_pool_2(prob.unsqueeze(0))[0] - prob) < 0.0001).cpu().numpy()
avg_mask_2 = avg_pool_2(prob.unsqueeze(0))[0].cpu().numpy()
vertical_mask = avg_mask_1 > avg_mask_2
horizontal_mask = ~vertical_mask
nms_mask = (vertical_mask & nms_mask_1) | (horizontal_mask & nms_mask_2)
for j in range(1, num_inst + 1):
full_idx = np.where((single_class_inst_mask == j))
full_lane_coord = np.vstack((full_idx[1], full_idx[0])).transpose()
confidence = prob[single_class_inst_mask == j].mean().item()
idx = np.where(nms_mask & (single_class_inst_mask == j))
if len(idx[0]) == 0:
continue
lane_coordinate = np.vstack((idx[1], idx[0])).transpose()
range_0 = np.max(full_lane_coord[:, 0]) - np.min(full_lane_coord[:, 0])
range_1 = np.max(full_lane_coord[:, 1]) - np.min(full_lane_coord[:, 1])
if range_0 > range_1:
lane_coordinate = sorted(lane_coordinate, key=lambda x: x[0])
else:
lane_coordinate = sorted(lane_coordinate, key=lambda x: x[1])
lane_coordinate = np.stack(lane_coordinate)
lane_coordinate = sort_points_by_dist(lane_coordinate)
lane_coordinate = lane_coordinate.astype('int32')
lane_coordinate = connect_by_direction(lane_coordinate, direction, step=7, per_deg=360 / angle_class)
simplified_coords.append(lane_coordinate)
confidences.append(confidence)
line_types.append(i-1)
return simplified_coords, confidences, line_types
|