File size: 2,777 Bytes
d4cbafd | 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 | # Motion Transformer (MTR): https://arxiv.org/abs/2209.13508
# Published at NeurIPS 2022
# Written by Shaoshuai Shi
# All Rights Reserved
import torch
import torch.nn as nn
from ..utils import common_layers
class PointNetPolylineEncoder(nn.Module):
def __init__(self, in_channels, hidden_dim, num_layers=3, num_pre_layers=1, out_channels=None):
super().__init__()
self.pre_mlps = common_layers.build_mlps(
c_in=in_channels,
mlp_channels=[hidden_dim] * num_pre_layers,
ret_before_act=False
)
self.mlps = common_layers.build_mlps(
c_in=hidden_dim * 2,
mlp_channels=[hidden_dim] * (num_layers - num_pre_layers),
ret_before_act=False
)
if out_channels is not None:
self.out_mlps = common_layers.build_mlps(
c_in=hidden_dim, mlp_channels=[hidden_dim, out_channels],
ret_before_act=True, without_norm=True
)
else:
self.out_mlps = None
def forward(self, polylines, polylines_mask):
"""
Args:
polylines (batch_size, num_polylines, num_points_each_polylines, C):
polylines_mask (batch_size, num_polylines, num_points_each_polylines):
Returns:
"""
batch_size, num_polylines, num_points_each_polylines, C = polylines.shape
# pre-mlp
polylines_feature_valid = self.pre_mlps(polylines[polylines_mask]) # (N, C)
polylines_feature = polylines.new_zeros(batch_size, num_polylines, num_points_each_polylines, polylines_feature_valid.shape[-1])
polylines_feature[polylines_mask] = polylines_feature_valid
# get global feature
pooled_feature = polylines_feature.max(dim=2)[0]
polylines_feature = torch.cat((polylines_feature, pooled_feature[:, :, None, :].repeat(1, 1, num_points_each_polylines, 1)), dim=-1)
# mlp
polylines_feature_valid = self.mlps(polylines_feature[polylines_mask])
feature_buffers = polylines_feature.new_zeros(batch_size, num_polylines, num_points_each_polylines, polylines_feature_valid.shape[-1])
feature_buffers[polylines_mask] = polylines_feature_valid
# max-pooling
feature_buffers = feature_buffers.max(dim=2)[0] # (batch_size, num_polylines, C)
# out-mlp
if self.out_mlps is not None:
valid_mask = (polylines_mask.sum(dim=-1) > 0)
feature_buffers_valid = self.out_mlps(feature_buffers[valid_mask]) # (N, C)
feature_buffers = feature_buffers.new_zeros(batch_size, num_polylines, feature_buffers_valid.shape[-1])
feature_buffers[valid_mask] = feature_buffers_valid
return feature_buffers
|