PointRSP / model /common.py
Mo-nan's picture
Upload 16 files
913ec88 verified
Raw
History Blame Contribute Delete
7.11 kB
import warnings
warnings.filterwarnings("ignore", category=FutureWarning)
import functools
from collections import OrderedDict
import spconv.pytorch as spconv
import torch
from spconv.pytorch.modules import SparseModule
from torch import nn
import torch.nn.functional as F
import torch
import torch.nn as nn
class CappedLayerNorm(nn.Module):
"""
LayerNorm with:
- capped scale: gamma ∈ (0, max_scale)
- capped bias: beta ∈ (-max_bias, max_bias)
- stable initialization: scale ≈ max_scale * 0.5 at start
"""
def __init__(self, dim, eps=1e-5, max_scale=1.8, max_bias=0.5):
super().__init__()
self.dim = dim
self.eps = eps
self.max_scale = max_scale
self.max_bias = max_bias
# raw learnable parameters
# g_raw = 0 → gamma = max_scale * sigmoid(0) = max_scale * 0.5
self.g_raw = nn.Parameter(torch.zeros(dim))
# b_raw = 0 → beta = tanh(0) * max_bias = 0
self.b_raw = nn.Parameter(torch.zeros(dim))
def forward(self, x):
# LayerNorm normalization
mean = x.mean(-1, keepdim=True)
var = x.var(-1, unbiased=False, keepdim=True)
x_norm = (x - mean) / torch.sqrt(var + self.eps)
# capped scale: gamma ∈ (0, max_scale)
gamma = torch.sigmoid(self.g_raw) * self.max_scale
# capped bias: beta ∈ (-max_bias, max_bias)
beta = torch.tanh(self.b_raw) * self.max_bias
return gamma * x_norm + beta
class ResidualConv(nn.Module):
def __init__(self, in_channels, out_channels):
super().__init__()
self.resi_ratio = 0.5
self.norm1 = nn.LayerNorm(in_channels)
self.linear1 = nn.Linear(in_channels, out_channels)
self.norm2 = nn.LayerNorm(out_channels)
self.linear2 = nn.Linear(out_channels, out_channels)
self.act = nn.ReLU()
def forward(self, x):
h = self.norm1(x)
h = self.linear1(h)
h = self.act(h)
h = self.norm2(h)
h = self.linear2(h)
return x + h * self.resi_ratio
class MLP(nn.Sequential):
def __init__(self, in_channels, out_channels, norm_fn=None, num_layers=2):
modules = []
for _ in range(num_layers - 1):
modules.append(nn.Linear(in_channels, in_channels))
if norm_fn:
modules.append(norm_fn(in_channels))
modules.append(nn.ReLU())
modules.append(nn.Linear(in_channels, out_channels))
return super().__init__(*modules)
def init_weights(self):
for m in self.modules():
if isinstance(m, nn.Linear):
nn.init.xavier_uniform_(m.weight)
nn.init.constant_(m.bias, 0)
nn.init.normal_(self[-1].weight, 0, 0.01)
nn.init.constant_(self[-1].bias, 0)
# current 1x1 conv in spconv2x has a bug. It will be removed after the bug is fixed
class Custom1x1Subm3d(spconv.SparseConv3d):
def forward(self, input):
features = torch.mm(input.features, self.weight.view(self.out_channels, self.in_channels).T)
if self.bias is not None:
features += self.bias
out_tensor = spconv.SparseConvTensor(features, input.indices, input.spatial_shape,
input.batch_size)
out_tensor.indice_dict = input.indice_dict
out_tensor.grid = input.grid
return out_tensor
class ResidualBlock(SparseModule):
def __init__(self, in_channels, out_channels, norm_fn, indice_key=None):
super().__init__()
if in_channels == out_channels:
self.i_branch = spconv.SparseSequential(nn.Identity())
else:
self.i_branch = spconv.SparseSequential(
Custom1x1Subm3d(in_channels, out_channels, kernel_size=1, bias=False))
self.conv_branch = spconv.SparseSequential(
norm_fn(in_channels), nn.ReLU(),
spconv.SubMConv3d(
in_channels,
out_channels,
kernel_size=3,
padding=1,
bias=False,
indice_key=indice_key), norm_fn(out_channels), nn.ReLU(),
spconv.SubMConv3d(
out_channels,
out_channels,
kernel_size=3,
padding=1,
bias=False,
indice_key=indice_key))
def forward(self, input):
identity = spconv.SparseConvTensor(input.features, input.indices, input.spatial_shape,
input.batch_size)
output = self.conv_branch(input)
out_feats = output.features + self.i_branch(identity).features
output = output.replace_feature(out_feats)
return output
class UBlock(nn.Module):
def __init__(self, nPlanes, norm_fn, block_reps, block, indice_key_id=1):
super().__init__()
self.nPlanes = nPlanes
blocks = {
'block{}'.format(i):
block(nPlanes[0], nPlanes[0], norm_fn, indice_key='subm{}'.format(indice_key_id))
for i in range(block_reps)
}
blocks = OrderedDict(blocks)
self.blocks = spconv.SparseSequential(blocks)
if len(nPlanes) > 1:
self.conv = spconv.SparseSequential(
norm_fn(nPlanes[0]), nn.ReLU(),
spconv.SparseConv3d(
nPlanes[0],
nPlanes[1],
kernel_size=2,
stride=2,
bias=False,
indice_key='spconv{}'.format(indice_key_id)))
self.u = UBlock(
nPlanes[1:], norm_fn, block_reps, block, indice_key_id=indice_key_id + 1)
self.deconv = spconv.SparseSequential(
norm_fn(nPlanes[1]), nn.ReLU(),
spconv.SparseInverseConv3d(
nPlanes[1],
nPlanes[0],
kernel_size=2,
bias=False,
indice_key='spconv{}'.format(indice_key_id)))
blocks_tail = {}
for i in range(block_reps):
blocks_tail['block{}'.format(i)] = block(
nPlanes[0] * (2 - i),
nPlanes[0],
norm_fn,
indice_key='subm{}'.format(indice_key_id))
blocks_tail = OrderedDict(blocks_tail)
self.blocks_tail = spconv.SparseSequential(blocks_tail)
def forward(self, input):
output = self.blocks(input)
identity = spconv.SparseConvTensor(output.features, output.indices, output.spatial_shape,
output.batch_size)
if len(self.nPlanes) > 1:
output_decoder = self.conv(output)
output_decoder = self.u(output_decoder)
output_decoder = self.deconv(output_decoder)
out_feats = torch.cat((identity.features, output_decoder.features), dim=1)
output = output.replace_feature(out_feats)
output = self.blocks_tail(output)
return output