"""LitePT backbone for pimm.
This is a pimm-local port of the Pointcept-style LitePT implementation. The
LitePT-specific helpers are kept in this file so the backbone can move as a
single module without importing implementation details from PTv3.
"""
from functools import partial
import torch
import torch.nn as nn
import torch.nn.functional as F
import spconv.pytorch as spconv
import torch_scatter
from addict import Dict
from timm.layers import DropPath
from pointrope import PointROPE
from pimm.models.builder import MODELS
from pimm.models.modules import PointModule, PointSequential
from pimm.models.utils.misc import offset2bincount
from pimm.models.utils.structure import Point
try:
from flash_attn import flash_attn_varlen_qkvpacked_func
except ImportError:
flash_attn_varlen_qkvpacked_func = None
class Embedding(PointModule):
def __init__(
self,
in_channels,
embed_channels,
norm_layer=None,
act_layer=None,
mask_token=False,
):
super().__init__()
self.in_channels = in_channels
self.embed_channels = embed_channels
self.stem = PointSequential(
conv=spconv.SubMConv3d(
in_channels,
embed_channels,
kernel_size=5,
padding=1,
bias=False,
indice_key="stem",
)
)
if norm_layer is not None:
self.stem.add(norm_layer(embed_channels), name="norm")
if act_layer is not None:
self.stem.add(act_layer(), name="act")
if mask_token:
self.mask_token = nn.Parameter(torch.zeros(1, embed_channels))
else:
self.mask_token = None
def forward(self, point: Point):
point = self.stem(point)
if self.mask_token is not None and "mask" in point.keys():
point.feat = torch.where(
point.mask.unsqueeze(-1),
self.mask_token.to(dtype=point.feat.dtype),
point.feat,
)
point.sparse_conv_feat = point.sparse_conv_feat.replace_feature(point.feat)
return point
class GridPooling(PointModule):
def __init__(
self,
in_channels,
out_channels,
stride=2,
norm_layer=None,
act_layer=None,
reduce="max",
shuffle_orders=True,
traceable=True,
re_serialization=False,
serialization_order="z",
):
super().__init__()
self.in_channels = in_channels
self.out_channels = out_channels
self.stride = stride
assert reduce in ["sum", "mean", "min", "max"]
self.reduce = reduce
self.shuffle_orders = shuffle_orders
self.traceable = traceable
self.re_serialization = re_serialization
self.serialization_order = serialization_order
self.proj = nn.Linear(in_channels, out_channels)
self.norm = PointSequential(norm_layer(out_channels)) if norm_layer else None
self.act = PointSequential(act_layer()) if act_layer else None
def forward(self, point: Point):
if "grid_coord" in point.keys():
grid_coord = point.grid_coord
elif {"coord", "grid_size"}.issubset(point.keys()):
grid_coord = torch.div(
point.coord - point.coord.min(0)[0],
point.grid_size,
rounding_mode="trunc",
).int()
else:
raise AssertionError(
"[grid_coord] or [coord, grid_size] should be included in the Point"
)
grid_coord = torch.div(grid_coord, self.stride, rounding_mode="trunc")
grid_coord = grid_coord | point.batch.view(-1, 1) << 48
grid_coord, cluster, counts = torch.unique(
grid_coord,
sorted=True,
return_inverse=True,
return_counts=True,
dim=0,
)
grid_coord = grid_coord & ((1 << 48) - 1)
_, indices = torch.sort(cluster)
idx_ptr = torch.cat([counts.new_zeros(1), torch.cumsum(counts, dim=0)])
head_indices = indices[idx_ptr[:-1]]
point_dict = Dict(
feat=torch_scatter.segment_csr(
self.proj(point.feat)[indices], idx_ptr, reduce=self.reduce
),
coord=torch_scatter.segment_csr(
point.coord[indices], idx_ptr, reduce="mean"
),
grid_coord=grid_coord,
batch=point.batch[head_indices],
)
if "origin_coord" in point.keys():
point_dict["origin_coord"] = torch_scatter.segment_csr(
point.origin_coord[indices], idx_ptr, reduce="mean"
)
if "condition" in point.keys():
point_dict["condition"] = point.condition
if "context" in point.keys():
point_dict["context"] = point.context
if "name" in point.keys():
point_dict["name"] = point.name
if "split" in point.keys():
point_dict["split"] = point.split
if "color" in point.keys():
point_dict["color"] = torch_scatter.segment_csr(
point.color[indices], idx_ptr, reduce="mean"
)
if "segment_motif" in point.keys():
point_dict["segment_motif"] = point.segment_motif[head_indices]
if "grid_size" in point.keys():
point_dict["grid_size"] = point.grid_size * self.stride
if "mask" in point.keys():
point_dict["mask"] = (
torch_scatter.segment_csr(
point.mask[indices].float(), idx_ptr, reduce="mean"
)
> 0.5
)
if self.traceable:
point_dict["pooling_inverse"] = cluster
point_dict["pooling_parent"] = point
point = Point(point_dict)
if self.norm is not None:
point = self.norm(point)
if self.act is not None:
point = self.act(point)
if self.re_serialization:
point.serialization(
order=self.serialization_order, shuffle_orders=self.shuffle_orders
)
point.sparsify()
return point
class GridUnpooling(PointModule):
def __init__(
self,
in_channels,
skip_channels,
out_channels,
norm_layer=None,
act_layer=None,
traceable=False,
):
super().__init__()
self.proj = PointSequential(nn.Linear(in_channels, out_channels))
self.proj_skip = PointSequential(nn.Linear(skip_channels, out_channels))
if norm_layer is not None:
self.proj.add(norm_layer(out_channels))
self.proj_skip.add(norm_layer(out_channels))
if act_layer is not None:
self.proj.add(act_layer())
self.proj_skip.add(act_layer())
self.traceable = traceable
def forward(self, point):
assert "pooling_parent" in point.keys()
assert "pooling_inverse" in point.keys()
parent = point.pop("pooling_parent")
inverse = point.pooling_inverse
feat = point.feat
parent = self.proj_skip(parent)
parent.feat = parent.feat + self.proj(point).feat[inverse]
parent.sparse_conv_feat = parent.sparse_conv_feat.replace_feature(parent.feat)
if self.traceable:
point.feat = feat
parent["unpooling_parent"] = point
parent["unpooling_inverse"] = inverse
return parent
class PointROPEAttention(PointModule):
def __init__(
self,
channels,
num_heads,
patch_size,
rope_freq,
qkv_bias=True,
qk_scale=None,
attn_drop=0.0,
proj_drop=0.0,
order_index=0,
):
super().__init__()
assert channels % num_heads == 0
self.channels = channels
self.num_heads = num_heads
self.scale = qk_scale or (channels // num_heads) ** -0.5
self.order_index = order_index
self.patch_size = patch_size
self.attn_drop = attn_drop
self.qkv = nn.Linear(channels, channels * 3, bias=qkv_bias)
self.proj = nn.Linear(channels, channels)
self.proj_drop = nn.Dropout(proj_drop)
self.rope = PointROPE(freq=rope_freq)
@torch.no_grad()
def get_padding_and_inverse(self, point):
pad_key = "pad"
unpad_key = "unpad"
cu_seqlens_key = "cu_seqlens_key"
if (
pad_key not in point.keys()
or unpad_key not in point.keys()
or cu_seqlens_key not in point.keys()
):
if self.patch_size == -1:
cu_seqlens = torch.cat([point.offset.new_zeros(1), point.offset]).int()
point[pad_key] = None
point[unpad_key] = None
point[cu_seqlens_key] = cu_seqlens
return point[pad_key], point[unpad_key], point[cu_seqlens_key]
offset = point.offset
bincount = offset2bincount(offset)
patch_size = self.patch_size
device = offset.device
bincount_pad = torch.div(
bincount + patch_size - 1, patch_size, rounding_mode="trunc"
) * patch_size
mask_pad = bincount > patch_size
bincount_pad = ~mask_pad * bincount + mask_pad * bincount_pad
offset_ = F.pad(offset, (1, 0))
offset_pad = F.pad(torch.cumsum(bincount_pad, dim=0), (1, 0))
n_total = offset_[-1]
n_pad_total = offset_pad[-1]
shift = offset_pad[:-1] - offset_[:-1]
idx_unpad = torch.arange(n_total, device=device)
unpad = idx_unpad + shift[torch.searchsorted(offset, idx_unpad, right=True)]
pad = torch.arange(n_pad_total, device=device)
remainder = bincount % patch_size
needs_pad = mask_pad & (remainder != 0)
if needs_pad.any():
pad_views = torch.where(needs_pad)[0]
r = remainder[pad_views]
copy_lens = patch_size - r
total_copies = copy_lens.sum()
view_of_copy = torch.arange(
len(pad_views), device=device
).repeat_interleave(copy_lens)
local_idx = torch.arange(total_copies, device=device)
copy_cumsum = F.pad(torch.cumsum(copy_lens, dim=0), (1, 0))
local_idx = local_idx - copy_cumsum[view_of_copy]
dst = (
offset_pad[pad_views[view_of_copy] + 1]
- patch_size
+ r[view_of_copy]
+ local_idx
)
pad[dst] = pad[dst - patch_size]
offset_pad_cumsum = torch.cumsum(bincount_pad, dim=0)
idx_pad = torch.arange(n_pad_total, device=device)
pad = pad - shift[
torch.searchsorted(offset_pad_cumsum, idx_pad, right=True)
]
patches_per_view = torch.div(
bincount_pad + patch_size - 1, patch_size, rounding_mode="trunc"
).int()
total_patches = patches_per_view.sum()
patches_cumsum = F.pad(torch.cumsum(patches_per_view, dim=0), (1, 0))
patch_idx = torch.arange(total_patches, device=device, dtype=torch.int32)
patch_view = torch.searchsorted(
patches_cumsum[1:], patch_idx, right=True
).int()
patch_local = patch_idx - patches_cumsum[patch_view].int()
cu_seqlens_vals = offset_pad[patch_view].int() + patch_local * patch_size
point[pad_key] = pad
point[unpad_key] = unpad
point[cu_seqlens_key] = F.pad(
cu_seqlens_vals.int(), (0, 1), value=int(n_pad_total)
)
return point[pad_key], point[unpad_key], point[cu_seqlens_key]
def forward(self, point):
heads = self.num_heads
channels = self.channels
pad, unpad, cu_seqlens = self.get_padding_and_inverse(point)
max_seqlen = int(cu_seqlens[-1]) if self.patch_size == -1 else self.patch_size
order = point.serialized_order[self.order_index]
if pad is not None:
order = order[pad]
inverse = unpad[point.serialized_inverse[self.order_index]]
qkv = self.qkv(point.feat)[order]
pos = point.grid_coord[order].reshape(-1, 3).unsqueeze(0)
q, k, v = qkv.half().chunk(3, dim=-1)
q = q.reshape(-1, heads, channels // heads).transpose(0, 1)[None]
k = k.reshape(-1, heads, channels // heads).transpose(0, 1)[None]
q = self.rope(q.float(), pos).to(q.dtype)
k = self.rope(k.float(), pos).to(k.dtype)
qkv_rotated = torch.stack(
[
q.squeeze(0).transpose(0, 1),
k.squeeze(0).transpose(0, 1),
v.reshape(-1, heads, channels // heads),
],
dim=1,
)
feat = flash_attn_varlen_qkvpacked_func(
qkv_rotated,
cu_seqlens,
max_seqlen=max_seqlen,
dropout_p=self.attn_drop if self.training else 0,
softmax_scale=self.scale,
).reshape(-1, channels)
feat = feat.to(qkv.dtype)
if pad is not None:
feat = feat[inverse]
point.feat = self.proj_drop(self.proj(feat))
return point
class MLP(nn.Module):
def __init__(
self,
in_channels,
hidden_channels=None,
out_channels=None,
act_layer=nn.GELU,
drop=0.0,
):
super().__init__()
out_channels = out_channels or in_channels
hidden_channels = hidden_channels or in_channels
self.fc1 = nn.Linear(in_channels, hidden_channels)
self.act = act_layer()
self.fc2 = nn.Linear(hidden_channels, out_channels)
self.drop = nn.Dropout(drop)
def forward(self, x):
x = self.fc1(x)
x = self.act(x)
x = self.drop(x)
x = self.fc2(x)
x = self.drop(x)
return x
class Block(PointModule):
def __init__(
self,
channels,
num_heads,
patch_size=48,
mlp_ratio=4.0,
qkv_bias=True,
qk_scale=None,
attn_drop=0.0,
proj_drop=0.0,
drop_path=0.0,
norm_layer=nn.LayerNorm,
act_layer=nn.GELU,
pre_norm=True,
order_index=0,
cpe_indice_key=None,
enable_conv=True,
enable_attn=True,
rope_freq=100.0,
):
super().__init__()
self.channels = channels
self.pre_norm = pre_norm
self.enable_conv = enable_conv
self.enable_attn = enable_attn
if self.enable_conv:
self.conv = PointSequential(
spconv.SubMConv3d(
channels,
channels,
kernel_size=3,
bias=True,
indice_key=cpe_indice_key,
),
nn.Linear(channels, channels),
norm_layer(channels),
)
else:
self.norm0 = PointSequential(norm_layer(channels))
if self.enable_attn:
self.norm1 = PointSequential(norm_layer(channels))
self.attn = PointROPEAttention(
channels=channels,
patch_size=patch_size,
rope_freq=rope_freq,
num_heads=num_heads,
qkv_bias=qkv_bias,
qk_scale=qk_scale,
attn_drop=attn_drop,
proj_drop=proj_drop,
order_index=order_index,
)
self.norm2 = PointSequential(norm_layer(channels))
self.mlp = PointSequential(
MLP(
in_channels=channels,
hidden_channels=int(channels * mlp_ratio),
out_channels=channels,
act_layer=act_layer,
drop=proj_drop,
)
)
self.drop_path = PointSequential(
DropPath(drop_path) if drop_path > 0.0 else nn.Identity()
)
def forward(self, point: Point):
if self.enable_conv:
shortcut = point.feat
point = self.conv(point)
point.feat = shortcut + point.feat
else:
point = self.norm0(point)
if self.enable_attn:
shortcut = point.feat
if self.pre_norm:
point = self.norm1(point)
point = self.drop_path(self.attn(point))
point.feat = shortcut + point.feat
if not self.pre_norm:
point = self.norm1(point)
shortcut = point.feat
if self.pre_norm:
point = self.norm2(point)
point = self.drop_path(self.mlp(point))
point.feat = shortcut + point.feat
if not self.pre_norm:
point = self.norm2(point)
point.sparse_conv_feat = point.sparse_conv_feat.replace_feature(point.feat)
return point
[docs]
@MODELS.register_module("LitePT")
class LitePT(PointModule):
def __init__(
self,
in_channels=4,
order=("z", "z-trans", "hilbert", "hilbert-trans"),
stride=(2, 2, 2, 2),
enc_depths=(2, 2, 2, 6, 2),
enc_channels=(36, 72, 144, 252, 504),
enc_num_head=(2, 4, 8, 14, 28),
enc_patch_size=(1024, 1024, 1024, 1024, 1024),
enc_conv=(True, True, True, False, False),
enc_attn=(False, False, False, True, True),
enc_rope_freq=(100.0, 100.0, 100.0, 100.0, 100.0),
dec_depths=(0, 0, 0, 0),
dec_channels=(72, 72, 144, 252),
dec_num_head=(4, 4, 8, 14),
dec_patch_size=(1024, 1024, 1024, 1024),
dec_conv=(False, False, False, False),
dec_attn=(False, False, False, False),
dec_rope_freq=(100.0, 100.0, 100.0, 100.0),
mlp_ratio=4,
qkv_bias=True,
qk_scale=None,
attn_drop=0.0,
proj_drop=0.0,
drop_path=0.3,
pre_norm=True,
shuffle_orders=True,
mask_token=False,
enc_mode=False,
freeze_encoder=False,
traceable=True,
):
super().__init__()
self.num_stages = len(enc_depths)
self.order = [order] if isinstance(order, str) else order
self.enc_mode = enc_mode
self.freeze_encoder = freeze_encoder
self.shuffle_orders = shuffle_orders
self.enc_conv = enc_conv
self.enc_attn = enc_attn
self.dec_conv = dec_conv
self.dec_attn = dec_attn
assert self.num_stages == len(stride) + 1
assert self.num_stages == len(enc_depths)
assert self.num_stages == len(enc_channels)
assert self.num_stages == len(enc_num_head)
assert self.num_stages == len(enc_patch_size)
assert self.enc_mode or self.num_stages == len(dec_depths) + 1
assert self.enc_mode or self.num_stages == len(dec_channels) + 1
assert self.enc_mode or self.num_stages == len(dec_num_head) + 1
assert self.enc_mode or self.num_stages == len(dec_patch_size) + 1
bn_layer = partial(nn.BatchNorm1d, eps=1e-3, momentum=0.01)
ln_layer = nn.LayerNorm
act_layer = nn.GELU
self.embedding = Embedding(
in_channels=in_channels,
embed_channels=enc_channels[0],
norm_layer=bn_layer,
act_layer=act_layer,
mask_token=mask_token,
)
enc_drop_path = [
x.item() for x in torch.linspace(0, drop_path, sum(enc_depths))
]
self.enc = PointSequential()
for s in range(self.num_stages):
enc_drop_path_ = enc_drop_path[
sum(enc_depths[:s]) : sum(enc_depths[: s + 1])
]
enc = PointSequential()
if s > 0:
enc.add(
GridPooling(
in_channels=enc_channels[s - 1],
out_channels=enc_channels[s],
stride=stride[s - 1],
norm_layer=bn_layer,
act_layer=act_layer,
traceable=traceable,
re_serialization=enc_attn[s],
serialization_order=self.order,
shuffle_orders=self.shuffle_orders,
),
name="down",
)
for i in range(enc_depths[s]):
enc.add(
Block(
channels=enc_channels[s],
num_heads=enc_num_head[s],
patch_size=enc_patch_size[s],
mlp_ratio=mlp_ratio,
qkv_bias=qkv_bias,
qk_scale=qk_scale,
attn_drop=attn_drop,
proj_drop=proj_drop,
drop_path=enc_drop_path_[i],
norm_layer=ln_layer,
act_layer=act_layer,
pre_norm=pre_norm,
order_index=i % len(self.order),
cpe_indice_key=f"stage{s}",
enable_conv=enc_conv[s],
enable_attn=enc_attn[s],
rope_freq=enc_rope_freq[s],
),
name=f"block{i}",
)
if len(enc) != 0:
self.enc.add(module=enc, name=f"enc{s}")
if not self.enc_mode:
dec_drop_path = [
x.item() for x in torch.linspace(0, drop_path, sum(dec_depths))
]
self.dec = PointSequential()
dec_channels = list(dec_channels) + [enc_channels[-1]]
for s in reversed(range(self.num_stages - 1)):
dec_drop_path_ = dec_drop_path[
sum(dec_depths[:s]) : sum(dec_depths[: s + 1])
]
dec_drop_path_.reverse()
dec = PointSequential()
dec.add(
GridUnpooling(
in_channels=dec_channels[s + 1],
skip_channels=enc_channels[s],
out_channels=dec_channels[s],
norm_layer=bn_layer,
act_layer=act_layer,
traceable=traceable,
),
name="up",
)
for i in range(dec_depths[s]):
dec.add(
Block(
channels=dec_channels[s],
num_heads=dec_num_head[s],
patch_size=dec_patch_size[s],
mlp_ratio=mlp_ratio,
qkv_bias=qkv_bias,
qk_scale=qk_scale,
attn_drop=attn_drop,
proj_drop=proj_drop,
drop_path=dec_drop_path_[i],
norm_layer=ln_layer,
act_layer=act_layer,
pre_norm=pre_norm,
order_index=i % len(self.order),
cpe_indice_key=f"stage{s}",
enable_conv=dec_conv[s],
enable_attn=dec_attn[s],
rope_freq=dec_rope_freq[s],
),
name=f"block{i}",
)
self.dec.add(module=dec, name=f"dec{s}")
if self.freeze_encoder:
for p in self.embedding.parameters():
p.requires_grad = False
for p in self.enc.parameters():
p.requires_grad = False
[docs]
def forward(self, data_dict):
point = Point(data_dict)
if self.enc_attn[0]:
point.serialization(order=self.order, shuffle_orders=self.shuffle_orders)
point.sparsify()
point = self.embedding(point)
point = self.enc(point)
if not self.enc_mode:
point = self.dec(point)
return point