changh95's picture
tt-model push diffusion-planner-p150 (container)
be62f78 verified
Raw History Blame Contribute Delete
22 kB
# SPDX-License-Identifier: Apache-2.0
"""C13 sparse rulebooks: spconv-faithful neighbour maps of sparse 3-D convolutions, built on the host, in the
*gather* form the device consumes (PLAN.md section 1.2 row C13; owned by the BEVFusion port; S:bevfusion:465-466,
S:ptv3:158).
A sparse tensor is a set of active voxels ``coords (N, 3)`` = ``(x, y, z)`` (batch 0 implied) in a grid of
``spatial_shape (X, Y, Z)``. Every sparse conv of the Autoware graphs (``autoware::GetIndicePairsImplicitGemm`` +
``ImplicitGemm``, spconv 2.3.8 ``MaskImplicitGemm``) is, per output voxel ``o``,
out[o] = sum_k W[:, k, :] @ in[ nmap[o, k] ] (taps with ``nmap[o, k] == -1`` contribute nothing)
with ``W`` the KRSC filter ``[Cout, k0, k1, k2, Cin]`` (``k0`` along x, ``k1`` along y, ``k2`` along z) and the tap
index ``k = (a * k1 + b) * k2 + c`` (row-major over (x, y, z), spconv's ``layout_rs``). This module builds
``nmap`` -- a pure gather table, no scatter, for sub-manifold and strided layers alike -- and the device-side index
forms of it:
- :class:`SparseLevel`: the active set of one resolution with a hashed lookup (sorted int64 keys + ``searchsorted``;
keys are exact for any coordinate, also outside the grid);
- :func:`subm_neighbor_map`: sub-manifold conv (output set == input set; spconv uses ``padding = (k // 2) *
dilation`` whatever the attribute says, ``spconv/csrc/sparse/indices.py:1522``). ``rule="spconv"`` (default)
reproduces spconv's **13-query + symmetric-write** construction (``generate_subm_conv_inds`` launches only the taps
``k <= K // 2``, ``indices.py:1532``; a hit also writes the mirrored pair ``pair[K-1-k][in] = out``, ``:865``; the
centre is always the voxel itself, ``:838-845``). Its bounds check is applied to the *queried* coordinate, so a tap
``k < K // 2`` is accepted only if the **neighbour** lies inside the grid and a tap ``k > K // 2`` only if the
**output voxel** does (the asymmetric-bounds rule of S:ptv3:158, verified bit-exact against spconv there). For
voxel sets inside the grid (every BEVFusion level) it equals ``rule="exact"``, the plain 27-tap lookup;
- :func:`strided_neighbor_map`: regular (strided) sparse conv: output ``o`` is active iff some active input ``i``
and tap ``k`` give ``i = o * s - p + k * d`` with ``0 <= o < out_shape`` (spconv ``generate_conv_inds``); outputs
come in ascending linear-key order (any order is equivalent: the next layers are permutation-equivariant); one
input per (output, tap), so the strided layer is a gather too;
- device forms (C28, probe P4): :func:`gather_index` (``-1`` -> an explicit zero sentinel row, capacity padding),
:func:`im2col_index` (the **tile-ordered** index of the one-gather im2col: ``(voxel block of 32, tap, voxel)``,
so ``ttnn.embedding(..., layout=TILE)`` of a ``[N/32, K * 32]`` index into a ``Cin % 32 == 0`` table gives a tensor
whose tile sequence *is* the ``[N, K * Cin]`` im2col matrix, read through ``ttnn.experimental.view``) and its
oracle :func:`im2col_numpy`; :func:`dense_gather_index` (to-dense as a gather: one ``[X * Y]`` canvas index per z
slice);
- capacities (PLAN.md 0.2: grid-independent constants chosen from data): :func:`select_capacity` picks the smallest
bucket that holds a count, and flags an overflow;
- oracles: :func:`sparse_conv_numpy` (the gather form above, float64 accumulation) and :func:`dense_conv3d_numpy` (a
brute-force dense 3-D cross-correlation restricted to the active outputs; spconv == dense Conv3d on the active set).
Not here yet (their consumer adds them, PTv3): the linear-key **alias** mode of S:ptv3:388-393 (spconv inserts
out-of-grid voxels under aliasing int32 keys; ``rule="alias"`` raises ``NotImplementedError``) and Morton /
serialization helpers.
numpy only; no side effects on import.
"""
from __future__ import annotations
from dataclasses import dataclass
from typing import Any, Dict, Iterable, Optional, Sequence, Tuple
import numpy as np
__all__ = [
"SparseLevel",
"RULES",
"kernel_offsets",
"conv_out_shape",
"subm_neighbor_map",
"strided_neighbor_map",
"pairs_count",
"gather_index",
"im2col_index",
"im2col_numpy",
"dense_gather_index",
"dense_numpy",
"select_capacity",
"sparse_conv_numpy",
"dense_conv3d_numpy",
"TILE_ROWS",
"neighbor_stats",
"iter_taps",
]
RULES = ("spconv", "exact") # subm neighbour rules ("alias": PTv3, not implemented yet)
TILE_ROWS = 32 # rows of one ttnn tile: the voxel block of the tile-ordered im2col index
def _triple(v: Any, name: str) -> Tuple[int, int, int]:
if isinstance(v, (int, np.integer)):
t = (int(v),) * 3
else:
t = tuple(int(x) for x in v)
if len(t) != 3:
raise ValueError(f"{name} needs 3 values, got {v!r}")
return t # type: ignore[return-value]
# ------------------------------------------------------------------------------------------------- levels
@dataclass
class SparseLevel:
"""The active voxels of one resolution: ``coords`` int64 ``(N, 3)`` ``(x, y, z)`` in ``spatial_shape (X, Y, Z)``.
Rows are the feature rows of the level (row ``i`` holds voxel ``coords[i]``). Coordinates outside the grid are
allowed (spconv carries them, S:ptv3:158); keys are exact for them (no aliasing). Duplicate coordinates are
refused (a spconv indice set is unique)."""
coords: np.ndarray
spatial_shape: Tuple[int, int, int]
def __post_init__(self) -> None:
c = np.asarray(self.coords)
if c.ndim != 2 or c.shape[1] != 3:
raise ValueError(f"coords must be (N, 3) (x, y, z), got {c.shape}")
if c.size and not np.issubdtype(c.dtype, np.integer):
raise ValueError(f"coords must be integers, got {c.dtype}")
self.coords = np.ascontiguousarray(c, dtype=np.int64)
self.spatial_shape = _triple(self.spatial_shape, "spatial_shape")
if min(self.spatial_shape) <= 0:
raise ValueError(f"spatial_shape must be positive, got {self.spatial_shape}")
keys = self.keys_of(self.coords)
order = np.argsort(keys, kind="stable")
self._sorted = keys[order]
self._perm = order
if len(self._sorted) > 1 and np.any(self._sorted[1:] == self._sorted[:-1]):
raise ValueError("coords hold duplicate voxels")
# exact keys over the bounding box of the grid and the voxels, widened by a margin, so no coordinate aliases
_MARGIN = 8
def _box(self) -> Tuple[np.ndarray, np.ndarray]:
box = getattr(self, "_box_cache", None)
if box is None:
shape = np.asarray(self.spatial_shape, dtype=np.int64)
lo, hi = np.zeros(3, np.int64), shape.copy()
if len(self.coords):
lo = np.minimum(lo, self.coords.min(axis=0))
hi = np.maximum(hi, self.coords.max(axis=0) + 1)
box = (lo - self._MARGIN, hi + self._MARGIN)
self._box_cache = box
return box
def keys_of(self, q: np.ndarray) -> np.ndarray:
"""Exact int64 keys of ``q (M, 3)``; coordinates outside the widened bounding box get key -1 (never found:
no voxel of the level lies there)."""
q = np.asarray(q, dtype=np.int64).reshape(-1, 3)
lo, hi = self._box()
dims = hi - lo
ok = np.all((q >= lo) & (q < hi), axis=1)
r = q - lo
k = (r[:, 0] * dims[1] + r[:, 1]) * dims[2] + r[:, 2]
return np.where(ok, k, -1)
@property
def num_active(self) -> int:
return int(len(self.coords))
def in_grid(self, q: Optional[np.ndarray] = None) -> np.ndarray:
"""``0 <= q < spatial_shape`` on every axis (default: the level's own coords)."""
q = self.coords if q is None else np.asarray(q, dtype=np.int64)
return np.all((q >= 0) & (q < np.asarray(self.spatial_shape, dtype=np.int64)), axis=1)
def lookup(self, q: np.ndarray) -> np.ndarray:
"""Rows of the voxels at ``q (M, 3)`` (int64), ``-1`` where no active voxel is."""
keys = self.keys_of(q)
if len(self._sorted) == 0:
return np.full(len(keys), -1, np.int64)
pos = np.searchsorted(self._sorted, keys).clip(max=len(self._sorted) - 1)
found = (self._sorted[pos] == keys) & (keys >= 0)
return np.where(found, self._perm[pos], -1)
def linear_index(self, order: str = "xyz") -> np.ndarray:
"""In-grid linear index of every voxel: ``(x * Y + y) * Z + z`` (``"xyz"``, spconv's for batch 0)."""
if order != "xyz":
raise ValueError("only the xyz order is defined")
X, Y, Z = self.spatial_shape
c = self.coords
return (c[:, 0] * Y + c[:, 1]) * Z + c[:, 2]
def describe(self) -> Dict[str, Any]:
return {"num_active": self.num_active, "spatial_shape": list(self.spatial_shape),
"out_of_grid": int((~self.in_grid()).sum()) if self.num_active else 0}
# ---------------------------------------------------------------------------------------------- rulebooks
def kernel_offsets(kernel: Any = 3, dilation: Any = 1) -> np.ndarray:
"""``(K, 3)`` int64 tap offsets ``(a * d0, b * d1, c * d2)`` in tap order ``k = (a * k1 + b) * k2 + c``."""
k = _triple(kernel, "kernel")
d = _triple(dilation, "dilation")
a, b, c = np.meshgrid(np.arange(k[0]), np.arange(k[1]), np.arange(k[2]), indexing="ij")
return np.stack([a.ravel() * d[0], b.ravel() * d[1], c.ravel() * d[2]], axis=1).astype(np.int64)
def conv_out_shape(spatial_shape: Any, kernel: Any, stride: Any, padding: Any, dilation: Any = 1
) -> Tuple[int, int, int]:
"""spconv / Conv3d output grid: ``(in + 2 p - d (k - 1) - 1) // s + 1`` per axis."""
s_in, k, s, p, d = (np.asarray(_triple(v, n)) for v, n in ((spatial_shape, "spatial_shape"), (kernel, "kernel"),
(stride, "stride"), (padding, "padding"),
(dilation, "dilation")))
out = (s_in + 2 * p - d * (k - 1) - 1) // s + 1
if np.any(out <= 0):
raise ValueError(f"empty output grid {tuple(out)}")
return tuple(int(v) for v in out) # type: ignore[return-value]
def subm_neighbor_map(level: SparseLevel, kernel: Any = 3, dilation: Any = 1, *, rule: str = "spconv") -> np.ndarray:
"""Sub-manifold conv neighbour map ``(N, K)`` int32: ``nmap[o, k]`` = the input row read by tap ``k`` of output
``o`` (the level's own rows), ``-1`` if none (module docstring for ``rule``)."""
if rule == "alias":
raise NotImplementedError("the linear-key alias mode (S:ptv3:388-393) is added by the PTv3 port")
if rule not in RULES:
raise ValueError(f"rule must be one of {RULES}")
k = _triple(kernel, "kernel")
if any(v % 2 == 0 for v in k):
raise ValueError(f"sub-manifold kernels are odd, got {k}")
d = np.asarray(_triple(dilation, "dilation"), dtype=np.int64)
pad = (np.asarray(k, dtype=np.int64) // 2) * d
offs = kernel_offsets(k, d) - pad # (K, 3) centred offsets
K = len(offs)
n = level.num_active
out = np.full((n, K), -1, dtype=np.int32)
if n == 0:
return out
c = level.coords
q = (c[:, None, :] + offs[None, :, :]).reshape(-1, 3)
rows = level.lookup(q).reshape(n, K)
if rule == "spconv":
centre = K // 2
own_in = level.in_grid() # (N,)
nb_in = level.in_grid(q).reshape(n, K)
low = np.arange(K) < centre
accept = np.where(low[None, :], nb_in, own_in[:, None])
rows = np.where(accept, rows, -1)
rows[:, centre] = np.arange(n) # the centre tap is the voxel itself
out[:] = rows
return out
def strided_neighbor_map(level: SparseLevel, kernel: Any, stride: Any, padding: Any, dilation: Any = 1
) -> Tuple[SparseLevel, np.ndarray]:
"""Regular sparse conv: ``(out_level, nmap (N_out, K) int32)``; ``nmap[o, k]`` = the input row with
``i = o * s - p + k * d`` (``-1``: none). Output voxels in ascending ``(x, y, z)`` linear-key order. The input
voxels must lie inside their grid (out-of-grid inputs would change spconv's set; refused)."""
k = _triple(kernel, "kernel")
s = np.asarray(_triple(stride, "stride"), dtype=np.int64)
p = np.asarray(_triple(padding, "padding"), dtype=np.int64)
d = np.asarray(_triple(dilation, "dilation"), dtype=np.int64)
out_shape = conv_out_shape(level.spatial_shape, k, s, p, d)
if level.num_active and not level.in_grid().all():
raise ValueError("strided rulebooks need every input voxel inside its grid")
offs = kernel_offsets(k, d) # (K, 3), uncentred: tap k at + k * d
K = len(offs)
c = level.coords
num = c[:, None, :] + p - offs[None, :, :] # = o * s for a hit
ok = np.all(num % s == 0, axis=2)
o = num // s
ok &= np.all((o >= 0) & (o < np.asarray(out_shape, dtype=np.int64)), axis=2)
ii, kk = np.nonzero(ok)
oc = o[ii, kk]
X, Y, Z = out_shape
okey = (oc[:, 0] * Y + oc[:, 1]) * Z + oc[:, 2]
uniq, inv = np.unique(okey, return_inverse=True)
out_coords = np.stack([uniq // (Y * Z), (uniq // Z) % Y, uniq % Z], axis=1).astype(np.int64)
out_level = SparseLevel(out_coords, out_shape)
nmap = np.full((len(uniq), K), -1, dtype=np.int32)
nmap[inv.reshape(-1), kk] = ii # one input per (output, tap)
return out_level, nmap
def pairs_count(nmap: np.ndarray) -> int:
"""Number of (input, output, tap) pairs of a neighbour map (the rulebook size; FLOPs = 2 Cin Cout pairs)."""
return int(np.count_nonzero(np.asarray(nmap) >= 0))
# ------------------------------------------------------------------------------------------- device forms
def gather_index(nmap: np.ndarray, *, sentinel: int, rows: Optional[int] = None) -> np.ndarray:
"""``(rows, K)`` uint32: ``nmap`` with ``-1`` -> ``sentinel`` (the zero row of the feature table) and the rows
past ``len(nmap)`` (capacity padding) all ``sentinel``. ``sentinel`` must exceed every valid row."""
nm = np.asarray(nmap)
n, K = nm.shape
rows = n if rows is None else int(rows)
if rows < n:
raise ValueError(f"{n} active rows exceed the capacity {rows}")
if n and nm.max() >= sentinel:
raise ValueError(f"row {int(nm.max())} >= sentinel {sentinel}")
if sentinel < 0 or sentinel >= 2 ** 32:
raise ValueError("sentinel must fit uint32")
out = np.full((rows, K), sentinel, dtype=np.uint32)
out[:n] = np.where(nm >= 0, nm, sentinel).astype(np.uint32)
return out
def im2col_index(nmap: np.ndarray, *, sentinel: int, rows: Optional[int] = None, block: int = TILE_ROWS
) -> np.ndarray:
"""The tile-ordered index of the one-gather im2col (C28, probe P4): ``(rows / block, K * block)`` uint32 with
``index[b, k * block + v] = gather_index[b * block + v, k]``. ``ttnn.embedding`` of it into a ``[1, 1, V, Cin]``
table (``Cin % 32 == 0``) with ``layout=TILE`` yields ``[rows / block, K * block, Cin]``, whose tiles, in memory
order, are those of the ``[rows, K * Cin]`` im2col matrix (row ``o``: the ``K`` tap rows of output ``o``
concatenated). ``rows`` (default ``len(nmap)`` rounded up) must be a multiple of ``block``."""
nm = np.asarray(nmap)
n, K = nm.shape
rows = -(-n // block) * block if rows is None else int(rows)
if rows % block:
raise ValueError(f"rows {rows} is not a multiple of the block {block}")
g = gather_index(nm, sentinel=sentinel, rows=rows) # (rows, K)
return np.ascontiguousarray(g.reshape(rows // block, block, K).transpose(0, 2, 1).reshape(rows // block,
K * block))
def im2col_numpy(table: np.ndarray, index: np.ndarray, *, kernel_volume: int, block: int = TILE_ROWS) -> np.ndarray:
"""Oracle of the device im2col: the ``[rows, K * C]`` matrix that :func:`im2col_index` describes, built from the
row table ``(V, C)`` (``index`` values are rows of ``table``)."""
t = np.asarray(table)
idx = np.asarray(index, dtype=np.int64)
nb = idx.shape[0]
K = int(kernel_volume)
if idx.shape[1] != K * block:
raise ValueError(f"index has {idx.shape[1]} columns, expected {K} x {block}")
g = idx.reshape(nb, K, block).transpose(0, 2, 1).reshape(nb * block, K) # back to (rows, K)
return t[g].reshape(nb * block, K * t.shape[1])
def dense_gather_index(level: SparseLevel, *, sentinel: int, split_axis: int = 2) -> np.ndarray:
"""To-dense as gathers: ``(S, A * B)`` uint32 with ``index[s, a * B + b]`` = the row of voxel at ``(a, b)`` in
slice ``s`` of ``split_axis`` (the two other axes in (x, y, z) order give ``a``, ``b``), ``sentinel`` where the
cell is empty. BEVFusion's to-dense: ``split_axis = 2`` (z), ``index[z, x * Y + y]``; the slices are concatenated
or interleaved on the channel axis afterwards."""
if split_axis not in (0, 1, 2):
raise ValueError("split_axis must be 0, 1 or 2")
if level.num_active and not level.in_grid().all():
raise ValueError("to-dense needs every voxel inside the grid")
if level.num_active and level.num_active > sentinel:
raise ValueError("sentinel must not be a valid row")
axes = [a for a in range(3) if a != split_axis]
S = level.spatial_shape[split_axis]
A, B = level.spatial_shape[axes[0]], level.spatial_shape[axes[1]]
out = np.full((S, A * B), sentinel, dtype=np.uint32)
c = level.coords
out[c[:, split_axis], c[:, axes[0]] * B + c[:, axes[1]]] = np.arange(level.num_active, dtype=np.uint32)
return out
def dense_numpy(features: np.ndarray, level: SparseLevel) -> np.ndarray:
"""``(X, Y, Z, C)`` dense tensor of the level's features (zeros elsewhere): ONNX ``ScatterND`` into
``[X, Y, Z, C]``, the reference of the to-dense gathers."""
f = np.asarray(features)
X, Y, Z = level.spatial_shape
out = np.zeros((X, Y, Z, f.shape[1]), dtype=f.dtype)
c = level.coords
out[c[:, 0], c[:, 1], c[:, 2]] = f[: level.num_active]
return out
# ------------------------------------------------------------------------------------------------ capacity
def select_capacity(count: int, buckets: Sequence[int]) -> Tuple[int, bool]:
"""The smallest bucket ``>= count`` and ``False``; the largest bucket and ``True`` (overflow: the caller keeps
that many and flags the frame) when none holds it. Buckets are grid-independent constants (PLAN.md 0.2)."""
b = sorted(int(v) for v in buckets)
if not b or b[0] <= 0:
raise ValueError("buckets must be positive")
for v in b:
if count <= v:
return v, False
return b[-1], True
# ------------------------------------------------------------------------------------------------- oracles
def sparse_conv_numpy(features: np.ndarray, nmap: np.ndarray, weight: np.ndarray, bias: Optional[np.ndarray] = None,
*, dtype: Any = np.float64) -> np.ndarray:
"""The gather form of the sparse conv: ``out[o] = sum_k W[:, k, :] @ features[nmap[o, k]] (+ bias)``.
``weight`` KRSC ``[Cout, k0, k1, k2, Cin]``; accumulation in ``dtype`` (float64 default)."""
f = np.asarray(features, dtype=dtype)
w = np.asarray(weight, dtype=dtype)
cout, cin = w.shape[0], w.shape[-1]
wk = w.reshape(cout, -1, cin) # (Cout, K, Cin)
nm = np.asarray(nmap)
out = np.zeros((len(nm), cout), dtype=dtype)
for k in range(nm.shape[1]):
sel = nm[:, k] >= 0
if sel.any():
out[sel] += f[nm[sel, k]] @ wk[:, k, :].T
if bias is not None:
out += np.asarray(bias, dtype=dtype)
return out
def dense_conv3d_numpy(features: np.ndarray, level: SparseLevel, weight: np.ndarray, out_coords: np.ndarray, *,
stride: Any = 1, padding: Any = 0, dilation: Any = 1) -> np.ndarray:
"""Brute-force dense 3-D cross-correlation (``torch.nn.Conv3d`` semantics, zero padding) of the densified input,
read at ``out_coords``: the definition a sparse conv must equal on its active outputs. Float64; small grids."""
w = np.asarray(weight, dtype=np.float64)
k = w.shape[1:4]
s = np.asarray(_triple(stride, "stride"))
p = np.asarray(_triple(padding, "padding"))
d = np.asarray(_triple(dilation, "dilation"))
X, Y, Z = level.spatial_shape
dense = dense_numpy(np.asarray(features, dtype=np.float64), level)
oc = np.asarray(out_coords, dtype=np.int64)
out = np.zeros((len(oc), w.shape[0]))
for a in range(k[0]):
for b in range(k[1]):
for c in range(k[2]):
i = oc * s - p + np.array([a, b, c]) * d
inside = np.all((i >= 0) & (i < np.array([X, Y, Z])), axis=1)
if inside.any():
ii = i[inside]
out[inside] += dense[ii[:, 0], ii[:, 1], ii[:, 2]] @ w[:, a, b, c, :].T
return out
def neighbor_stats(nmap: np.ndarray) -> Dict[str, Any]:
"""``{rows, pairs, mean_taps, max_taps}`` of a neighbour map (diagnostics, PORT_LOG tables)."""
nm = np.asarray(nmap)
taps = (nm >= 0).sum(axis=1) if len(nm) else np.zeros(0)
return {"rows": int(len(nm)), "pairs": int(taps.sum()), "mean_taps": float(taps.mean()) if len(nm) else 0.0,
"max_taps": int(taps.max()) if len(nm) else 0}
def iter_taps(kernel: Any) -> Iterable[Tuple[int, Tuple[int, int, int]]]:
"""``(k, (a, b, c))`` in tap order."""
k = _triple(kernel, "kernel")
for a in range(k[0]):
for b in range(k[1]):
for c in range(k[2]):
yield (a * k[1] + b) * k[2] + c, (a, b, c)