| 168 | |
| 169 | |
| 170 | def forward(self, x): |
| 171 | # coordinates |
| 172 | p = x[:, :, 3:].contiguous() |
| 173 | |
| 174 | B, N, _ = p.shape[:3] |
| 175 | # idx = self.sample_fn(p, int(N * self.sample_ratio)).long() |
| 176 | idx = self.sample_fn(p, self.sample_number).long() |
| 177 | center_p = torch.gather(p, 1, idx.unsqueeze(-1).expand(-1, -1, 3)) |
| 178 | # query neighbors. |
| 179 | _, fj = self.grouper(center_p, p, x.permute(0, 2, 1).contiguous()) # [B, N, 6] -> [B, 6, N] -> [B, 6, 1024, 32] |
| 180 | |
| 181 | # [B, 6, 1024] -> [B, channels, 1024, 1] |
| 182 | fj = self.conv1(fj).max(dim=-1, keepdim=True)[0] |
| 183 | |
| 184 | return fj |
| 185 | |
| 186 | |
| 187 | if __name__ == '__main__': |