multimodalart's picture
multimodalart HF Staff
Upload folder using huggingface_hub
9550667 verified
Raw
History Blame Contribute Delete
5.49 kB
"""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