"""AQ3D — Adaptive Query Transformer for 3D Instance Segmentation. Vendored, dependency-trimmed port of the official implementation (https://github.com/kenomo/aq3d, MIT, Keno Moenck & Thorsten Schuppstuhl) so it runs on ZeroGPU. The Volt-B backbone comes from https://github.com/YilmazKadir/Volt. The only deviations from upstream are the compiled-extension replacements in ``nnutils`` (torch_scatter / torch_geometric.fps / flash_attn -> plain PyTorch); module names, layer order and every hyper-parameter follow ``configs/model/aqtd_volt_scannet200.yaml`` exactly so the released checkpoint loads with ``strict=True``. """ import math from functools import partial import torch import torch.nn as nn import torch.nn.functional as F from nnutils import fps, scatter_mean, scatter_softmax, scatter_sum, varlen_qkvpacked_attention # =========================================================================== # # src/models/components/attn.py # =========================================================================== # class MultiHeadAttention(nn.Module): def __init__(self, embed_dim=256, v_dim=None, num_heads=8, dropout=0.0, q_proj=True, k_proj=True, v_proj=True): super().__init__() self.num_heads = num_heads self.embed_dim = embed_dim self.v_dim = v_dim if v_dim is not None else embed_dim self.head_dim = embed_dim // num_heads self.v_head_dim = self.v_dim // num_heads assert self.head_dim * num_heads == embed_dim assert self.v_head_dim * num_heads == self.v_dim self.q_proj, self.k_proj, self.v_proj = q_proj, k_proj, v_proj if q_proj: self.q_proj_weight = nn.Parameter(torch.empty(embed_dim, embed_dim)) self.q_proj_bias = nn.Parameter(torch.empty(embed_dim)) if k_proj: self.k_proj_weight = nn.Parameter(torch.empty(embed_dim, embed_dim)) self.k_proj_bias = nn.Parameter(torch.empty(embed_dim)) if v_proj: self.v_proj_weight = nn.Parameter(torch.empty(self.v_dim, self.v_dim)) self.v_proj_bias = nn.Parameter(torch.empty(self.v_dim)) self.out_proj = nn.Linear(self.v_dim, self.v_dim, bias=True) self.dropout = nn.Dropout(dropout) def forward(self, query, key, value, key_padding_mask=None, attn_mask=None): B, q_len, _ = query.shape _, k_len, _ = key.shape v_len = k_len q = query.transpose(0, 1) k = key.transpose(0, 1) v = value.transpose(0, 1) if key_padding_mask is None: key_padding_mask = torch.zeros((B, k_len), dtype=torch.bool, device=q.device) if self.q_proj: q = F.linear(q, self.q_proj_weight, self.q_proj_bias) if self.k_proj: k = F.linear(k, self.k_proj_weight, self.k_proj_bias) if self.v_proj: v = F.linear(v, self.v_proj_weight, self.v_proj_bias) key_padding_mask = key_padding_mask.unsqueeze(1).repeat_interleave(q_len, dim=1) if attn_mask is None: attn_mask = key_padding_mask else: attn_mask = attn_mask.logical_or(key_padding_mask) attn_mask = attn_mask.repeat_interleave(self.num_heads, dim=0) attn_mask_float = torch.zeros_like(attn_mask, dtype=q.dtype, device=q.device) attn_mask_float = attn_mask_float.masked_fill(attn_mask, float("-inf")) q_sdpa = q.transpose(0, 1).view(B, q_len, self.num_heads, self.head_dim).transpose(1, 2) k_sdpa = k.transpose(0, 1).view(B, k_len, self.num_heads, self.head_dim).transpose(1, 2) v_sdpa = v.transpose(0, 1).view(B, v_len, self.num_heads, self.v_head_dim).transpose(1, 2) attn_mask_sdpa = attn_mask_float.view(B, self.num_heads, q_len, k_len) out = F.scaled_dot_product_attention(q_sdpa, k_sdpa, v_sdpa, attn_mask=attn_mask_sdpa, is_causal=False) out = out.transpose(1, 2).reshape(B, q_len, self.v_dim) return self.out_proj(out), None # =========================================================================== # # src/models/components/modules.py # =========================================================================== # class RoPE(nn.Module): """Axial rotary positional embedding over metric 3-D coordinates.""" def __init__(self, theta=100.0, head_split=(12, 12, 8), grid_size=0.1, max_grid_size=(1024, 1024, 512)): super().__init__() freqs = [1.0 / theta ** torch.linspace(0, 1, head_split[i] // 2) for i in range(3)] self.grid_size = grid_size self.head_split = head_split self.max_grid_size = max_grid_size for name, f, m in zip("xyz", freqs, max_grid_size): self.register_buffer(f"cis_cache_{name}", self._precompute(f, m), persistent=False) @staticmethod def _precompute(freqs, max_pos): freqs_pos = torch.outer(torch.arange(max_pos).float(), freqs) return torch.polar(torch.ones_like(freqs_pos), freqs_pos) def forward(self, x, coords): indices = torch.div(coords, self.grid_size, rounding_mode="floor").long() indices = indices.clamp(min=0) # upstream asserts here; clamping keeps out-of-domain (very large) scenes # running instead of hard-crashing the demo for a in range(3): indices[..., a] = indices[..., a].clamp(max=self.max_grid_size[a] - 1) cis = torch.cat([self.cis_cache_x[indices[..., 0]], self.cis_cache_y[indices[..., 1]], self.cis_cache_z[indices[..., 2]]], dim=-1).unsqueeze(2) x_ = torch.view_as_complex(x.float().reshape(*x.shape[:-1], -1, 2)) return torch.view_as_real(x_ * cis).flatten(-2).to(x.dtype) class CosineClassifier(nn.Module): def __init__(self, in_features, out_features, scale=20.0): super().__init__() self.weight = nn.Parameter(torch.Tensor(out_features, in_features)) self.scale = scale nn.init.xavier_uniform_(self.weight) def forward(self, x): return F.linear(F.normalize(x, p=2, dim=-1), F.normalize(self.weight, p=2, dim=-1)) * self.scale class FFN(nn.Module): def __init__(self, d_model=256, output_dim=None, hidden_dim=1024, dropout=0.0, activation_fn=nn.GELU, use_residual=True, use_norm=True, num_layers=2): super().__init__() self.num_layers = num_layers output_dim = output_dim or d_model h = [hidden_dim] * (num_layers - 1) self.layers = nn.ModuleList(nn.Linear(n, k) for n, k in zip([d_model] + h, h + [output_dim])) self.use_residual = use_residual if use_residual: self.fast_path = nn.Linear(d_model, output_dim) if d_model != output_dim else nn.Identity() self.use_norm = use_norm self.activation_fn = activation_fn() self.norm = nn.LayerNorm(output_dim) self.dropout = nn.Dropout(dropout) def forward(self, x): input_x = x for i, layer in enumerate(self.layers): x = layer(x) if i < self.num_layers - 1: x = self.dropout(self.activation_fn(x)) x = self.dropout(x) if self.use_residual: x = x + self.fast_path(input_x) if self.use_norm: x = self.norm(x) return x # =========================================================================== # # src/models/components/aqtd/modules.py # =========================================================================== # class SelfAttentionLayer(nn.Module): def __init__(self, d_model=256, nhead=8, dropout=0.0, rope=None): super().__init__() self.qc_in_proj = nn.Linear(d_model, d_model) self.kc_in_proj = nn.Linear(d_model, d_model) self.attn = MultiHeadAttention(embed_dim=d_model, v_dim=d_model, num_heads=nhead, dropout=dropout, q_proj=False, k_proj=False) self.norm = nn.LayerNorm(d_model) self.dropout = nn.Dropout(dropout) self.nhead = nhead self.head_dim = d_model // nhead self.rope = rope def forward(self, q_c, q_coords, B, key_padding_mask=None, scene_ranges_min=None, **kwargs): B = q_c.shape[0] tgt_len = src_len = q_c.shape[1] coords = q_coords - scene_ranges_min qc = self.qc_in_proj(q_c).view(B, tgt_len, self.nhead, self.head_dim) kc = self.kc_in_proj(q_c).view(B, src_len, self.nhead, self.head_dim) q = (self.rope(qc, coords) if self.rope is not None else qc).flatten(2) k = (self.rope(kc, coords) if self.rope is not None else kc).flatten(2) out, _ = self.attn(q, k, q_c, key_padding_mask=key_padding_mask) return self.norm(self.dropout(out) + q_c) class CrossAttentionLayer(nn.Module): def __init__(self, d_model=256, nhead=8, dropout=0.0, attn_mask_thres=0.1, with_query_pos=False, rope=None): super().__init__() self.qc_in_proj = nn.Linear(d_model, d_model) self.kc_in_proj = nn.Linear(d_model, d_model) self.with_query_pos = with_query_pos if with_query_pos: self.qp_in_proj = nn.Linear(d_model, d_model) self.kp_in_proj = nn.Linear(d_model, d_model) self.attn = MultiHeadAttention(embed_dim=d_model * 2 if with_query_pos else d_model, v_dim=d_model, num_heads=nhead, dropout=dropout, q_proj=False, k_proj=False) self.nhead = nhead self.head_dim = d_model // nhead self.norm = nn.LayerNorm(d_model) self.dropout = nn.Dropout(dropout) self.attn_mask_thres = attn_mask_thres self.rope = rope def forward(self, q_c, q_p, k_c, k_p, v, key_padding_mask, q_coords, kv_coords, pred_masks, B, scene_ranges_min=None, **kwargs): device = q_c.device src_len = k_c.shape[1] tgt_len = q_c.shape[1] if pred_masks is not None: attn_mask = torch.ones((B, tgt_len, src_len), dtype=torch.bool, device=device) for i in range(B): inv = (pred_masks[i].sigmoid() < self.attn_mask_thres).bool() inv[torch.where(inv.sum(-1) == inv.shape[-1])] = False attn_mask[i, :inv.shape[0], :inv.shape[1]] = inv else: attn_mask = None q_coords_ = q_coords - scene_ranges_min kv_coords_ = kv_coords - scene_ranges_min qc = self.qc_in_proj(q_c).view(B, tgt_len, self.nhead, self.head_dim) qc_r = self.rope(qc, q_coords_) if self.rope is not None else qc if self.with_query_pos: qp = self.qp_in_proj(q_p).view(B, tgt_len, self.nhead, self.head_dim) q = torch.cat((qc_r, qp), dim=-1).flatten(2) else: q = qc_r.flatten(2) kc = self.kc_in_proj(k_c).view(B, src_len, self.nhead, self.head_dim) kc_r = self.rope(kc, kv_coords_) if self.rope is not None else kc if self.with_query_pos: kp = self.kp_in_proj(k_p).view(B, src_len, self.nhead, self.head_dim) k = torch.cat((kc_r, kp), dim=-1).flatten(2) else: k = kc_r.flatten(2) out, _ = self.attn(q, k, v, key_padding_mask=key_padding_mask, attn_mask=attn_mask) return self.norm(self.dropout(out) + q_c) # =========================================================================== # # src/models/components/aqtd/query_decoder.py # =========================================================================== # class QueryDecoder(nn.Module): def __init__(self, num_layer=6, num_query=100, num_query_ratio=0.6, max_query=True, num_class=198, in_channel=128, d_model=384, dropout_head=0.0, dropout_layer=0.0, query_init="feat", query_pos_init="adaptive", cosine_classifier=True, refinement_cross_attention=True, refinement_cross_attention_layer_indices=(1, 3, 5), refinement_cross_attention_layer=None, detach_query_pos=True, cross_attention_layer=None, self_attention_layer=None, ffn_layer=None, activation_fn=nn.ReLU): super().__init__() self.num_layer = num_layer self.d_model = d_model self.num_query = num_query self.num_query_ratio = num_query_ratio self.max_query = max_query self.dropout_layer = torch.linspace(0, dropout_layer, num_layer).tolist()[::-1] self.refinement_cross_attention = refinement_cross_attention self.detach_query_pos = detach_query_pos self.query_init = query_init self.query_pos_init = query_pos_init if query_init == "feat": self.query_feat_proj = nn.Sequential(nn.Linear(in_channel, d_model), nn.LayerNorm(d_model), activation_fn()) self.feat_proj = nn.Sequential(nn.Linear(in_channel, d_model), nn.LayerNorm(d_model), activation_fn()) self.mask_proj = nn.Sequential(nn.Linear(in_channel, d_model), activation_fn(), nn.Linear(d_model, d_model)) self.cross_attn_layers = nn.ModuleList() self.self_attn_layers = nn.ModuleList() self.ffn_layers = nn.ModuleList() for _ in range(num_layer): self.self_attn_layers.append(self_attention_layer(d_model=d_model)) self.cross_attn_layers.append(cross_attention_layer(d_model=d_model)) self.ffn_layers.append(ffn_layer(d_model=d_model)) self.refinement_cross_attention_layer_indices = list(refinement_cross_attention_layer_indices) if refinement_cross_attention: self.refinement_cross_attn_layers = nn.ModuleList() self.refinement_ffn_layers = nn.ModuleList() for _ in self.refinement_cross_attention_layer_indices: self.refinement_cross_attn_layers.append(refinement_cross_attention_layer(d_model=d_model)) self.refinement_ffn_layers.append(ffn_layer(d_model=d_model)) self.abs_pos_encoder = None self.abs_pos_encoder_proj = None self.query_pos_delta_head = nn.Sequential( nn.Linear(d_model, d_model), activation_fn(), nn.Linear(d_model, d_model), activation_fn(), nn.Dropout(dropout_head), nn.Linear(d_model, 3)) self.out_norm = nn.LayerNorm(d_model) self.out_cls = nn.Sequential( nn.Linear(d_model, d_model), activation_fn(), nn.Dropout(dropout_head), CosineClassifier(d_model, num_class + 1) if cosine_classifier else nn.Linear(d_model, num_class + 1)) self.out_score = nn.Sequential( nn.Linear(d_model, d_model), activation_fn(), nn.Dropout(dropout_head), nn.Linear(d_model, 1)) self.out_center = nn.Sequential( nn.Linear(d_model, d_model), activation_fn(), nn.Dropout(dropout_head), nn.Linear(d_model, 3)) @staticmethod def get_mask(query, mask_feats, batch_offsets): pred_masks = [] for i in range(len(batch_offsets) - 1): start_id, end_id = batch_offsets[i], batch_offsets[i + 1] pred_masks.append(torch.einsum("nd,md->nm", query[i], mask_feats[start_id:end_id])) return pred_masks def prediction_head(self, query, query_pos, mask_feats, batch_offsets, scene_ranges_max, scene_ranges_min): pred_masks = self.get_mask(query, mask_feats, batch_offsets) pred_labels = self.out_cls(query) pred_scores = self.out_score(query) pred_spatials = self.out_center(query) pred_spatials = query_pos * (scene_ranges_max - scene_ranges_min) + scene_ranges_min + pred_spatials return pred_labels, pred_scores, pred_masks, pred_spatials def get_query(self, B, batch_offsets, batch, device, dtype, kv_pos_xyz, query_feats=None): num_queris = (batch["superpoint_len"].to(device) * self.num_query_ratio).int() max_num_query = num_queris.max().item() query = torch.zeros(B, max_num_query, self.d_model, device=device, dtype=dtype) query_padding_mask = torch.ones(B, max_num_query, dtype=torch.bool, device=device) query_pos_norm = ((torch.randn(B, max_num_query, 3, device=device, dtype=dtype) + 0.5) * 0.5).clamp(0, 1) for b in range(B): start_id, end_id = batch_offsets[b], batch_offsets[b + 1] sp_xyz = kv_pos_xyz[start_id:end_id] ratio = torch.clamp(num_queris[b] / sp_xyz.size(0), max=0.99).item() fps_idx = fps(sp_xyz, ratio=ratio, random_start=True) query_pos_norm_b = ((sp_xyz[fps_idx] - sp_xyz.min(0).values) / (sp_xyz.max(0).values - sp_xyz.min(0).values)) len_b = min(num_queris[b].item(), query_pos_norm_b.size(0)) query_pos_norm_b = query_pos_norm_b[:len_b] query_padding_mask[b, :len_b] = False if self.query_init == "feat": query[b, :len_b] = query_feats[start_id:end_id][fps_idx][:len_b] query_pos_norm[b, :len_b] = query_pos_norm_b return query, query_pos_norm, query_padding_mask def forward(self, x, batch): dtype = x.dtype device = x.device batch_offsets = F.pad(batch["batched_superpoint_offset"], (1, 0)) B = len(batch_offsets) - 1 inst_feats = self.feat_proj(x) mask_feats = self.mask_proj(x) query_feats = self.query_feat_proj(x) if self.query_init == "feat" else None kv_pos_xyz = scatter_mean(batch["coord_full"], batch["batched_superpoint"], dim=0) query, query_pos_norm, query_padding_mask = self.get_query( B, batch_offsets, batch, device, dtype, kv_pos_xyz, query_feats) max_len = batch["superpoint_len"].max() key_padding_mask = torch.ones(B, max_len, dtype=torch.bool, device=device) for i in range(B): key_padding_mask[i, :batch["superpoint_len"][i]] = False kv_batched = torch.zeros(B, max_len, self.d_model, device=device, dtype=dtype) mask_feats_batched = torch.zeros(B, max_len, self.d_model, device=device, dtype=dtype) kv_pos_embedd_batched = torch.zeros(B, max_len, self.d_model, device=device, dtype=dtype) kv_pos_xyz_batched = torch.zeros(B, max_len, 3, device=device, dtype=dtype) scene_ranges_min, scene_ranges_max = [], [] for b in range(B): s, e = batch_offsets[b], batch_offsets[b + 1] kv_batched[b, :e - s] = inst_feats[s:e] mask_feats_batched[b, :e - s] = mask_feats[s:e] kv_pos_xyz_batched[b, :e - s] = kv_pos_xyz[s:e] scene_ranges_min.append(kv_pos_xyz[s:e].min(0).values) scene_ranges_max.append(kv_pos_xyz[s:e].max(0).values) scene_ranges_min = torch.stack(scene_ranges_min, 0).unsqueeze(0).permute(1, 0, 2) scene_ranges_max = torch.stack(scene_ranges_max, 0).unsqueeze(0).permute(1, 0, 2) pred_masks = None for layer_i in range(self.num_layer): query_pos_xyz = query_pos_norm * (scene_ranges_max - scene_ranges_min) + scene_ranges_min query = self.self_attn_layers[layer_i]( q_c=query, q_coords=query_pos_xyz, B=B, key_padding_mask=query_padding_mask, scene_ranges_min=scene_ranges_min) query = self.cross_attn_layers[layer_i]( q_c=query, q_p=None, k_c=kv_batched, k_p=kv_pos_embedd_batched, v=kv_batched, key_padding_mask=key_padding_mask, q_coords=query_pos_xyz, kv_coords=kv_pos_xyz_batched, pred_masks=pred_masks, B=B, scene_ranges_min=scene_ranges_min) query = self.ffn_layers[layer_i](query) if self.refinement_cross_attention and layer_i in self.refinement_cross_attention_layer_indices: ri = self.refinement_cross_attention_layer_indices.index(layer_i) mask_feats_batched = self.refinement_cross_attn_layers[ri]( q_c=mask_feats_batched, q_p=None, k_c=query, k_p=None, v=query, key_padding_mask=query_padding_mask, q_coords=kv_pos_xyz_batched, kv_coords=query_pos_xyz, pred_masks=None, B=B, scene_ranges_min=scene_ranges_min) mask_feats_batched = self.refinement_ffn_layers[ri](mask_feats_batched) query_norm = self.out_norm(query) if layer_i < self.num_layer - 1: if self.refinement_cross_attention and layer_i in self.refinement_cross_attention_layer_indices: mask_feats = torch.cat( [mask_feats_batched[b, :batch_offsets[b + 1] - batch_offsets[b]] for b in range(B)], dim=0) pred_masks = self.get_mask(query_norm, mask_feats, batch_offsets) query_pos_delta = self.query_pos_delta_head(query_norm) new_query_pos = (query_pos_norm * (scene_ranges_max - scene_ranges_min) + scene_ranges_min + query_pos_delta) new_query_pos_norm = (new_query_pos - scene_ranges_min) / (scene_ranges_max - scene_ranges_min) query_pos_norm = new_query_pos_norm.detach() if self.detach_query_pos else new_query_pos_norm # only the last layer is used at inference time if self.refinement_cross_attention: mask_feats = torch.cat( [mask_feats_batched[b, :batch_offsets[b + 1] - batch_offsets[b]] for b in range(B)], dim=0) pred_labels, pred_scores, pred_masks, pred_spatials = self.prediction_head( query_norm, query_pos_norm, mask_feats, batch_offsets, scene_ranges_max, scene_ranges_min) keep = [~query_padding_mask[b] for b in range(B)] return { "labels": [pred_labels[b][keep[b]] for b in range(B)], "scores": [pred_scores[b][keep[b]] for b in range(B)], "masks": [pred_masks[b][keep[b]] for b in range(B)], "spatials": [pred_spatials[b][keep[b]] for b in range(B)], } # =========================================================================== # # src/models/components/volt/{volt_base,decoder}.py # =========================================================================== # class Mlp(nn.Module): def __init__(self, in_features, hidden_features, act_layer=nn.GELU): super().__init__() self.fc1 = nn.Linear(in_features, hidden_features) self.act = act_layer() self.fc2 = nn.Linear(hidden_features, in_features) def forward(self, x): return self.fc2(self.act(self.fc1(x))) class Tokenizer(nn.Module): def __init__(self, in_channels, out_channels, kernel_size): super().__init__() self.kernel_size = kernel_size self.out_channels = out_channels self.proj = nn.Linear(kernel_size ** 3 * in_channels, out_channels) def forward(self, features, indices): K = self.kernel_size coarse_indices_per_voxel = indices // indices.new_tensor([1, K, K, K]) coarse_indices, inverse = torch.unique(coarse_indices_per_voxel, dim=0, sorted=True, return_inverse=True) offset = indices[:, 1:] % K offset_id = offset[:, 0] * K * K + offset[:, 1] * K + offset[:, 2] patches = features.new_zeros(coarse_indices.shape[0], K ** 3, features.shape[1]) patches[inverse, offset_id] = features return self.proj(patches.flatten(1)), coarse_indices, inverse, offset_id class VoltRoPE(nn.Module): def __init__(self, theta=100.0, freq_split=(12, 12, 8), max_grid_size=(1024, 1024, 512)): super().__init__() self.max_grid_size = max_grid_size for name, n, m in zip("xyz", freq_split, max_grid_size): freqs = 1.0 / theta ** torch.linspace(0, 1, n) self.register_buffer(f"cis_cache_{name}", self._precompute(freqs, m), persistent=False) @staticmethod def _precompute(freqs, max_pos): freqs_pos = torch.outer(torch.arange(max_pos).float(), freqs) return torch.polar(torch.ones_like(freqs_pos), freqs_pos) def compute_axial_cis_efficient(self, indices): idx = indices.clone() for a in range(3): idx[:, a] = idx[:, a].clamp(0, self.max_grid_size[a] - 1) return torch.cat([self.cis_cache_x[idx[:, 0]], self.cis_cache_y[idx[:, 1]], self.cis_cache_z[idx[:, 2]]], dim=-1).unsqueeze(0) class RoPE_Attention(nn.Module): def __init__(self, dim=768, num_heads=12, qk_norm=False): super().__init__() self.num_heads = num_heads self.h_dim = dim // num_heads self.qkv = nn.Linear(dim, 3 * dim) self.proj = nn.Linear(dim, dim) self.q_norm = nn.LayerNorm(self.h_dim) if qk_norm else nn.Identity() self.k_norm = nn.LayerNorm(self.h_dim) if qk_norm else nn.Identity() @staticmethod def apply_rotary_emb(q, k, freqs_cis): q_ = torch.view_as_complex(q.float().reshape(*q.shape[:-1], -1, 2)) k_ = torch.view_as_complex(k.float().reshape(*k.shape[:-1], -1, 2)) q_out = torch.view_as_real(q_ * freqs_cis).flatten(2) k_out = torch.view_as_real(k_ * freqs_cis).flatten(2) return q_out.type_as(q), k_out.type_as(k) def forward(self, x, freqs_cis, cu_seqlens, max_seqlen): N, C = x.shape qkv = self.qkv(x).view(N, 3, self.num_heads, self.h_dim).permute(1, 2, 0, 3) q, k, v = qkv.unbind(dim=0) q, k = self.q_norm(q).to(q.dtype), self.k_norm(k).to(k.dtype) q, k = self.apply_rotary_emb(q, k, freqs_cis) qkv = torch.stack([q, k, v], dim=0).permute(2, 0, 1, 3) qkv_dtype = qkv.dtype # upstream runs this through FlashAttention-2 in fp16 attn_dtype = torch.float16 if qkv.is_cuda else torch.float32 x = varlen_qkvpacked_attention(qkv.to(attn_dtype), cu_seqlens, max_seqlen) return self.proj(x.reshape(-1, C).to(qkv_dtype)) class Block(nn.Module): def __init__(self, dim=768, num_heads=12, mlp_ratio=4.0, qk_norm=False, act_layer=nn.GELU, norm_layer=nn.LayerNorm): super().__init__() self.norm1 = norm_layer(dim) self.attn = RoPE_Attention(dim=dim, num_heads=num_heads, qk_norm=qk_norm) self.ls1 = nn.Identity() self.drop_path1 = nn.Identity() self.norm2 = norm_layer(dim) self.mlp = Mlp(in_features=dim, hidden_features=int(dim * mlp_ratio), act_layer=act_layer) self.ls2 = nn.Identity() self.drop_path2 = nn.Identity() def forward(self, x, freqs_cis, cu_seq_lens, max_seqlen): x = x + self.drop_path1(self.ls1(self.attn(self.norm1(x), freqs_cis, cu_seq_lens, max_seqlen))) x = x + self.drop_path2(self.ls2(self.mlp(self.norm2(x)))) return x class Detokenizer(nn.Module): def __init__(self, in_channels, out_channels, kernel_size): super().__init__() self.kernel_size = kernel_size self.out_channels = out_channels self.proj = nn.Linear(in_channels, kernel_size ** 3 * out_channels, bias=False) self.bias = nn.Parameter(torch.zeros(out_channels)) def forward(self, coarse_features, inverse, offset_id): K = self.kernel_size all_offsets = self.proj(coarse_features).view(-1, K ** 3, self.out_channels) return all_offsets[inverse, offset_id] + self.bias class VoltDecoder(nn.Module): def __init__(self, in_channels, out_channels, kernel_size, norm_layer=partial(nn.BatchNorm1d, eps=1e-3, momentum=0.01)): super().__init__() act_layer = nn.GELU self.pre = nn.Sequential(norm_layer(in_channels), act_layer(), nn.Linear(in_channels, out_channels, bias=False), norm_layer(out_channels), act_layer()) self.unembed = Detokenizer(out_channels, out_channels, kernel_size=kernel_size) self.post = nn.Sequential(norm_layer(out_channels), act_layer()) def forward(self, x, inverse, offset_id): return self.post(self.unembed(self.pre(x), inverse, offset_id)) class Volt(nn.Module): def __init__(self, in_channels=6, embed_dim=768, depth=12, num_heads=12, mlp_ratio=4, qk_norm=True, stride=5, kernel_size=5, out_channels=128): super().__init__() assert stride == kernel_size self.tokenizer = Tokenizer(in_channels, embed_dim, kernel_size) self.blocks = nn.Sequential(*[ Block(dim=embed_dim, num_heads=num_heads, mlp_ratio=mlp_ratio, qk_norm=qk_norm) for _ in range(depth)]) self.pos_enc = VoltRoPE() self.decoder = VoltDecoder(in_channels=embed_dim, out_channels=out_channels, kernel_size=kernel_size) @staticmethod def compute_seqlens(batch_indices): points_per_batch = torch.bincount(batch_indices + 1) cu_seqlens = torch.cumsum(points_per_batch, dim=0, dtype=torch.int32) seq_lens = cu_seqlens[1:] - cu_seqlens[:-1] return cu_seqlens, seq_lens.max().item() def forward(self, data_dict): grid_coord = data_dict["coord_grid"] feat = data_dict["feat"] indices = torch.cat([data_dict["batch_indices"].unsqueeze(-1).int(), grid_coord.int()], dim=1).contiguous() features, indices, inverse, offset_id = self.tokenizer(feat, indices) cu_seqlens, max_seqlen = self.compute_seqlens(indices[:, 0]) freqs_cis = self.pos_enc.compute_axial_cis_efficient(indices[:, 1:]) for blk in self.blocks: features = blk(features, freqs_cis, cu_seqlens, max_seqlen) return self.decoder(features, inverse, offset_id) # =========================================================================== # # src/models/base_instance_former.py # =========================================================================== # class AQ3D(nn.Module): """AQ3D with the Volt-B backbone, configured for ScanNet200 (198 classes).""" def __init__(self, num_classes=198, in_features=6, mid_features=128): super().__init__() self.num_classes = num_classes rope = partial(RoPE, theta=100.0, head_split=[16, 16, 16], grid_size=0.05, max_grid_size=[512, 512, 256]) self.backbone = Volt(in_channels=in_features, embed_dim=768, depth=12, num_heads=12, mlp_ratio=4, qk_norm=True, stride=5, kernel_size=5, out_channels=mid_features) self.decoder = QueryDecoder( num_layer=6, max_query=True, num_query_ratio=0.6, query_init="feat", query_pos_init="adaptive", dropout_head=0.1, dropout_layer=0.2, cosine_classifier=True, refinement_cross_attention=True, refinement_cross_attention_layer_indices=[1, 3, 5], num_class=num_classes, in_channel=mid_features, d_model=384, activation_fn=nn.ReLU, self_attention_layer=partial(SelfAttentionLayer, nhead=8, dropout=0.0, rope=rope()), cross_attention_layer=partial(CrossAttentionLayer, nhead=8, dropout=0.0, attn_mask_thres=0.1, rope=rope()), refinement_cross_attention_layer=partial(CrossAttentionLayer, nhead=8, dropout=0.0, rope=rope()), ffn_layer=partial(FFN, hidden_dim=1024, dropout=0.0, activation_fn=nn.GELU), ) self.pool_attn = nn.Sequential( nn.Linear(mid_features, mid_features), nn.LayerNorm(mid_features), nn.ReLU(), nn.Linear(mid_features, mid_features), nn.LayerNorm(mid_features), nn.ReLU(), nn.Linear(mid_features, mid_features)) def forward(self, batch): feat = self.backbone(batch) feat = feat[batch["batched_inverse"]] scores = self.pool_attn(feat) weights = scatter_softmax(scores, batch["batched_superpoint"], dim=0) feat = scatter_sum(feat * weights, batch["batched_superpoint"], dim=0) return self.decoder(feat, batch)