Spaces:
Running on Zero
Running on Zero
File size: 7,645 Bytes
be88765 | 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 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 | import torch
import torch.nn as nn
import torch.nn.functional as F
import numpy as np
from pytorch3d.ops import ball_query
from . import patchnet as pn
from .utils.common import cal_Z_axis, l2_norm, RodsRotatFormula, sphere_query, var_to_invar, get_voxel_coordinate
class MiniSpinNet(nn.Module):
def __init__(
self,
des_r: float = 3.0,
num_points_per_patch: int = 512,
rad_n: int = 3,
azi_n: int = 20,
ele_n: int = 7,
delta: float = 0.8,
voxel_sample: int = 10,
is_aligned_to_global_z: bool = True,
):
super(MiniSpinNet, self).__init__()
self.des_r = des_r
self.patch_sample = num_points_per_patch
self.rad_n = rad_n
self.azi_n = azi_n
self.ele_n = ele_n
self.delta = delta
self.voxel_sample = voxel_sample
self.is_aligned_to_global_z = is_aligned_to_global_z
self.pnt_layer = nn.Sequential(
nn.Conv2d(3, 16, kernel_size=(1, 1), stride=(1, 1)),
nn.BatchNorm2d(16),
nn.ReLU(True),
)
self.pool_layer = nn.Sequential(
nn.Conv2d(32, 16, kernel_size=(1, 1), stride=(1, 1)),
nn.BatchNorm2d(16),
nn.ReLU(True),
nn.Conv2d(16, 1, kernel_size=(1, 1), stride=(1, 1)),
nn.BatchNorm2d(1),
nn.ReLU(True),
)
self.conv_net = pn.Cylindrical_Net(inchan=16, dim=32)
# self.conv_net = pn.Cylindrical_UNet(inchan=16, dim=32)
def forward(self, pts, kpts, des_r, is_aligned_to_global_z=True, z_axis=None, is_aug=False):
# extract patches
init_patch = self.select_patches(pts, kpts, vicinity=des_r, patch_sample=self.patch_sample)
init_patch = init_patch.squeeze(0)
# print('init_patch:', init_patch)
# print('kpts:', kpts)
# align with reference axis
patches, rand_axis, R = self.axis_align(init_patch, is_aligned_to_global_z, z_axis)
patches = self.normalize(patches, des_r)
# print('patches:', patches)
# by default, we do not apply any SO(2) rotation augmentation
aug_rotation = np.eye(3)[None].repeat(patches.shape[0], axis=0)
aug_rotation = torch.FloatTensor(aug_rotation).to(patches.device)
patches = patches @ aug_rotation.transpose(-1, -2)
rand_axis = (rand_axis.unsqueeze(1) @ aug_rotation.transpose(-1, -2)).squeeze(1)
# spatial point transformer
inv_patches = self.SPT(patches, 1, self.delta / self.rad_n)
# Vanilla SpinNet
new_points = inv_patches.permute(0, 3, 1, 2) # (B, C_in, npoint, nsample+1), input features
x = self.pnt_layer(new_points)
x = F.max_pool2d(x, kernel_size=(1, x.shape[-1])).squeeze(3) # (B, C_in, npoint)
del new_points
x = x.view(x.shape[0], x.shape[1], self.rad_n, self.ele_n, self.azi_n)
x, mid = self.conv_net(x)
w = self.pool_layer(x)
f = F.avg_pool2d(x * w, kernel_size=(x.shape[2], x.shape[3]))
f = F.normalize(f.view(f.shape[0], -1), p=2, dim=1)
x = F.normalize(x, p=2, dim=1)
return {'desc': f,
'equi': x,
'rand_axis': rand_axis,
'R': R,
'patches': patches,
'aug_rotation': aug_rotation}
def select_patches(self, pts, refer_pts, vicinity, patch_sample=1024):
# pts: (B, N, 3)
# refer_pts: (B, K, 3), key points
B, N, C = pts.shape
# shuffle pts if pts is not orderless
index = np.random.choice(N, N, replace=False)
pts = pts[:, index]
# Use PyTorch3D's ball_query instead of pnt2.ball_query and pnt2.grouping_operation
# ball_query returns (dists, idx, nn) where nn contains the neighbor points
dists, group_idx, new_points = ball_query(
p1=refer_pts, # query points (centers)
p2=pts, # points to search in
K=patch_sample, # maximum number of neighbors
radius=vicinity, # radius within which to search
return_nn=True # return the neighbor points directly
)
# dists as the squared distance
# Sort distances in decreasing order for each reference point
# sorted_dists, _ = torch.sort(dists, dim=-1, descending=True)
# print('sorted_dists.shape:', sorted_dists.shape)
# print('sorted_dists:', sorted_dists)
# new_points from ball_query has shape (B, P1, K, D) where D is the point dimension
# We want it in shape (B, P1, K, D) to match the original format
# No permute needed since ball_query already returns the correct format
# Create mask for invalid neighbors (where group_idx == -1)
invalid_mask = (group_idx == -1).float() # 1 where no valid neighbor found
# Expand masks to match coordinate dimensions
invalid_mask = invalid_mask.unsqueeze(3).repeat([1, 1, 1, C])
# Create reference points repeated for each patch position
new_pts = refer_pts.unsqueeze(2).repeat([1, 1, patch_sample, 1])
# Fill invalid positions with reference points, keep valid neighbors as they are
local_patches = new_points * (1 - invalid_mask) + new_pts * invalid_mask
del invalid_mask
del new_points
del group_idx
del new_pts
del pts
return local_patches
def axis_align(self, input, is_aligned_to_global_z, z_axis=None):
center = input[:, -1, :3]
delta_x = input[:, :, :3] - center.unsqueeze(1) # (B, npoint, 3), normalized coordinates
if not is_aligned_to_global_z:
if z_axis is None:
z_axis = cal_Z_axis(delta_x, ref_point=center)
z_axis = l2_norm(z_axis, axis=1)
else:
z_axis = z_axis[0]
R = RodsRotatFormula(z_axis, torch.FloatTensor([0, 0, 1]).expand_as(z_axis))
delta_x = torch.matmul(delta_x, R)
# for calculate gt lable
rand_axis = torch.zeros_like(center)
rand_axis[:, -1] = 1
rand_axis = torch.cross(z_axis, rand_axis)
rand_axis = F.normalize(rand_axis, p=2, dim=-1)
else:
rand_axis = torch.zeros_like(center)
rand_axis[:, 0] = 1
R = torch.eye(3).to(center.device)
R = R[None].repeat([center.shape[0], 1, 1])
return delta_x, rand_axis, R
def SPT(self, delta_x, des_r, voxel_r):
# partition the local surface along elevator, azimuth, radial dimensions
S2_xyz = torch.FloatTensor(get_voxel_coordinate(radius=des_r,
rad_n=self.rad_n,
azi_n=self.azi_n,
ele_n=self.ele_n))
pts_xyz = S2_xyz.view(1, -1, 3).repeat([delta_x.shape[0], 1, 1]).cuda()
# query points in sphere
new_points = sphere_query(delta_x, pts_xyz, radius=voxel_r,
nsample=self.voxel_sample)
# transform rotation-variant coords into rotation-invariant coords
new_points = var_to_invar(new_points, self.rad_n, self.azi_n, self.ele_n)
return new_points
def normalize(self, pts, radius):
delta_x = pts / (torch.ones_like(pts).to(pts.device) * radius)
return delta_x
def get_parameter(self):
return list(self.parameters())
|