"""Dependency-free replacements for the compiled extensions AQ3D relies on. The reference implementation of AQ3D (https://github.com/kenomo/aq3d) uses ``torch_scatter``, ``torch_geometric.nn.pool.fps`` and ``flash_attn``. None of those ship Blackwell (sm_120) wheels usable inside a ZeroGPU Space, so this module re-implements exactly the operators the inference path needs with plain PyTorch / NumPy. Semantics are matched 1:1 with the originals for the shapes that occur during inference (1-D index, ``dim=0``). """ from typing import Optional import numpy as np import torch import torch.nn.functional as F # --------------------------------------------------------------------------- # # torch_scatter replacements (1-D index, reduce over dim 0) # --------------------------------------------------------------------------- # def _prepare(src: torch.Tensor, index: torch.Tensor, dim: int, dim_size: Optional[int]): if dim < 0: dim = src.dim() + dim assert dim == 0, "only dim=0 scatters are used by AQ3D inference" assert index.dim() == 1, "only 1-D indices are used by AQ3D inference" if dim_size is None: dim_size = int(index.max()) + 1 if index.numel() else 0 idx = index.view(-1, *([1] * (src.dim() - 1))).expand_as(src) size = (dim_size,) + tuple(src.shape[1:]) return idx, size, dim_size def scatter_sum(src: torch.Tensor, index: torch.Tensor, dim: int = 0, dim_size: Optional[int] = None) -> torch.Tensor: idx, size, _ = _prepare(src, index, dim, dim_size) out = src.new_zeros(size) return out.scatter_add_(0, idx, src) scatter_add = scatter_sum def scatter_mean(src: torch.Tensor, index: torch.Tensor, dim: int = 0, dim_size: Optional[int] = None) -> torch.Tensor: _, _, n = _prepare(src, index, dim, dim_size) out = scatter_sum(src, index, dim, n) count = torch.zeros(n, dtype=src.dtype, device=src.device) count.scatter_add_(0, index, torch.ones_like(index, dtype=src.dtype)) count = count.clamp(min=1).view(-1, *([1] * (src.dim() - 1))) return out / count def scatter_max(src: torch.Tensor, index: torch.Tensor, dim: int = 0, dim_size: Optional[int] = None): idx, size, _ = _prepare(src, index, dim, dim_size) out = src.new_full(size, float("-inf")) out.scatter_reduce_(0, idx, src, reduce="amax", include_self=True) out = torch.where(torch.isneginf(out), torch.zeros_like(out), out) return out, None def scatter_softmax(src: torch.Tensor, index: torch.Tensor, dim: int = 0, dim_size: Optional[int] = None) -> torch.Tensor: idx, size, n = _prepare(src, index, dim, dim_size) max_value = src.new_full(size, float("-inf")) max_value.scatter_reduce_(0, idx, src, reduce="amax", include_self=True) max_value = torch.where(torch.isneginf(max_value), torch.zeros_like(max_value), max_value) exped = (src - max_value.gather(0, idx)).exp() denom = scatter_sum(exped, index, 0, n) return exped / (denom.gather(0, idx) + 1e-16) # --------------------------------------------------------------------------- # # torch_geometric.nn.pool.fps replacement # --------------------------------------------------------------------------- # def fps(x: torch.Tensor, ratio: float = 0.5, random_start: bool = True) -> torch.Tensor: """Farthest point sampling over a single (non-batched) point set. Matches ``torch_geometric.nn.pool.fps`` for a single batch element: returns ``ceil(ratio * N)`` indices in the order they were selected, with a random first point when ``random_start`` is set. """ n = int(x.size(0)) k = int(np.ceil(ratio * n)) k = max(1, min(n, k)) start = int(torch.randint(0, n, (1,)).item()) if random_start else 0 pts = x.detach().float().cpu().numpy() idx = np.empty(k, dtype=np.int64) idx[0] = start dist = np.full(n, np.inf, dtype=np.float32) last = pts[start] for i in range(1, k): d = ((pts - last) ** 2).sum(-1) np.minimum(dist, d, out=dist) nxt = int(dist.argmax()) idx[i] = nxt last = pts[nxt] return torch.from_numpy(idx).to(x.device) # --------------------------------------------------------------------------- # # flash_attn.flash_attn_varlen_qkvpacked_func replacement # --------------------------------------------------------------------------- # def varlen_qkvpacked_attention(qkv: torch.Tensor, cu_seqlens: torch.Tensor, max_seqlen: int) -> torch.Tensor: """``qkv``: [total_tokens, 3, heads, head_dim] -> [total_tokens, heads, dim]. Equivalent to ``flash_attn.flash_attn_varlen_qkvpacked_func`` (no dropout, non-causal, softmax_scale = 1/sqrt(head_dim)), implemented with PyTorch SDPA, which dispatches to the memory-efficient / flash kernels on CUDA. """ total, three, heads, head_dim = qkv.shape assert three == 3 out = torch.empty(total, heads, head_dim, dtype=qkv.dtype, device=qkv.device) bounds = cu_seqlens.tolist() for i in range(len(bounds) - 1): s, e = int(bounds[i]), int(bounds[i + 1]) if e <= s: continue q, k, v = qkv[s:e].permute(1, 2, 0, 3).unbind(0) # each [H, L, D] o = F.scaled_dot_product_attention(q.unsqueeze(0), k.unsqueeze(0), v.unsqueeze(0)) out[s:e] = o.squeeze(0).transpose(0, 1) return out