multimodalart's picture
multimodalart HF Staff
Upload folder using huggingface_hub
9550667 verified
Raw
History Blame Contribute Delete
12.2 kB
"""Mesh -> AQ3D input tensors -> instance predictions -> colored GLB.
Mirrors the official ScanNet200 validation pipeline of
https://github.com/kenomo/aq3d :
preprocessing : vertex normals (area weighted), vertex colors in [0, 1]
transforms : MeanCoord -> NormalizeColor(-1, 1) -> Copy(coord -> coord_full)
voxelisation : GridSample(grid_size=0.02, train=False) with FNV hashing and
scatter-mean pooling of color / coord / normal
superpoints : segmentator.segment_mesh(kThresh=0.01, segMinVerts=20)
post-processing: superpoint NMS (0.8) -> adaptive top-k -> mask scores ->
score / point-count thresholds
"""
import colorsys
from typing import Dict, List, Tuple
import numpy as np
import torch
import trimesh
import superpoints as spp
from labels import CLASS_COLORS, CLASS_NAMES
GRID_SIZE = 0.02
K_THRESH = 0.01
SEG_MIN_VERTS = 20
NMS_SPP_THRES = 0.8
ADAPTIVE_TOPK_RATIO = 0.99
NPOINT_THRES = 100
MAX_VERTICES = 700_000
# --------------------------------------------------------------------------- #
# mesh loading
# --------------------------------------------------------------------------- #
def load_mesh(path: str) -> trimesh.Trimesh:
obj = trimesh.load(path, process=False, force="mesh")
if isinstance(obj, trimesh.Scene):
parts = [g for g in obj.geometry.values() if isinstance(g, trimesh.Trimesh)]
if not parts:
raise ValueError("No triangle mesh found in the uploaded file.")
obj = trimesh.util.concatenate(parts)
if not isinstance(obj, trimesh.Trimesh):
raise ValueError("The uploaded file does not contain a triangle mesh.")
if obj.faces is None or len(obj.faces) == 0:
raise ValueError(
"The uploaded file is a point cloud (no triangle faces). AQ3D needs a "
"surface mesh, because its superpoints come from a mesh graph "
"segmentation. Please upload a reconstructed mesh (.ply / .obj / .glb)."
)
if len(obj.vertices) > MAX_VERTICES:
raise ValueError(
f"Mesh has {len(obj.vertices):,} vertices; please downsample it below "
f"{MAX_VERTICES:,} vertices first."
)
return obj
def mesh_vertex_colors(mesh: trimesh.Trimesh) -> np.ndarray:
"""Per-vertex RGB in [0, 1]; bakes textures down when needed."""
visual = mesh.visual
try:
if hasattr(visual, "to_color"):
visual = visual.to_color()
except Exception:
pass
colors = getattr(visual, "vertex_colors", None)
if colors is None or len(colors) != len(mesh.vertices):
return np.full((len(mesh.vertices), 3), 0.5, dtype=np.float32)
rgb = np.asarray(colors, dtype=np.float32)[:, :3] / 255.0
if not np.isfinite(rgb).all():
rgb = np.nan_to_num(rgb, nan=0.5)
return rgb
def area_weighted_vertex_normals(vertices: np.ndarray, faces: np.ndarray) -> np.ndarray:
"""``datasets/utils.py::vertex_normal`` from the AQ3D repository."""
v = vertices.astype(np.float64)
vec = np.cross(v[faces[:, 1]] - v[faces[:, 0]], v[faces[:, 2]] - v[faces[:, 0]])
length = np.sqrt((vec ** 2).sum(1, keepdims=True)) + 1.0e-8
nf = (vec / length) * (length * 0.5) # unit normal scaled by triangle area
nv = np.zeros_like(v)
idx = faces.reshape(-1)
vals = np.repeat(nf, 3, axis=0)
for a in range(3):
nv[:, a] = np.bincount(idx, weights=vals[:, a], minlength=v.shape[0])
nv /= np.sqrt((nv ** 2).sum(1, keepdims=True)) + 1.0e-8
return nv.astype(np.float32)
def orient_and_scale(vertices: np.ndarray, up_axis: str, scale: float,
auto_fit: bool) -> Tuple[np.ndarray, str, float]:
"""Bring an arbitrary mesh into the ScanNet convention: Z-up, metres."""
v = vertices.astype(np.float32).copy()
if up_axis == "Auto":
extent = v.max(0) - v.min(0)
detected = "XYZ"[int(np.argmin(extent))]
up_axis = detected
if up_axis == "Y":
v = np.stack([v[:, 0], -v[:, 2], v[:, 1]], axis=1)
elif up_axis == "X":
v = np.stack([v[:, 1], v[:, 2], v[:, 0]], axis=1)
v = v * float(scale)
applied = float(scale)
if auto_fit:
extent = v.max(0) - v.min(0)
horizontal = float(max(extent[0], extent[1]))
if horizontal > 1e-6 and not (1.5 <= horizontal <= 30.0):
factor = 8.0 / horizontal
v = v * factor
applied *= factor
return v, up_axis, applied
# --------------------------------------------------------------------------- #
# voxelisation (pointcept GridSample, test mode)
# --------------------------------------------------------------------------- #
def _fnv_hash_vec(arr: np.ndarray) -> np.ndarray:
arr = arr.astype(np.uint64, copy=True)
hashed = np.uint64(14695981039346656037) * np.ones(arr.shape[0], dtype=np.uint64)
for j in range(arr.shape[1]):
hashed *= np.uint64(1099511628211)
hashed = np.bitwise_xor(hashed, arr[:, j])
return hashed
def _scatter_mean_np(src: np.ndarray, index: np.ndarray, n: int) -> np.ndarray:
out = np.zeros((n, src.shape[1]), dtype=np.float64)
for a in range(src.shape[1]):
out[:, a] = np.bincount(index, weights=src[:, a], minlength=n)
counts = np.maximum(np.bincount(index, minlength=n), 1)
return (out / counts[:, None]).astype(np.float32)
def build_batch(vertices: np.ndarray, faces: np.ndarray, rgb01: np.ndarray,
device: torch.device) -> Tuple[Dict[str, torch.Tensor], np.ndarray]:
normals = area_weighted_vertex_normals(vertices, faces)
superpoints = np.ascontiguousarray(
spp.segment_mesh(vertices, faces, K_THRESH, SEG_MIN_VERTS))
coord = vertices.astype(np.float32) - vertices.astype(np.float32).mean(0) # MeanCoord
color = rgb01.astype(np.float32) * 2.0 - 1.0 # NormalizeColor
coord_full = coord.copy() # Copy
grid_coord = np.floor(coord / GRID_SIZE).astype(np.int64)
grid_coord -= grid_coord.min(0)
key = _fnv_hash_vec(grid_coord)
idx_sort = np.argsort(key)
key_sort = key[idx_sort]
_, inverse_sorted, count = np.unique(key_sort, return_inverse=True, return_counts=True)
inverse = np.zeros(coord.shape[0], dtype=np.int64)
inverse[idx_sort] = inverse_sorted.reshape(-1)
num_voxels = int(count.shape[0])
color_v = _scatter_mean_np(color, inverse, num_voxels)
normal_v = _scatter_mean_np(normals, inverse, num_voxels)
idx_unique = idx_sort[np.cumsum(np.insert(count, 0, 0)[:-1])]
coord_grid = grid_coord[idx_unique]
feat = np.concatenate([color_v, normal_v], axis=1)
num_sp = int(superpoints.max()) + 1
t = lambda a, d=torch.float32: torch.as_tensor(a).to(device=device, dtype=d)
batch = {
"coord_grid": t(coord_grid, torch.long),
"feat": t(feat),
"batch_indices": torch.zeros(num_voxels, dtype=torch.long, device=device),
"batched_inverse": t(inverse, torch.long),
"batched_superpoint": t(superpoints, torch.long),
"superpoint_len": torch.tensor([num_sp], dtype=torch.long, device=device),
"batched_superpoint_offset": torch.tensor([num_sp], dtype=torch.long, device=device),
"coord_full": t(coord_full),
}
return batch, superpoints
# --------------------------------------------------------------------------- #
# post-processing (src/models/base_instance_prediction.py)
# --------------------------------------------------------------------------- #
@torch.no_grad()
def decode_predictions(out: Dict, superpoints_np: np.ndarray,
num_classes: int = 198) -> Tuple[np.ndarray, np.ndarray, torch.Tensor]:
labels = out["labels"][0]
masks = out["masks"][0]
scores = torch.softmax(labels.float(), dim=-1)[:, :-1]
# superpoint-level NMS. Identical to upstream, but the pairwise union / IoU
# matrices are formed row-wise instead of all at once -- with ~25k queries the
# dense versions would be several GB each.
nms_score = scores.max(-1)[0]
mask_f = (masks > 0).float()
intersection = mask_f @ mask_f.t()
del mask_f
areas = intersection.diagonal().clone()
idxs = torch.argsort(nms_score, descending=True)
keep = []
while idxs.numel() > 0:
i = idxs[0]
keep.append(i.item())
if idxs.numel() == 1:
break
rest = idxs[1:]
inter = intersection[i, rest]
iou = inter / (areas[i] + areas[rest] - inter + 1e-6)
idxs = rest[iou < NMS_SPP_THRES]
del intersection, areas
keep = torch.tensor(keep, dtype=torch.long, device=scores.device)
masks = masks[keep]
scores = scores[keep]
# adaptive top-k over the flattened (query x class) score matrix
num_superpoints = masks.shape[-1]
topk = min(int(num_superpoints * ADAPTIVE_TOPK_RATIO), scores.numel())
flat_labels = torch.arange(num_classes, device=scores.device).unsqueeze(0)
flat_labels = flat_labels.repeat(scores.shape[0], 1).flatten(0, 1)
scores, topk_idx = scores.flatten(0, 1).topk(topk, sorted=False)
out_labels = flat_labels[topk_idx]
topk_idx = torch.div(topk_idx, num_classes, rounding_mode="floor")
masks = masks[topk_idx]
masks_binary = masks > 0
mask_scores = ((masks.sigmoid() * masks_binary).sum(1)
/ (masks_binary.sum(1) + 1e-6))
scores = scores * mask_scores
masks_binary = masks_binary.cpu()
scores = scores.cpu()
out_labels = out_labels.cpu()
sp = torch.from_numpy(superpoints_np)
spp_sizes = torch.bincount(sp, minlength=masks_binary.shape[1]).float()
npoints = (masks_binary.float() * spp_sizes).sum(1)
keep2 = npoints > NPOINT_THRES
scores, out_labels, masks_binary = scores[keep2], out_labels[keep2], masks_binary[keep2]
npoints = npoints[keep2]
order = torch.argsort(scores, descending=True)
return (out_labels[order].numpy(), scores[order].numpy(),
masks_binary[order], npoints[order].numpy())
# --------------------------------------------------------------------------- #
# visualisation
# --------------------------------------------------------------------------- #
def _instance_color(class_idx: int, nth: int) -> Tuple[int, int, int]:
"""Class color from SCANNET_COLOR_MAP_200, lightened per repeated instance."""
base = np.array(CLASS_COLORS[class_idx], dtype=np.float32) / 255.0
h, l, s = colorsys.rgb_to_hls(*base.tolist())
l = float(np.clip(l + ((nth % 4) - 1.5) * 0.13, 0.22, 0.85))
s = float(np.clip(s + ((nth % 3) - 1) * 0.10, 0.35, 1.0))
r, g, b = colorsys.hls_to_rgb(h, l, s)
return int(r * 255), int(g * 255), int(b * 255)
def colorize(mesh_vertices: np.ndarray, faces: np.ndarray, superpoints: np.ndarray,
labels: np.ndarray, scores: np.ndarray, masks_binary: torch.Tensor,
npoints: np.ndarray, threshold: float, max_instances: int
) -> Tuple[trimesh.Trimesh, List[List]]:
keep = np.where(scores >= threshold)[0][:max_instances]
colors = np.full((mesh_vertices.shape[0], 4), 205, dtype=np.uint8)
colors[:, 3] = 255
# one color per kept instance, walking in descending-score order
per_class_count: Dict[int, int] = {}
assigned: List[Tuple[int, int, int]] = []
rows: List[List] = []
for rank, i in enumerate(keep):
cls = int(labels[i])
nth = per_class_count.get(cls, 0)
per_class_count[cls] = nth + 1
rgb = _instance_color(cls, nth)
assigned.append(rgb)
rows.append([rank + 1, CLASS_NAMES[cls], round(float(scores[i]), 3),
int(npoints[i]), "#{:02x}{:02x}{:02x}".format(*rgb)])
# paint low -> high score so the most confident instance wins overlaps
for rank in reversed(range(len(keep))):
sel = masks_binary[keep[rank]].numpy()[superpoints]
colors[sel, :3] = assigned[rank]
# ScanNet is Z-up; glTF viewers are Y-up
v = mesh_vertices
display = np.stack([v[:, 0], v[:, 2], -v[:, 1]], axis=1)
out_mesh = trimesh.Trimesh(vertices=display, faces=faces, vertex_colors=colors,
process=False)
return out_mesh, rows