| from __future__ import annotations |
|
|
| import math |
| from typing import Tuple |
|
|
| import torch |
|
|
|
|
| def sparse_downsample(feats: torch.Tensor, coords: torch.Tensor, factor: Tuple[int, int, int]): |
| """Reference matching AniGen SparseDownsample (mean pooling + inverse map).""" |
| coord = list(coords.unbind(dim=-1)) |
| for i, f in enumerate(factor): |
| coord[i + 1] = coord[i + 1] // int(f) |
| maxs = [int(coord[i + 1].max().item()) + 1 for i in range(3)] |
| off = torch.cumprod(torch.tensor(maxs[::-1], dtype=torch.int64), 0).tolist()[::-1] + [1] |
| code = sum(c.to(torch.int64) * int(o) for c, o in zip(coord, off)) |
| unique_code, inverse = code.unique(return_inverse=True) |
| out = torch.zeros((unique_code.shape[0], feats.shape[1]), device=feats.device, dtype=feats.dtype) |
| out = torch.scatter_reduce( |
| out, |
| 0, |
| inverse[:, None].expand(-1, feats.shape[1]), |
| feats, |
| reduce="mean", |
| ) |
| out_coords = torch.stack( |
| [unique_code // off[0]] + [(unique_code // off[i + 1]) % maxs[i] for i in range(3)], |
| dim=-1, |
| ).to(torch.int32) |
| return out, out_coords, inverse.to(torch.int32) |
|
|
|
|
| def sparse_upsample(feats: torch.Tensor, target_coords: torch.Tensor, inverse: torch.Tensor): |
| return feats[inverse.to(torch.long)], target_coords |
|
|
|
|
| def sparse_subdivide(feats: torch.Tensor, coords: torch.Tensor): |
| offsets = torch.tensor( |
| [[0, x, y, z] for x in (0, 1) for y in (0, 1) for z in (0, 1)], |
| device=coords.device, |
| dtype=coords.dtype, |
| ) |
| out_coords = coords.clone() |
| out_coords[:, 1:] *= 2 |
| out_coords = (out_coords[:, None, :] + offsets[None, :, :]).flatten(0, 1) |
| out_feats = feats[:, None, :].expand(feats.shape[0], 8, feats.shape[1]).flatten(0, 1) |
| return out_feats, out_coords |
|
|
|
|
| def window_partition(coords: torch.Tensor, window_size: int, shift: Tuple[int, int, int]): |
| shifted = coords.clone().detach() |
| shifted[:, 1:] += torch.tensor(shift, device=coords.device, dtype=torch.int32)[None] |
| max_coords = shifted[:, 1:].max(dim=0).values.tolist() |
| num_windows = [math.ceil((int(v) + 1) / window_size) for v in max_coords] |
| offset = torch.cumprod(torch.tensor([1] + num_windows[::-1]), dim=0).tolist()[::-1] |
| shifted[:, 1:] //= int(window_size) |
| indices = (shifted * torch.tensor(offset, device=coords.device, dtype=torch.int32)[None]).sum(dim=1) |
| fwd = torch.argsort(indices) |
| bwd = torch.empty_like(fwd) |
| bwd[fwd] = torch.arange(fwd.shape[0], device=coords.device) |
| seq_lens = torch.bincount(indices) |
| mask = seq_lens != 0 |
| return fwd, bwd, seq_lens[mask].to(torch.int32) |
|
|
|
|
| def sparse_window_attention(qkv: torch.Tensor, coords: torch.Tensor, window_size: int, shift=(0, 0, 0)): |
| """Reference path matching AniGen FlashAttention window semantics.""" |
| import flash_attn |
|
|
| fwd, bwd, seq_lens = window_partition(coords, window_size, tuple(int(x) for x in shift)) |
| qkv_sorted = qkv[fwd] |
| cu = torch.cat( |
| [torch.zeros(1, device=qkv.device, dtype=torch.int32), torch.cumsum(seq_lens, 0, dtype=torch.int32)], |
| 0, |
| ) |
| out = flash_attn.flash_attn_varlen_qkvpacked_func(qkv_sorted, cu, int(seq_lens.max().item())) |
| return out[bwd] |
|
|