Spaces:
Running on Zero
Running on Zero
Upload 56 files
Browse filesThis view is limited to 50 files because it contains too many changes. See raw diff
- .gitattributes +4 -0
- app.py +706 -0
- assets/example_mesh/crocodile.glb +3 -0
- assets/example_mesh/dragon.glb +3 -0
- assets/example_mesh/spaceman.glb +3 -0
- assets/teaser.png +3 -0
- dataset/mesh_render.py +265 -0
- dataset/topo_dataset.py +88 -0
- dataset/utils.py +283 -0
- dataset/voxel_dataset.py +211 -0
- models/__init__.py +9 -0
- models/dino_encoder.py +73 -0
- models/flow_sampler.py +58 -0
- models/offset_head.py +22 -0
- models/topo_autoencoder.py +297 -0
- models/topo_flow.py +400 -0
- models/vdf_encoder.py +55 -0
- models/vertex_autoencoder.py +595 -0
- models/vertex_structured_flow.py +147 -0
- models/voxel_encoder.py +183 -0
- modules/attention.py +160 -0
- modules/norm.py +41 -0
- modules/pointnet.py +330 -0
- modules/sparse/__init__.py +130 -0
- modules/sparse/attention/__init__.py +27 -0
- modules/sparse/attention/full_attn.py +238 -0
- modules/sparse/attention/modules.py +214 -0
- modules/sparse/attention/serialized_attn.py +217 -0
- modules/sparse/attention/windowed_attn.py +158 -0
- modules/sparse/basic.py +482 -0
- modules/sparse/blocks.py +71 -0
- modules/sparse/conv/__init__.py +44 -0
- modules/sparse/conv/conv_spconv.py +107 -0
- modules/sparse/conv/conv_torchsparse.py +60 -0
- modules/sparse/linear.py +38 -0
- modules/sparse/nonlinearity.py +58 -0
- modules/sparse/norm.py +81 -0
- modules/sparse/spatial.py +158 -0
- modules/sparse/transformer/__init__.py +26 -0
- modules/sparse/transformer/bases.py +234 -0
- modules/sparse/transformer/blocks.py +165 -0
- modules/sparse/transformer/modulated.py +119 -0
- modules/transformer/__init__.py +24 -0
- modules/transformer/blocks.py +276 -0
- modules/transformer/hybrid.py +236 -0
- modules/utils.py +145 -0
- requirements.txt +23 -0
- scripts/ckpt_download.py +90 -0
- scripts/e2e_inference.py +380 -0
- scripts/tflow_inference.py +221 -0
.gitattributes
CHANGED
|
@@ -41,3 +41,7 @@ examples/example-04.jpg filter=lfs diff=lfs merge=lfs -text
|
|
| 41 |
examples/example-05.pdf filter=lfs diff=lfs merge=lfs -text
|
| 42 |
examples/2.jpg filter=lfs diff=lfs merge=lfs -text
|
| 43 |
examples/4.jpg filter=lfs diff=lfs merge=lfs -text
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 41 |
examples/example-05.pdf filter=lfs diff=lfs merge=lfs -text
|
| 42 |
examples/2.jpg filter=lfs diff=lfs merge=lfs -text
|
| 43 |
examples/4.jpg filter=lfs diff=lfs merge=lfs -text
|
| 44 |
+
assets/example_mesh/crocodile.glb filter=lfs diff=lfs merge=lfs -text
|
| 45 |
+
assets/example_mesh/dragon.glb filter=lfs diff=lfs merge=lfs -text
|
| 46 |
+
assets/example_mesh/spaceman.glb filter=lfs diff=lfs merge=lfs -text
|
| 47 |
+
assets/teaser.png filter=lfs diff=lfs merge=lfs -text
|
app.py
ADDED
|
@@ -0,0 +1,706 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""
|
| 2 |
+
LATO.2 Gradio App — Image-to-3D Mesh Generation
|
| 3 |
+
=================================================
|
| 4 |
+
Factorized 3D Mesh Generation with Vertex and Topology Flow.
|
| 5 |
+
|
| 6 |
+
Launches a Gradio interface that accepts an input image (or mesh),
|
| 7 |
+
runs the full V-Flow → T-Flow pipeline, and displays the result
|
| 8 |
+
with Rerun 3D viewer + GLB download.
|
| 9 |
+
|
| 10 |
+
Usage:
|
| 11 |
+
python app.py [--share] [--port 7860]
|
| 12 |
+
"""
|
| 13 |
+
|
| 14 |
+
import argparse
|
| 15 |
+
import os
|
| 16 |
+
import sys
|
| 17 |
+
import tempfile
|
| 18 |
+
import time
|
| 19 |
+
import uuid
|
| 20 |
+
from pathlib import Path
|
| 21 |
+
|
| 22 |
+
# ── project root on sys.path ──────────────────────────────────────────────────
|
| 23 |
+
ROOT = os.path.dirname(os.path.abspath(__file__))
|
| 24 |
+
sys.path.insert(0, ROOT)
|
| 25 |
+
os.environ.setdefault("PYTORCH_CUDA_ALLOC_CONF", "expandable_segments:True")
|
| 26 |
+
os.environ.setdefault("XDG_RUNTIME_DIR", os.path.join(tempfile.gettempdir(), "runtime-root"))
|
| 27 |
+
os.makedirs(os.environ["XDG_RUNTIME_DIR"], exist_ok=True)
|
| 28 |
+
os.environ.setdefault("EGL_PLATFORM", "surfaceless")
|
| 29 |
+
|
| 30 |
+
import gradio as gr
|
| 31 |
+
import numpy as np
|
| 32 |
+
import rerun as rr
|
| 33 |
+
import torch
|
| 34 |
+
import trimesh
|
| 35 |
+
from gradio_rerun import Rerun
|
| 36 |
+
from PIL import Image
|
| 37 |
+
|
| 38 |
+
from dataset.utils import (
|
| 39 |
+
MESH_EXTENSIONS,
|
| 40 |
+
dedup_quantized_mesh,
|
| 41 |
+
extract_active_voxels,
|
| 42 |
+
quantize_mesh_clustering,
|
| 43 |
+
)
|
| 44 |
+
from models import (
|
| 45 |
+
DinoV2Encoder,
|
| 46 |
+
OffsetHead,
|
| 47 |
+
TopoFlowEulerSampler,
|
| 48 |
+
TopologySiTFlow,
|
| 49 |
+
TopologyVAE,
|
| 50 |
+
VertexSLatFlowModel,
|
| 51 |
+
VertFlowEulerCfgSampler,
|
| 52 |
+
VertexVAE,
|
| 53 |
+
VoxelFieldConditioner,
|
| 54 |
+
)
|
| 55 |
+
from modules.sparse import SparseTensor
|
| 56 |
+
import utils.logging as logging
|
| 57 |
+
from utils.inference import (
|
| 58 |
+
build_voxel_fields,
|
| 59 |
+
decode_vertices,
|
| 60 |
+
edges_to_faces,
|
| 61 |
+
pad_verts,
|
| 62 |
+
)
|
| 63 |
+
from utils.load import load_latov2_model
|
| 64 |
+
|
| 65 |
+
|
| 66 |
+
# ═══════════════════════════════════════════════════════════════════════════════
|
| 67 |
+
# Global model state (lazy-loaded once on first inference)
|
| 68 |
+
# ═══════════════════════════════════════════════════════════════════════════════
|
| 69 |
+
_models = {}
|
| 70 |
+
_configs = {}
|
| 71 |
+
_device = None
|
| 72 |
+
|
| 73 |
+
OUTPUT_DIR = os.path.join(ROOT, "gradio_outputs")
|
| 74 |
+
os.makedirs(OUTPUT_DIR, exist_ok=True)
|
| 75 |
+
|
| 76 |
+
|
| 77 |
+
def _load_models():
|
| 78 |
+
"""Load all LATO.2 sub-models once (idempotent)."""
|
| 79 |
+
global _models, _configs, _device
|
| 80 |
+
|
| 81 |
+
if _models:
|
| 82 |
+
return # already loaded
|
| 83 |
+
|
| 84 |
+
_device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
| 85 |
+
device = _device
|
| 86 |
+
ckpt = os.path.join(ROOT, "ckpt")
|
| 87 |
+
|
| 88 |
+
logging.info("Loading LATO.2 models …")
|
| 89 |
+
|
| 90 |
+
# Stage 1 — vertex generation
|
| 91 |
+
vflow, vflow_cfg = load_latov2_model(VertexSLatFlowModel, os.path.join(ckpt, "vflow.pt"), device)
|
| 92 |
+
vvae, vvae_cfg = load_latov2_model(VertexVAE, os.path.join(ckpt, "vvae.pt"), device)
|
| 93 |
+
offset_head, _ = load_latov2_model(OffsetHead, os.path.join(ckpt, "offset_head.pt"), device)
|
| 94 |
+
|
| 95 |
+
# Stage 2 — topology generation
|
| 96 |
+
tflow, tflow_cfg = load_latov2_model(TopologySiTFlow, os.path.join(ckpt, "tflow.pt"), device)
|
| 97 |
+
tvae, _ = load_latov2_model(TopologyVAE, os.path.join(ckpt, "tvae.pt"), device)
|
| 98 |
+
voxel_encoder, venc_cfg = load_latov2_model(VoxelFieldConditioner, os.path.join(ckpt, "voxel_encoder.pt"), device)
|
| 99 |
+
|
| 100 |
+
# DINO-v2 image encoder
|
| 101 |
+
dino = (
|
| 102 |
+
DinoV2Encoder(
|
| 103 |
+
model_name=vflow_cfg["dino_version"],
|
| 104 |
+
hub_dir=os.path.join(ckpt, "dinov2"),
|
| 105 |
+
img_res=vflow_cfg["image_resolution"],
|
| 106 |
+
)
|
| 107 |
+
.to(device)
|
| 108 |
+
.eval()
|
| 109 |
+
)
|
| 110 |
+
|
| 111 |
+
_models = dict(
|
| 112 |
+
vflow=vflow, vvae=vvae, offset_head=offset_head,
|
| 113 |
+
tflow=tflow, tvae=tvae, voxel_encoder=voxel_encoder,
|
| 114 |
+
dino=dino,
|
| 115 |
+
vertex_sampler=VertFlowEulerCfgSampler(),
|
| 116 |
+
topo_sampler=TopoFlowEulerSampler(),
|
| 117 |
+
)
|
| 118 |
+
_configs = dict(vflow=vflow_cfg, vvae=vvae_cfg, tflow=tflow_cfg, venc=venc_cfg)
|
| 119 |
+
logging.info("All models loaded ✓")
|
| 120 |
+
|
| 121 |
+
|
| 122 |
+
# ═══════════════════════════════════════════════════════════════════════════════
|
| 123 |
+
# Helper: prepare input mesh → active voxels + quantized vertices
|
| 124 |
+
# ═══════════════════════════════════════════════════════════════════════════════
|
| 125 |
+
|
| 126 |
+
def _prepare_mesh_data(mesh_path: str, resolution: int, min_resolution: int):
|
| 127 |
+
"""Quantize and extract voxels from a mesh file for the pipeline."""
|
| 128 |
+
quantized = quantize_mesh_clustering(mesh_path, resolution=resolution)
|
| 129 |
+
if quantized is None:
|
| 130 |
+
raise ValueError("Input mesh is empty or could not be loaded.")
|
| 131 |
+
v_int, offsets, faces = quantized
|
| 132 |
+
if len(faces) < 1 or len(v_int) < 3:
|
| 133 |
+
raise ValueError("Mesh is degenerate after quantization.")
|
| 134 |
+
|
| 135 |
+
gt_int, gt_offsets, gt_faces = dedup_quantized_mesh(v_int, offsets, faces, resolution)
|
| 136 |
+
if len(gt_int) < 3 or len(gt_faces) < 1:
|
| 137 |
+
raise ValueError("Too few vertices/faces after deduplication.")
|
| 138 |
+
|
| 139 |
+
quant_v = gt_int.astype(np.float64) / (resolution - 1.0) - 0.5
|
| 140 |
+
quant_v = np.clip(quant_v, -0.5 + 1e-6, 0.5 - 1e-6).astype(np.float32)
|
| 141 |
+
|
| 142 |
+
min_active = extract_active_voxels(quant_v, gt_faces, min_resolution)
|
| 143 |
+
return gt_int, gt_offsets, gt_faces, quant_v, min_active
|
| 144 |
+
|
| 145 |
+
|
| 146 |
+
def _render_mesh_to_image(mesh_path: str, resolution: int, img_res: int = 518,
|
| 147 |
+
azimuth: float = 45.0, elevation: float = 30.0):
|
| 148 |
+
"""Render a conditioning view from a mesh (used when no user image is provided)."""
|
| 149 |
+
quantized = quantize_mesh_clustering(mesh_path, resolution=resolution)
|
| 150 |
+
if quantized is None:
|
| 151 |
+
return None
|
| 152 |
+
v_int, _, faces = quantized
|
| 153 |
+
render_v = v_int.astype(np.float64) / resolution - 0.5
|
| 154 |
+
|
| 155 |
+
from dataset.mesh_render import WhiteModelRenderer
|
| 156 |
+
renderer = WhiteModelRenderer(
|
| 157 |
+
img_res=img_res,
|
| 158 |
+
mesh_color=(0.78, 0.78, 0.82),
|
| 159 |
+
bg_color=(0.0, 0.0, 0.0),
|
| 160 |
+
up_axis="y",
|
| 161 |
+
add_ground=False,
|
| 162 |
+
shadow=True,
|
| 163 |
+
crop_to_object=True,
|
| 164 |
+
crop_padding=1.2,
|
| 165 |
+
)
|
| 166 |
+
imgs, _ = renderer.render(
|
| 167 |
+
np.asarray(render_v, dtype=np.float64),
|
| 168 |
+
np.asarray(faces, dtype=np.int64),
|
| 169 |
+
num_views=1,
|
| 170 |
+
azimuths=[azimuth],
|
| 171 |
+
elevations=[elevation],
|
| 172 |
+
)
|
| 173 |
+
return imgs[0] # (H, W, 3) uint8
|
| 174 |
+
|
| 175 |
+
|
| 176 |
+
# ═══════════════════════════════════════════════════════════════════════════════
|
| 177 |
+
# Core generation pipeline
|
| 178 |
+
# ═══════════════════════════════════════════════════════════════════════════════
|
| 179 |
+
|
| 180 |
+
def generate_mesh(
|
| 181 |
+
input_image: np.ndarray | None,
|
| 182 |
+
input_mesh_path: str | None,
|
| 183 |
+
vert_num: int,
|
| 184 |
+
cfg_strength: float,
|
| 185 |
+
vflow_steps: int,
|
| 186 |
+
tflow_steps: int,
|
| 187 |
+
seed: int,
|
| 188 |
+
progress=gr.Progress(track_tqdm=True),
|
| 189 |
+
):
|
| 190 |
+
"""
|
| 191 |
+
Main generation function.
|
| 192 |
+
- input_image: user-uploaded image (H, W, 3) uint8 — used as DINOv2 conditioning
|
| 193 |
+
- input_mesh_path: reference mesh file — provides the voxel scaffold
|
| 194 |
+
If only an image is supplied, the user must also supply a reference mesh for
|
| 195 |
+
the voxel scaffold (or we use one of the bundled examples).
|
| 196 |
+
"""
|
| 197 |
+
_load_models() # ensure models are loaded
|
| 198 |
+
|
| 199 |
+
device = _device
|
| 200 |
+
m = _models
|
| 201 |
+
c = _configs
|
| 202 |
+
|
| 203 |
+
torch.manual_seed(seed)
|
| 204 |
+
np.random.seed(seed)
|
| 205 |
+
|
| 206 |
+
res = c["vvae"]["resolution"]
|
| 207 |
+
min_res = c["vvae"]["min_resolution"]
|
| 208 |
+
latent_dim = c["vflow"]["latent_dim"]
|
| 209 |
+
density_max = c["vflow"]["max_vertex_num"]
|
| 210 |
+
z_dim = int(c["tflow"]["args"]["z_dim"])
|
| 211 |
+
max_vertices = int(c["tflow"]["args"]["max_vertices"])
|
| 212 |
+
latent_scale = float(c["tflow"]["latent_scale"])
|
| 213 |
+
voxel_res = int(c["venc"]["resolution"])
|
| 214 |
+
inference_threshold = 0.5
|
| 215 |
+
|
| 216 |
+
run_id = str(uuid.uuid4())[:8]
|
| 217 |
+
|
| 218 |
+
# ── Resolve mesh scaffold ─────────────────────────────────────────────
|
| 219 |
+
if input_mesh_path is None or not os.path.isfile(input_mesh_path):
|
| 220 |
+
raise gr.Error(
|
| 221 |
+
"A reference mesh file is required to provide the voxel scaffold. "
|
| 222 |
+
"Please upload a .glb / .obj / .ply / .stl mesh."
|
| 223 |
+
)
|
| 224 |
+
|
| 225 |
+
progress(0.05, desc="Quantizing mesh & extracting voxels …")
|
| 226 |
+
gt_int, gt_offsets, gt_faces, quant_v, min_active = _prepare_mesh_data(
|
| 227 |
+
input_mesh_path, res, min_res
|
| 228 |
+
)
|
| 229 |
+
|
| 230 |
+
# ── Resolve conditioning image ────────────────────────────────────────
|
| 231 |
+
if input_image is not None:
|
| 232 |
+
cond_img = np.asarray(input_image, dtype=np.uint8)
|
| 233 |
+
if cond_img.ndim == 2:
|
| 234 |
+
cond_img = np.stack([cond_img] * 3, axis=-1)
|
| 235 |
+
elif cond_img.shape[-1] == 4:
|
| 236 |
+
cond_img = cond_img[:, :, :3]
|
| 237 |
+
else:
|
| 238 |
+
progress(0.08, desc="Rendering conditioning view from mesh …")
|
| 239 |
+
cond_img = _render_mesh_to_image(input_mesh_path, res)
|
| 240 |
+
if cond_img is None:
|
| 241 |
+
raise gr.Error("Could not render a conditioning view from the mesh.")
|
| 242 |
+
|
| 243 |
+
# ── Compute density conditioning ──────────────────────────────────────
|
| 244 |
+
clamped = float(min(max(vert_num, 200), 5000))
|
| 245 |
+
density = torch.tensor([clamped], dtype=torch.float32, device=device)
|
| 246 |
+
density = density / density_max * 1000.0
|
| 247 |
+
|
| 248 |
+
# ── Stage 1: V-Flow �� V-VAE ──────────────────────────────────────────
|
| 249 |
+
progress(0.12, desc="Running V-Flow (vertex generation) …")
|
| 250 |
+
min_active_batched = torch.cat(
|
| 251 |
+
[torch.zeros(min_active.shape[0], 1, dtype=torch.int32), min_active], dim=1
|
| 252 |
+
)
|
| 253 |
+
|
| 254 |
+
with torch.no_grad():
|
| 255 |
+
cond = m["dino"](cond_img).float()
|
| 256 |
+
neg_cond = torch.zeros_like(cond)
|
| 257 |
+
|
| 258 |
+
with torch.no_grad(), torch.autocast("cuda", dtype=torch.bfloat16):
|
| 259 |
+
min_active_coords = min_active_batched.to(device)
|
| 260 |
+
noise = SparseTensor(
|
| 261 |
+
coords=min_active_coords.int(),
|
| 262 |
+
feats=torch.randn(
|
| 263 |
+
min_active_coords.shape[0], latent_dim, device=device
|
| 264 |
+
),
|
| 265 |
+
)
|
| 266 |
+
z_pred = m["vertex_sampler"].sample(
|
| 267 |
+
model=m["vflow"],
|
| 268 |
+
noise=noise,
|
| 269 |
+
cond=cond,
|
| 270 |
+
neg_cond=neg_cond,
|
| 271 |
+
steps=vflow_steps,
|
| 272 |
+
cfg_strength=cfg_strength,
|
| 273 |
+
rescale_t=1.0,
|
| 274 |
+
density=density,
|
| 275 |
+
)
|
| 276 |
+
pred_coords, pred_offsets = decode_vertices(
|
| 277 |
+
m["vvae"], m["offset_head"], z_pred, inference_threshold
|
| 278 |
+
)
|
| 279 |
+
|
| 280 |
+
progress(0.55, desc="Decoding vertices …")
|
| 281 |
+
|
| 282 |
+
# Filter to batch 0
|
| 283 |
+
pred_sel = pred_coords[:, 0] == 0
|
| 284 |
+
vert_int = pred_coords[pred_sel, 1:].long()
|
| 285 |
+
vert_off = pred_offsets[pred_sel]
|
| 286 |
+
num_pred = int(vert_int.shape[0])
|
| 287 |
+
|
| 288 |
+
if num_pred < 3:
|
| 289 |
+
raise gr.Error(f"Only {num_pred} vertices generated. Try different parameters.")
|
| 290 |
+
if num_pred > max_vertices:
|
| 291 |
+
raise gr.Error(
|
| 292 |
+
f"Generated {num_pred} vertices exceeds T-Flow max ({max_vertices}). "
|
| 293 |
+
"Try reducing the vertex count."
|
| 294 |
+
)
|
| 295 |
+
|
| 296 |
+
# ── Stage 2: T-Flow → T-VAE ──────────────────────────────────────────
|
| 297 |
+
progress(0.60, desc="Running T-Flow (topology generation) …")
|
| 298 |
+
|
| 299 |
+
with torch.no_grad():
|
| 300 |
+
verts, mask, lengths = pad_verts([vert_int], device)
|
| 301 |
+
voxel_list = [min_active.long()]
|
| 302 |
+
field = build_voxel_fields(voxel_list, voxel_res, device)
|
| 303 |
+
cond_vox = m["voxel_encoder"](field)
|
| 304 |
+
|
| 305 |
+
z0 = torch.randn(verts.shape[0], verts.shape[1], z_dim, device=device)
|
| 306 |
+
z_flow = m["topo_sampler"].sample(
|
| 307 |
+
model=m["tflow"],
|
| 308 |
+
noise=z0,
|
| 309 |
+
verts=verts,
|
| 310 |
+
mask=mask,
|
| 311 |
+
cond=cond_vox,
|
| 312 |
+
steps=tflow_steps,
|
| 313 |
+
)
|
| 314 |
+
z = z_flow.float() / latent_scale
|
| 315 |
+
|
| 316 |
+
with torch.autocast("cuda", dtype=torch.bfloat16):
|
| 317 |
+
edges_list = m["tvae"].decode(
|
| 318 |
+
z,
|
| 319 |
+
verts=verts,
|
| 320 |
+
verts_mask=mask,
|
| 321 |
+
chunk_size=20000,
|
| 322 |
+
threshold=0.0,
|
| 323 |
+
)
|
| 324 |
+
|
| 325 |
+
progress(0.85, desc="Assembling faces & exporting …")
|
| 326 |
+
faces = edges_to_faces(edges_list[0], lengths[0], fill_quad_rings=True)
|
| 327 |
+
|
| 328 |
+
if faces.shape[0] == 0:
|
| 329 |
+
raise gr.Error("No faces were generated. Try different parameters or a different input.")
|
| 330 |
+
|
| 331 |
+
# ── Build final mesh ──────────────────────────────────────────────────
|
| 332 |
+
vert_np = vert_int.numpy()
|
| 333 |
+
off_np = vert_off.numpy()
|
| 334 |
+
vert_with_offset = (
|
| 335 |
+
vert_np.astype(np.float64) / res
|
| 336 |
+
- 0.5
|
| 337 |
+
+ off_np.astype(np.float64) / (res * 2.0)
|
| 338 |
+
)
|
| 339 |
+
mesh = trimesh.Trimesh(vertices=vert_with_offset, faces=faces, process=False)
|
| 340 |
+
|
| 341 |
+
# ── Export to GLB ─────────────────────────────────────────────────────
|
| 342 |
+
glb_filename = f"lato2_{run_id}.glb"
|
| 343 |
+
glb_path = os.path.join(OUTPUT_DIR, glb_filename)
|
| 344 |
+
mesh.export(glb_path, file_type="glb")
|
| 345 |
+
|
| 346 |
+
# ── Build Rerun visualization ─────────────────────────────────────────
|
| 347 |
+
progress(0.92, desc="Building 3D visualization …")
|
| 348 |
+
rr_data = _build_rerun_stream(mesh, cond_img, run_id)
|
| 349 |
+
|
| 350 |
+
progress(1.0, desc="Done ✓")
|
| 351 |
+
return rr_data, glb_path
|
| 352 |
+
|
| 353 |
+
|
| 354 |
+
# ═══════════════════════════════════════════════════════════════════════════════
|
| 355 |
+
# Rerun 3D Visualization
|
| 356 |
+
# ═══════════════════════════════════════════════════════════════════════════════
|
| 357 |
+
|
| 358 |
+
def _build_rerun_stream(mesh: trimesh.Trimesh, cond_img: np.ndarray, run_id: str):
|
| 359 |
+
"""Create an .rrd byte stream with the generated mesh + conditioning image."""
|
| 360 |
+
rrd_path = os.path.join(OUTPUT_DIR, f"lato2_{run_id}.rrd")
|
| 361 |
+
|
| 362 |
+
rr.init("LATO.2 — 3D Mesh Generation", spawn=False)
|
| 363 |
+
rec = rr.new_recording(application_id="LATO.2", recording_id=run_id)
|
| 364 |
+
|
| 365 |
+
vertices = np.asarray(mesh.vertices, dtype=np.float32)
|
| 366 |
+
faces = np.asarray(mesh.faces, dtype=np.uint32)
|
| 367 |
+
|
| 368 |
+
# Compute vertex normals for nicer shading
|
| 369 |
+
if mesh.vertex_normals is not None and len(mesh.vertex_normals) > 0:
|
| 370 |
+
normals = np.asarray(mesh.vertex_normals, dtype=np.float32)
|
| 371 |
+
else:
|
| 372 |
+
normals = None
|
| 373 |
+
|
| 374 |
+
# Log the generated mesh
|
| 375 |
+
rec.log(
|
| 376 |
+
"world/generated_mesh",
|
| 377 |
+
rr.Mesh3D(
|
| 378 |
+
vertex_positions=vertices,
|
| 379 |
+
triangle_indices=faces,
|
| 380 |
+
vertex_normals=normals,
|
| 381 |
+
),
|
| 382 |
+
)
|
| 383 |
+
|
| 384 |
+
# Log the conditioning image
|
| 385 |
+
if cond_img is not None:
|
| 386 |
+
rec.log("conditioning_image", rr.Image(cond_img))
|
| 387 |
+
|
| 388 |
+
# Log mesh stats as text
|
| 389 |
+
rec.log(
|
| 390 |
+
"world/stats",
|
| 391 |
+
rr.TextDocument(
|
| 392 |
+
f"Vertices: {len(vertices)}\n"
|
| 393 |
+
f"Faces: {len(faces)}\n"
|
| 394 |
+
f"Bounding box: {vertices.min(axis=0).tolist()} → {vertices.max(axis=0).tolist()}"
|
| 395 |
+
),
|
| 396 |
+
)
|
| 397 |
+
|
| 398 |
+
rrd_bytes = rec.memory_recording()
|
| 399 |
+
return rrd_bytes
|
| 400 |
+
|
| 401 |
+
|
| 402 |
+
# ═══════════════════════════════════════════════════════════════════════════════
|
| 403 |
+
# Gradio UI
|
| 404 |
+
# ═══════════════════════════════════════════════════════════════════════════════
|
| 405 |
+
|
| 406 |
+
TITLE = "LATO.2: Factorized 3D Mesh Generation"
|
| 407 |
+
DESCRIPTION = """
|
| 408 |
+
**LATO.2** factorizes mesh generation into a **Vertex Flow (V-Flow)** for vertex positions
|
| 409 |
+
and a **Topology Flow (T-Flow)** for connectivity prediction.
|
| 410 |
+
|
| 411 |
+
### How to use
|
| 412 |
+
1. **Upload a conditioning image** — this drives the DINOv2 shape conditioning.
|
| 413 |
+
2. **Upload a reference mesh** (.glb / .obj / .ply / .stl) — this provides the coarse voxel scaffold.
|
| 414 |
+
*If you skip the image, a rendered view of the mesh is used as conditioning instead.*
|
| 415 |
+
3. Adjust parameters and click **🚀 Generate 3D Mesh**.
|
| 416 |
+
4. Explore the result in the **Rerun 3D Viewer** and download the **.glb** file.
|
| 417 |
+
"""
|
| 418 |
+
|
| 419 |
+
EXAMPLES_DIR = os.path.join(ROOT, "assets", "example_mesh")
|
| 420 |
+
|
| 421 |
+
CSS = """
|
| 422 |
+
/* ── Dark premium theme overrides ──────────────────────────────────── */
|
| 423 |
+
.gradio-container {
|
| 424 |
+
max-width: 1400px !important;
|
| 425 |
+
margin: auto;
|
| 426 |
+
}
|
| 427 |
+
|
| 428 |
+
#app-title {
|
| 429 |
+
text-align: center;
|
| 430 |
+
background: linear-gradient(135deg, #667eea 0%, #764ba2 100%);
|
| 431 |
+
-webkit-background-clip: text;
|
| 432 |
+
-webkit-text-fill-color: transparent;
|
| 433 |
+
background-clip: text;
|
| 434 |
+
font-size: 2.4rem;
|
| 435 |
+
font-weight: 800;
|
| 436 |
+
letter-spacing: -0.02em;
|
| 437 |
+
margin-bottom: 0.2em;
|
| 438 |
+
font-family: 'Inter', 'Segoe UI', sans-serif;
|
| 439 |
+
}
|
| 440 |
+
|
| 441 |
+
#app-subtitle {
|
| 442 |
+
text-align: center;
|
| 443 |
+
color: #9ca3af;
|
| 444 |
+
font-size: 1.05rem;
|
| 445 |
+
margin-top: -0.5em;
|
| 446 |
+
margin-bottom: 1.5em;
|
| 447 |
+
}
|
| 448 |
+
|
| 449 |
+
.generate-btn {
|
| 450 |
+
background: linear-gradient(135deg, #667eea 0%, #764ba2 100%) !important;
|
| 451 |
+
border: none !important;
|
| 452 |
+
color: white !important;
|
| 453 |
+
font-weight: 700 !important;
|
| 454 |
+
font-size: 1.15rem !important;
|
| 455 |
+
padding: 14px 32px !important;
|
| 456 |
+
border-radius: 12px !important;
|
| 457 |
+
box-shadow: 0 4px 15px rgba(102, 126, 234, 0.4) !important;
|
| 458 |
+
transition: all 0.3s ease !important;
|
| 459 |
+
letter-spacing: 0.02em !important;
|
| 460 |
+
}
|
| 461 |
+
|
| 462 |
+
.generate-btn:hover {
|
| 463 |
+
transform: translateY(-2px) !important;
|
| 464 |
+
box-shadow: 0 8px 25px rgba(102, 126, 234, 0.55) !important;
|
| 465 |
+
}
|
| 466 |
+
|
| 467 |
+
.param-accordion {
|
| 468 |
+
border: 1px solid rgba(102, 126, 234, 0.2) !important;
|
| 469 |
+
border-radius: 12px !important;
|
| 470 |
+
margin-top: 8px !important;
|
| 471 |
+
}
|
| 472 |
+
|
| 473 |
+
.download-btn {
|
| 474 |
+
background: linear-gradient(135deg, #11998e 0%, #38ef7d 100%) !important;
|
| 475 |
+
border: none !important;
|
| 476 |
+
color: white !important;
|
| 477 |
+
font-weight: 700 !important;
|
| 478 |
+
font-size: 1.05rem !important;
|
| 479 |
+
padding: 12px 28px !important;
|
| 480 |
+
border-radius: 12px !important;
|
| 481 |
+
box-shadow: 0 4px 15px rgba(17, 153, 142, 0.35) !important;
|
| 482 |
+
transition: all 0.3s ease !important;
|
| 483 |
+
}
|
| 484 |
+
|
| 485 |
+
.download-btn:hover {
|
| 486 |
+
transform: translateY(-2px) !important;
|
| 487 |
+
box-shadow: 0 8px 25px rgba(17, 153, 142, 0.5) !important;
|
| 488 |
+
}
|
| 489 |
+
|
| 490 |
+
.info-badge {
|
| 491 |
+
display: inline-block;
|
| 492 |
+
background: rgba(102, 126, 234, 0.12);
|
| 493 |
+
color: #667eea;
|
| 494 |
+
padding: 4px 12px;
|
| 495 |
+
border-radius: 20px;
|
| 496 |
+
font-size: 0.85rem;
|
| 497 |
+
font-weight: 600;
|
| 498 |
+
margin: 4px 0;
|
| 499 |
+
}
|
| 500 |
+
|
| 501 |
+
footer { display: none !important; }
|
| 502 |
+
"""
|
| 503 |
+
|
| 504 |
+
def build_app():
|
| 505 |
+
"""Construct the Gradio Blocks app."""
|
| 506 |
+
theme = gr.themes.Soft(
|
| 507 |
+
primary_hue=gr.themes.colors.indigo,
|
| 508 |
+
secondary_hue=gr.themes.colors.purple,
|
| 509 |
+
neutral_hue=gr.themes.colors.gray,
|
| 510 |
+
font=gr.themes.GoogleFont("Inter"),
|
| 511 |
+
).set(
|
| 512 |
+
body_background_fill="*neutral_950",
|
| 513 |
+
body_background_fill_dark="*neutral_950",
|
| 514 |
+
block_background_fill="*neutral_900",
|
| 515 |
+
block_background_fill_dark="*neutral_900",
|
| 516 |
+
block_border_width="0px",
|
| 517 |
+
block_shadow="0 2px 12px rgba(0,0,0,0.3)",
|
| 518 |
+
input_background_fill="*neutral_800",
|
| 519 |
+
input_background_fill_dark="*neutral_800",
|
| 520 |
+
)
|
| 521 |
+
|
| 522 |
+
with gr.Blocks(
|
| 523 |
+
theme=theme,
|
| 524 |
+
css=CSS,
|
| 525 |
+
title="LATO.2 — Image to 3D Mesh",
|
| 526 |
+
) as app:
|
| 527 |
+
# ── Header ────────────────────────────────────────────────────────
|
| 528 |
+
gr.HTML(
|
| 529 |
+
'<h1 id="app-title">LATO.2</h1>'
|
| 530 |
+
'<p id="app-subtitle">Factorized 3D Mesh Generation with Vertex & Topology Flow</p>'
|
| 531 |
+
)
|
| 532 |
+
gr.Markdown(DESCRIPTION)
|
| 533 |
+
|
| 534 |
+
with gr.Row(equal_height=False):
|
| 535 |
+
# ── LEFT: Inputs ──────────────────────────────────────────────
|
| 536 |
+
with gr.Column(scale=1, min_width=380):
|
| 537 |
+
gr.Markdown("### 📥 Inputs")
|
| 538 |
+
|
| 539 |
+
input_image = gr.Image(
|
| 540 |
+
label="Conditioning Image (optional)",
|
| 541 |
+
type="numpy",
|
| 542 |
+
height=280,
|
| 543 |
+
sources=["upload", "clipboard"],
|
| 544 |
+
elem_id="input-image",
|
| 545 |
+
)
|
| 546 |
+
input_mesh = gr.File(
|
| 547 |
+
label="Reference Mesh (.glb / .obj / .ply / .stl)",
|
| 548 |
+
file_types=[".glb", ".gltf", ".obj", ".ply", ".stl", ".off"],
|
| 549 |
+
type="filepath",
|
| 550 |
+
elem_id="input-mesh",
|
| 551 |
+
)
|
| 552 |
+
|
| 553 |
+
with gr.Accordion("⚙️ Generation Parameters", open=True, elem_classes="param-accordion"):
|
| 554 |
+
vert_num = gr.Slider(
|
| 555 |
+
label="Target Vertex Count",
|
| 556 |
+
minimum=200,
|
| 557 |
+
maximum=5000,
|
| 558 |
+
value=2000,
|
| 559 |
+
step=100,
|
| 560 |
+
info="Number of vertices in the generated mesh (200–5000)",
|
| 561 |
+
)
|
| 562 |
+
cfg_strength = gr.Slider(
|
| 563 |
+
label="CFG Strength",
|
| 564 |
+
minimum=0.0,
|
| 565 |
+
maximum=10.0,
|
| 566 |
+
value=3.0,
|
| 567 |
+
step=0.5,
|
| 568 |
+
info="Classifier-free guidance strength",
|
| 569 |
+
)
|
| 570 |
+
with gr.Row():
|
| 571 |
+
vflow_steps = gr.Slider(
|
| 572 |
+
label="V-Flow Steps",
|
| 573 |
+
minimum=4,
|
| 574 |
+
maximum=64,
|
| 575 |
+
value=24,
|
| 576 |
+
step=4,
|
| 577 |
+
info="Euler steps for vertex flow",
|
| 578 |
+
)
|
| 579 |
+
tflow_steps = gr.Slider(
|
| 580 |
+
label="T-Flow Steps",
|
| 581 |
+
minimum=10,
|
| 582 |
+
maximum=100,
|
| 583 |
+
value=50,
|
| 584 |
+
step=5,
|
| 585 |
+
info="Euler steps for topology flow",
|
| 586 |
+
)
|
| 587 |
+
seed = gr.Number(
|
| 588 |
+
label="Random Seed",
|
| 589 |
+
value=42,
|
| 590 |
+
precision=0,
|
| 591 |
+
info="Seed for reproducibility",
|
| 592 |
+
)
|
| 593 |
+
|
| 594 |
+
generate_btn = gr.Button(
|
| 595 |
+
"🚀 Generate 3D Mesh",
|
| 596 |
+
variant="primary",
|
| 597 |
+
size="lg",
|
| 598 |
+
elem_classes="generate-btn",
|
| 599 |
+
elem_id="generate-btn",
|
| 600 |
+
)
|
| 601 |
+
|
| 602 |
+
# ── Example meshes ────────────────────────────────────────
|
| 603 |
+
if os.path.isdir(EXAMPLES_DIR):
|
| 604 |
+
example_files = sorted(
|
| 605 |
+
os.path.join(EXAMPLES_DIR, f)
|
| 606 |
+
for f in os.listdir(EXAMPLES_DIR)
|
| 607 |
+
if os.path.splitext(f)[1].lower() in MESH_EXTENSIONS
|
| 608 |
+
)
|
| 609 |
+
if example_files:
|
| 610 |
+
gr.Markdown("### 📂 Example Meshes")
|
| 611 |
+
gr.Examples(
|
| 612 |
+
examples=[[None, f, 2000, 3.0, 24, 50, 42] for f in example_files],
|
| 613 |
+
inputs=[input_image, input_mesh, vert_num, cfg_strength, vflow_steps, tflow_steps, seed],
|
| 614 |
+
label="Click to load an example",
|
| 615 |
+
cache_examples=False,
|
| 616 |
+
)
|
| 617 |
+
|
| 618 |
+
# ── RIGHT: Outputs ────────────────────────────────────────────
|
| 619 |
+
with gr.Column(scale=2, min_width=600):
|
| 620 |
+
gr.Markdown("### 🖼️ 3D Output — Rerun Viewer")
|
| 621 |
+
|
| 622 |
+
rerun_viewer = Rerun(
|
| 623 |
+
streaming=False,
|
| 624 |
+
height=560,
|
| 625 |
+
elem_id="rerun-viewer",
|
| 626 |
+
)
|
| 627 |
+
|
| 628 |
+
gr.Markdown("---")
|
| 629 |
+
gr.Markdown("### 📦 Download")
|
| 630 |
+
|
| 631 |
+
glb_output = gr.File(
|
| 632 |
+
label="Generated GLB File",
|
| 633 |
+
type="filepath",
|
| 634 |
+
elem_id="glb-output",
|
| 635 |
+
interactive=False,
|
| 636 |
+
)
|
| 637 |
+
|
| 638 |
+
download_btn = gr.DownloadButton(
|
| 639 |
+
label="⬇️ Download GLB File",
|
| 640 |
+
size="lg",
|
| 641 |
+
elem_classes="download-btn",
|
| 642 |
+
elem_id="download-btn",
|
| 643 |
+
visible=False,
|
| 644 |
+
)
|
| 645 |
+
|
| 646 |
+
# ── Event wiring ──────────────────────────────────────────────────
|
| 647 |
+
|
| 648 |
+
def on_generate(image, mesh_path, vn, cfg, vfs, tfs, s):
|
| 649 |
+
rr_data, glb_path = generate_mesh(
|
| 650 |
+
input_image=image,
|
| 651 |
+
input_mesh_path=mesh_path,
|
| 652 |
+
vert_num=int(vn),
|
| 653 |
+
cfg_strength=float(cfg),
|
| 654 |
+
vflow_steps=int(vfs),
|
| 655 |
+
tflow_steps=int(tfs),
|
| 656 |
+
seed=int(s),
|
| 657 |
+
)
|
| 658 |
+
return (
|
| 659 |
+
rr_data,
|
| 660 |
+
glb_path,
|
| 661 |
+
gr.update(value=glb_path, visible=True),
|
| 662 |
+
)
|
| 663 |
+
|
| 664 |
+
generate_btn.click(
|
| 665 |
+
fn=on_generate,
|
| 666 |
+
inputs=[input_image, input_mesh, vert_num, cfg_strength, vflow_steps, tflow_steps, seed],
|
| 667 |
+
outputs=[rerun_viewer, glb_output, download_btn],
|
| 668 |
+
)
|
| 669 |
+
|
| 670 |
+
# ── Footer ────────────────────────────────────────────────────────
|
| 671 |
+
gr.HTML(
|
| 672 |
+
'<div style="text-align:center; color:#6b7280; padding:20px 0 10px; font-size:0.85rem;">'
|
| 673 |
+
'🔬 LATO.2 — Factorized 3D Mesh Generation with Vertex & Topology Flow<br>'
|
| 674 |
+
'<span style="color:#9ca3af;">Hang Long, Tianhao Zhao et al. • '
|
| 675 |
+
'<a href="https://arxiv.org/abs/2607.10623" target="_blank" '
|
| 676 |
+
'style="color:#667eea; text-decoration:none;">arXiv 2607.10623</a> • '
|
| 677 |
+
'<a href="https://huggingface.co/0x4c48/LATO.2" target="_blank" '
|
| 678 |
+
'style="color:#667eea; text-decoration:none;">🤗 Model</a></span>'
|
| 679 |
+
'</div>'
|
| 680 |
+
)
|
| 681 |
+
|
| 682 |
+
return app
|
| 683 |
+
|
| 684 |
+
|
| 685 |
+
# ═══════════════════════════════════════════════════════════════════════════════
|
| 686 |
+
# Entry point
|
| 687 |
+
# ═══════════════════════════════════════════════════════════════════════════════
|
| 688 |
+
|
| 689 |
+
def parse_app_args():
|
| 690 |
+
p = argparse.ArgumentParser(description="LATO.2 Gradio App")
|
| 691 |
+
p.add_argument("--port", type=int, default=7860, help="Port to serve on")
|
| 692 |
+
p.add_argument("--share", action="store_true", help="Create a public Gradio link")
|
| 693 |
+
p.add_argument("--server_name", default="0.0.0.0", help="Server bind address")
|
| 694 |
+
return p.parse_args()
|
| 695 |
+
|
| 696 |
+
|
| 697 |
+
if __name__ == "__main__":
|
| 698 |
+
args = parse_app_args()
|
| 699 |
+
app = build_app()
|
| 700 |
+
app.queue(max_size=4)
|
| 701 |
+
app.launch(
|
| 702 |
+
server_name=args.server_name,
|
| 703 |
+
server_port=args.port,
|
| 704 |
+
share=args.share,
|
| 705 |
+
show_error=True,
|
| 706 |
+
)
|
assets/example_mesh/crocodile.glb
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:6f6e3db8a580db37a6e9145f68b3cc765143a5b081d6dc6ce6713949cbad21ca
|
| 3 |
+
size 1721020
|
assets/example_mesh/dragon.glb
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:347d0a6f76b3b6afa34b13b75de14e73ab2e45ef5fdeb17c62a7d8216d358fd2
|
| 3 |
+
size 1846248
|
assets/example_mesh/spaceman.glb
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:20a48887b1d682b91f8a19e9ec9ec83e9bf0c402ef062ccdd26085a626b2eca1
|
| 3 |
+
size 1516348
|
assets/teaser.png
ADDED
|
Git LFS Details
|
dataset/mesh_render.py
ADDED
|
@@ -0,0 +1,265 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from typing import List, Optional, Sequence, Tuple, Union
|
| 2 |
+
|
| 3 |
+
import numpy as np
|
| 4 |
+
|
| 5 |
+
try:
|
| 6 |
+
import open3d as o3d
|
| 7 |
+
except Exception as _e:
|
| 8 |
+
o3d = None
|
| 9 |
+
_OPEN3D_IMPORT_ERROR = _e
|
| 10 |
+
|
| 11 |
+
|
| 12 |
+
ColorLike = Union[Sequence[float], np.ndarray]
|
| 13 |
+
|
| 14 |
+
|
| 15 |
+
def _to_rgb01(color: ColorLike) -> Tuple[float, float, float]:
|
| 16 |
+
c = np.asarray(color, dtype=np.float64).reshape(-1)[:3]
|
| 17 |
+
if c.max() > 1.0 + 1e-6:
|
| 18 |
+
c = c / 255.0
|
| 19 |
+
return float(c[0]), float(c[1]), float(c[2])
|
| 20 |
+
|
| 21 |
+
|
| 22 |
+
def _axis_index(up_axis: str) -> int:
|
| 23 |
+
return {"x": 0, "y": 1, "z": 2}[up_axis.lower()]
|
| 24 |
+
|
| 25 |
+
|
| 26 |
+
def _orbit_eye(
|
| 27 |
+
center: np.ndarray,
|
| 28 |
+
distance: float,
|
| 29 |
+
azimuth_deg: float,
|
| 30 |
+
elevation_deg: float,
|
| 31 |
+
up_axis: str,
|
| 32 |
+
) -> Tuple[np.ndarray, np.ndarray]:
|
| 33 |
+
az = np.deg2rad(azimuth_deg)
|
| 34 |
+
el = np.deg2rad(elevation_deg)
|
| 35 |
+
ce = np.cos(el)
|
| 36 |
+
|
| 37 |
+
horiz = distance * ce
|
| 38 |
+
vert = distance * np.sin(el)
|
| 39 |
+
ai = _axis_index(up_axis)
|
| 40 |
+
offset = np.zeros(3, dtype=np.float64)
|
| 41 |
+
|
| 42 |
+
plane_axes = [i for i in range(3) if i != ai]
|
| 43 |
+
offset[plane_axes[0]] = horiz * np.cos(az)
|
| 44 |
+
offset[plane_axes[1]] = horiz * np.sin(az)
|
| 45 |
+
offset[ai] = vert
|
| 46 |
+
eye = center + offset
|
| 47 |
+
up = np.zeros(3, dtype=np.float64)
|
| 48 |
+
up[ai] = 1.0
|
| 49 |
+
return eye, up
|
| 50 |
+
|
| 51 |
+
|
| 52 |
+
class WhiteModelRenderer:
|
| 53 |
+
def __init__(
|
| 54 |
+
self,
|
| 55 |
+
img_res: int = 512,
|
| 56 |
+
mesh_color: ColorLike = (0.78, 0.78, 0.82),
|
| 57 |
+
bg_color: ColorLike = (1.0, 1.0, 1.0),
|
| 58 |
+
up_axis: str = "y",
|
| 59 |
+
add_ground: bool = True,
|
| 60 |
+
shadow: bool = True,
|
| 61 |
+
elevation_range: Tuple[float, float] = (15.0, 40.0),
|
| 62 |
+
azimuth_range: Tuple[float, float] = (0.0, 360.0),
|
| 63 |
+
camera_distance: float = 1.8,
|
| 64 |
+
fov: float = 50.0,
|
| 65 |
+
ground_color: ColorLike = (0.92, 0.92, 0.92),
|
| 66 |
+
sun_intensity: float = 90000.0,
|
| 67 |
+
ambient_intensity: float = 32000.0,
|
| 68 |
+
crop_to_object: bool = False,
|
| 69 |
+
crop_padding: float = 1.2,
|
| 70 |
+
):
|
| 71 |
+
if o3d is None:
|
| 72 |
+
raise ImportError(
|
| 73 |
+
f"open3d is required for WhiteModelRenderer but failed to import: {_OPEN3D_IMPORT_ERROR}"
|
| 74 |
+
)
|
| 75 |
+
self.img_res = int(img_res)
|
| 76 |
+
self.mesh_color = _to_rgb01(mesh_color)
|
| 77 |
+
self.bg_color = _to_rgb01(bg_color)
|
| 78 |
+
self.up_axis = up_axis.lower()
|
| 79 |
+
self.add_ground = add_ground
|
| 80 |
+
self.shadow = shadow
|
| 81 |
+
self.elevation_range = elevation_range
|
| 82 |
+
self.azimuth_range = azimuth_range
|
| 83 |
+
self.camera_distance = float(camera_distance)
|
| 84 |
+
self.fov = float(fov)
|
| 85 |
+
self.ground_color = _to_rgb01(ground_color)
|
| 86 |
+
self.sun_intensity = float(sun_intensity)
|
| 87 |
+
self.ambient_intensity = float(ambient_intensity)
|
| 88 |
+
self.crop_to_object = crop_to_object
|
| 89 |
+
self.crop_padding = float(crop_padding)
|
| 90 |
+
|
| 91 |
+
self._renderer = None
|
| 92 |
+
self._rng = np.random.default_rng()
|
| 93 |
+
|
| 94 |
+
def _ensure_renderer(self):
|
| 95 |
+
if self._renderer is None:
|
| 96 |
+
self._renderer = o3d.visualization.rendering.OffscreenRenderer(
|
| 97 |
+
self.img_res, self.img_res
|
| 98 |
+
)
|
| 99 |
+
return self._renderer
|
| 100 |
+
|
| 101 |
+
def _make_o3d_mesh(self, vertices: np.ndarray, faces: np.ndarray):
|
| 102 |
+
mesh = o3d.geometry.TriangleMesh()
|
| 103 |
+
mesh.vertices = o3d.utility.Vector3dVector(
|
| 104 |
+
np.asarray(vertices, dtype=np.float64)
|
| 105 |
+
)
|
| 106 |
+
mesh.triangles = o3d.utility.Vector3iVector(np.asarray(faces, dtype=np.int32))
|
| 107 |
+
mesh.compute_vertex_normals()
|
| 108 |
+
return mesh
|
| 109 |
+
|
| 110 |
+
def _make_ground(self, mesh_min: np.ndarray, mesh_max: np.ndarray):
|
| 111 |
+
ai = _axis_index(self.up_axis)
|
| 112 |
+
center = (mesh_min + mesh_max) / 2.0
|
| 113 |
+
extent = float(np.max(mesh_max - mesh_min))
|
| 114 |
+
size = max(extent * 6.0, 4.0)
|
| 115 |
+
|
| 116 |
+
plane_axes = [i for i in range(3) if i != ai]
|
| 117 |
+
bottom = mesh_min[ai] - extent * 0.02
|
| 118 |
+
|
| 119 |
+
corners_2d = (
|
| 120 |
+
np.array([[-0.5, -0.5], [0.5, -0.5], [0.5, 0.5], [-0.5, 0.5]]) * size
|
| 121 |
+
)
|
| 122 |
+
verts = np.zeros((4, 3), dtype=np.float64)
|
| 123 |
+
for k, (a, b) in enumerate(corners_2d):
|
| 124 |
+
verts[k, plane_axes[0]] = center[plane_axes[0]] + a
|
| 125 |
+
verts[k, plane_axes[1]] = center[plane_axes[1]] + b
|
| 126 |
+
verts[k, ai] = bottom
|
| 127 |
+
tris = np.array([[0, 1, 2], [0, 2, 3]], dtype=np.int32)
|
| 128 |
+
ground = o3d.geometry.TriangleMesh()
|
| 129 |
+
ground.vertices = o3d.utility.Vector3dVector(verts)
|
| 130 |
+
ground.triangles = o3d.utility.Vector3iVector(tris)
|
| 131 |
+
ground.compute_vertex_normals()
|
| 132 |
+
return ground
|
| 133 |
+
|
| 134 |
+
def _lit_material(self, rgb: Tuple[float, float, float], roughness: float = 0.85):
|
| 135 |
+
mat = o3d.visualization.rendering.MaterialRecord()
|
| 136 |
+
mat.shader = "defaultLit"
|
| 137 |
+
mat.base_color = [rgb[0], rgb[1], rgb[2], 1.0]
|
| 138 |
+
mat.base_roughness = roughness
|
| 139 |
+
mat.base_metallic = 0.0
|
| 140 |
+
mat.base_reflectance = 0.4
|
| 141 |
+
return mat
|
| 142 |
+
|
| 143 |
+
def _setup_scene(
|
| 144 |
+
self,
|
| 145 |
+
vertices: np.ndarray,
|
| 146 |
+
faces: np.ndarray,
|
| 147 |
+
mesh_color: Tuple[float, float, float],
|
| 148 |
+
):
|
| 149 |
+
renderer = self._ensure_renderer()
|
| 150 |
+
scene = renderer.scene
|
| 151 |
+
scene.clear_geometry()
|
| 152 |
+
scene.set_background(
|
| 153 |
+
[self.bg_color[0], self.bg_color[1], self.bg_color[2], 1.0]
|
| 154 |
+
)
|
| 155 |
+
|
| 156 |
+
mesh = self._make_o3d_mesh(vertices, faces)
|
| 157 |
+
scene.add_geometry("mesh", mesh, self._lit_material(mesh_color))
|
| 158 |
+
|
| 159 |
+
mesh_min = np.asarray(vertices, dtype=np.float64).min(axis=0)
|
| 160 |
+
mesh_max = np.asarray(vertices, dtype=np.float64).max(axis=0)
|
| 161 |
+
if self.add_ground:
|
| 162 |
+
ground = self._make_ground(mesh_min, mesh_max)
|
| 163 |
+
scene.add_geometry(
|
| 164 |
+
"ground", ground, self._lit_material(self.ground_color, roughness=0.95)
|
| 165 |
+
)
|
| 166 |
+
|
| 167 |
+
ai = _axis_index(self.up_axis)
|
| 168 |
+
sun_dir = np.array([0.35, 0.35, 0.35])
|
| 169 |
+
sun_dir[ai] = -1.0
|
| 170 |
+
sun_dir = sun_dir / np.linalg.norm(sun_dir)
|
| 171 |
+
|
| 172 |
+
scene.scene.set_sun_light(sun_dir.tolist(), [1.0, 1.0, 1.0], self.sun_intensity)
|
| 173 |
+
scene.scene.enable_sun_light(True)
|
| 174 |
+
scene.scene.set_indirect_light_intensity(self.ambient_intensity)
|
| 175 |
+
|
| 176 |
+
center = (mesh_min + mesh_max) / 2.0
|
| 177 |
+
return center
|
| 178 |
+
|
| 179 |
+
def _object_mask_from_depth(self):
|
| 180 |
+
renderer = self._renderer
|
| 181 |
+
depth = np.asarray(renderer.render_to_depth_image(z_in_view_space=True))
|
| 182 |
+
mask = np.isfinite(depth) & (depth > 0)
|
| 183 |
+
return mask
|
| 184 |
+
|
| 185 |
+
def _crop_resize_to_object(self, rgb: np.ndarray, mask: np.ndarray) -> np.ndarray:
|
| 186 |
+
from PIL import Image as _Image
|
| 187 |
+
|
| 188 |
+
ys, xs = np.where(mask)
|
| 189 |
+
if xs.size == 0:
|
| 190 |
+
out = _Image.fromarray(rgb).resize(
|
| 191 |
+
(self.img_res, self.img_res), _Image.LANCZOS
|
| 192 |
+
)
|
| 193 |
+
return np.asarray(out)
|
| 194 |
+
|
| 195 |
+
x0, y0, x1, y1 = xs.min(), ys.min(), xs.max(), ys.max()
|
| 196 |
+
cx, cy = (x0 + x1) / 2.0, (y0 + y1) / 2.0
|
| 197 |
+
size = int(max(x1 - x0, y1 - y0) * self.crop_padding)
|
| 198 |
+
size = max(size, 1)
|
| 199 |
+
half = size // 2
|
| 200 |
+
bx0, by0, bx1, by1 = (
|
| 201 |
+
int(round(cx - half)),
|
| 202 |
+
int(round(cy - half)),
|
| 203 |
+
int(round(cx - half)) + size,
|
| 204 |
+
int(round(cy - half)) + size,
|
| 205 |
+
)
|
| 206 |
+
|
| 207 |
+
H, W = rgb.shape[:2]
|
| 208 |
+
canvas = np.zeros((size, size, 3), dtype=np.uint8)
|
| 209 |
+
sx0, sy0 = max(0, bx0), max(0, by0)
|
| 210 |
+
sx1, sy1 = min(W, bx1), min(H, by1)
|
| 211 |
+
if sx1 > sx0 and sy1 > sy0:
|
| 212 |
+
canvas[sy0 - by0 : sy1 - by0, sx0 - bx0 : sx1 - bx0] = rgb[sy0:sy1, sx0:sx1]
|
| 213 |
+
|
| 214 |
+
out = _Image.fromarray(canvas).resize(
|
| 215 |
+
(self.img_res, self.img_res), _Image.LANCZOS
|
| 216 |
+
)
|
| 217 |
+
return np.ascontiguousarray(np.asarray(out))
|
| 218 |
+
|
| 219 |
+
def render(
|
| 220 |
+
self,
|
| 221 |
+
vertices: np.ndarray,
|
| 222 |
+
faces: np.ndarray,
|
| 223 |
+
num_views: int = 1,
|
| 224 |
+
mesh_color: Optional[ColorLike] = None,
|
| 225 |
+
azimuths: Optional[Sequence[float]] = None,
|
| 226 |
+
elevations: Optional[Sequence[float]] = None,
|
| 227 |
+
seed: Optional[int] = None,
|
| 228 |
+
) -> Tuple[List[np.ndarray], List[dict]]:
|
| 229 |
+
rng = np.random.default_rng(seed) if seed is not None else self._rng
|
| 230 |
+
rgb = self.mesh_color if mesh_color is None else _to_rgb01(mesh_color)
|
| 231 |
+
|
| 232 |
+
center = self._setup_scene(vertices, faces, rgb)
|
| 233 |
+
renderer = self._renderer
|
| 234 |
+
|
| 235 |
+
images: List[np.ndarray] = []
|
| 236 |
+
params: List[dict] = []
|
| 237 |
+
for v in range(num_views):
|
| 238 |
+
if azimuths is not None:
|
| 239 |
+
az = float(azimuths[v])
|
| 240 |
+
else:
|
| 241 |
+
az = float(rng.uniform(*self.azimuth_range))
|
| 242 |
+
if elevations is not None:
|
| 243 |
+
el = float(elevations[v])
|
| 244 |
+
else:
|
| 245 |
+
el = float(rng.uniform(*self.elevation_range))
|
| 246 |
+
|
| 247 |
+
eye, up = _orbit_eye(center, self.camera_distance, az, el, self.up_axis)
|
| 248 |
+
renderer.setup_camera(self.fov, center.tolist(), eye.tolist(), up.tolist())
|
| 249 |
+
|
| 250 |
+
img = renderer.render_to_image()
|
| 251 |
+
arr = np.asarray(img)
|
| 252 |
+
if arr.ndim == 3 and arr.shape[2] == 4:
|
| 253 |
+
arr = arr[:, :, :3]
|
| 254 |
+
arr = arr.astype(np.uint8)
|
| 255 |
+
|
| 256 |
+
if self.crop_to_object:
|
| 257 |
+
mask = self._object_mask_from_depth()
|
| 258 |
+
arr = self._crop_resize_to_object(arr, mask)
|
| 259 |
+
|
| 260 |
+
images.append(np.ascontiguousarray(arr))
|
| 261 |
+
params.append(
|
| 262 |
+
{"azimuth": az, "elevation": el, "distance": self.camera_distance}
|
| 263 |
+
)
|
| 264 |
+
|
| 265 |
+
return images, params
|
dataset/topo_dataset.py
ADDED
|
@@ -0,0 +1,88 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from __future__ import annotations
|
| 2 |
+
|
| 3 |
+
import os
|
| 4 |
+
import traceback
|
| 5 |
+
from typing import Dict, List, Optional
|
| 6 |
+
|
| 7 |
+
import numpy as np
|
| 8 |
+
import torch
|
| 9 |
+
|
| 10 |
+
from dataset.utils import (
|
| 11 |
+
MESH_EXTENSIONS,
|
| 12 |
+
dedup_quantized_mesh,
|
| 13 |
+
extract_active_voxels,
|
| 14 |
+
quantize_mesh_clustering,
|
| 15 |
+
)
|
| 16 |
+
|
| 17 |
+
|
| 18 |
+
class TopoVoxelDataset(torch.utils.data.Dataset):
|
| 19 |
+
def __init__(
|
| 20 |
+
self,
|
| 21 |
+
root_dir: str,
|
| 22 |
+
num_discrete: int = 1024,
|
| 23 |
+
voxel_res: int = 64,
|
| 24 |
+
max_vertices: Optional[int] = None,
|
| 25 |
+
num_samples: Optional[int] = None,
|
| 26 |
+
):
|
| 27 |
+
self.root_dir = root_dir
|
| 28 |
+
self.num_discrete = int(num_discrete)
|
| 29 |
+
self.voxel_res = int(voxel_res)
|
| 30 |
+
self.max_vertices = max_vertices
|
| 31 |
+
self.files = sorted(
|
| 32 |
+
f
|
| 33 |
+
for f in os.listdir(root_dir)
|
| 34 |
+
if os.path.splitext(f)[1].lower() in MESH_EXTENSIONS
|
| 35 |
+
)
|
| 36 |
+
if num_samples is not None:
|
| 37 |
+
self.files = self.files[:num_samples]
|
| 38 |
+
if not self.files:
|
| 39 |
+
raise ValueError(f"no mesh files ({MESH_EXTENSIONS}) under {root_dir}")
|
| 40 |
+
|
| 41 |
+
def __len__(self) -> int:
|
| 42 |
+
return len(self.files)
|
| 43 |
+
|
| 44 |
+
def __getitem__(self, idx: int) -> Dict:
|
| 45 |
+
name = os.path.splitext(self.files[idx])[0]
|
| 46 |
+
path = os.path.join(self.root_dir, self.files[idx])
|
| 47 |
+
res = self.num_discrete
|
| 48 |
+
try:
|
| 49 |
+
quantized = quantize_mesh_clustering(path, resolution=res)
|
| 50 |
+
if quantized is None:
|
| 51 |
+
return {"name": name, "error": "empty mesh"}
|
| 52 |
+
v_int, offsets, faces = quantized
|
| 53 |
+
if len(faces) < 1 or len(v_int) < 3:
|
| 54 |
+
return {"name": name, "error": "degenerate mesh after quantization"}
|
| 55 |
+
|
| 56 |
+
gt_int, _, gt_faces = dedup_quantized_mesh(v_int, offsets, faces, res)
|
| 57 |
+
num_gt = len(gt_int)
|
| 58 |
+
if num_gt < 3 or len(gt_faces) < 1:
|
| 59 |
+
return {"name": name, "error": f"too few vertices/faces ({num_gt})"}
|
| 60 |
+
if self.max_vertices is not None and num_gt > self.max_vertices:
|
| 61 |
+
return {
|
| 62 |
+
"name": name,
|
| 63 |
+
"error": f"vertex count {num_gt} exceeds max_vertices={self.max_vertices}",
|
| 64 |
+
}
|
| 65 |
+
|
| 66 |
+
quant_v = gt_int.astype(np.float64) / (res - 1.0) - 0.5
|
| 67 |
+
quant_v = np.clip(quant_v, -0.5 + 1e-6, 0.5 - 1e-6).astype(np.float32)
|
| 68 |
+
voxel_coords = extract_active_voxels(quant_v, gt_faces, self.voxel_res)
|
| 69 |
+
|
| 70 |
+
return {
|
| 71 |
+
"name": name,
|
| 72 |
+
"vertices": torch.from_numpy(gt_int.astype(np.int64)),
|
| 73 |
+
"voxel_coords": voxel_coords.long(),
|
| 74 |
+
}
|
| 75 |
+
except Exception as e:
|
| 76 |
+
return {"name": name, "error": f"{e}\n{traceback.format_exc()}"}
|
| 77 |
+
|
| 78 |
+
|
| 79 |
+
def collate_fn(batch: List[Dict]) -> Dict:
|
| 80 |
+
errors = [b for b in batch if "error" in b]
|
| 81 |
+
good = [b for b in batch if "error" not in b]
|
| 82 |
+
collated: Dict = {"errors": errors}
|
| 83 |
+
if not good:
|
| 84 |
+
return collated
|
| 85 |
+
collated["name"] = [b["name"] for b in good]
|
| 86 |
+
collated["vertices"] = [b["vertices"] for b in good]
|
| 87 |
+
collated["voxel_coords"] = [b["voxel_coords"] for b in good]
|
| 88 |
+
return collated
|
dataset/utils.py
ADDED
|
@@ -0,0 +1,283 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import numpy as np
|
| 2 |
+
import torch
|
| 3 |
+
import trimesh
|
| 4 |
+
from trimesh import grouping
|
| 5 |
+
|
| 6 |
+
from o_voxel.convert import mesh_to_flexible_dual_grid
|
| 7 |
+
|
| 8 |
+
MESH_EXTENSIONS = {".obj", ".glb", ".gltf", ".ply", ".stl", ".off"}
|
| 9 |
+
|
| 10 |
+
|
| 11 |
+
def quantize_mesh_clustering(mesh_path: str, resolution: int = 1024):
|
| 12 |
+
mesh = trimesh.load(mesh_path, process=False, force="mesh")
|
| 13 |
+
if mesh is None or len(mesh.vertices) == 0 or len(mesh.faces) == 0:
|
| 14 |
+
return None
|
| 15 |
+
|
| 16 |
+
vertices = np.asarray(mesh.vertices, dtype=np.float64)
|
| 17 |
+
faces = np.asarray(mesh.faces, dtype=np.int64)
|
| 18 |
+
|
| 19 |
+
bbox_min, bbox_max = vertices.min(axis=0), vertices.max(axis=0)
|
| 20 |
+
center = (bbox_min + bbox_max) / 2.0
|
| 21 |
+
max_extent = max(float((bbox_max - bbox_min).max()), 1e-7)
|
| 22 |
+
normalized_v = (vertices - center) / max_extent + 0.5 # [0, 1]
|
| 23 |
+
|
| 24 |
+
v_grid = np.clip(np.floor(normalized_v * resolution), 0, resolution - 1).astype(
|
| 25 |
+
np.int64
|
| 26 |
+
)
|
| 27 |
+
v_hash = (
|
| 28 |
+
v_grid[:, 0] * resolution * resolution
|
| 29 |
+
+ v_grid[:, 1] * resolution
|
| 30 |
+
+ v_grid[:, 2]
|
| 31 |
+
)
|
| 32 |
+
|
| 33 |
+
unique_hashes, inverse = np.unique(v_hash, return_inverse=True)
|
| 34 |
+
num_clusters = len(unique_hashes)
|
| 35 |
+
|
| 36 |
+
v_sum = np.zeros((num_clusters, 3), dtype=np.float64)
|
| 37 |
+
counts = np.zeros(num_clusters, dtype=np.float64)
|
| 38 |
+
np.add.at(v_sum, inverse, normalized_v)
|
| 39 |
+
np.add.at(counts, inverse, 1)
|
| 40 |
+
v_mean = v_sum / counts[:, None]
|
| 41 |
+
|
| 42 |
+
v_int = np.clip(np.floor(v_mean * resolution), 0, resolution - 1).astype(np.int32)
|
| 43 |
+
|
| 44 |
+
# Offset of the cluster mean relative to the voxel center, scaled to (-1, 1).
|
| 45 |
+
voxel_center = (v_int.astype(np.float64) + 0.5) / float(resolution)
|
| 46 |
+
offsets = ((v_mean - voxel_center) * 2.0 * resolution).astype(np.float32)
|
| 47 |
+
offsets = np.clip(offsets, -1.0, 1.0)
|
| 48 |
+
|
| 49 |
+
new_faces = inverse[faces]
|
| 50 |
+
valid = (
|
| 51 |
+
(new_faces[:, 0] != new_faces[:, 1])
|
| 52 |
+
& (new_faces[:, 1] != new_faces[:, 2])
|
| 53 |
+
& (new_faces[:, 2] != new_faces[:, 0])
|
| 54 |
+
)
|
| 55 |
+
return v_int, offsets, new_faces[valid]
|
| 56 |
+
|
| 57 |
+
|
| 58 |
+
def realign_offsets(
|
| 59 |
+
orig_int: np.ndarray, orig_off: np.ndarray, kept_int: np.ndarray, resolution: int
|
| 60 |
+
) -> np.ndarray:
|
| 61 |
+
base = resolution * resolution
|
| 62 |
+
orig_i = orig_int.astype(np.int64)
|
| 63 |
+
kept_i = kept_int.astype(np.int64)
|
| 64 |
+
orig_h = orig_i[:, 0] * base + orig_i[:, 1] * resolution + orig_i[:, 2]
|
| 65 |
+
kept_h = kept_i[:, 0] * base + kept_i[:, 1] * resolution + kept_i[:, 2]
|
| 66 |
+
|
| 67 |
+
order = np.argsort(orig_h)
|
| 68 |
+
sorted_h = orig_h[order]
|
| 69 |
+
sorted_off = orig_off.astype(np.float64)[order]
|
| 70 |
+
cumsum = np.concatenate([np.zeros((1, 3)), np.cumsum(sorted_off, axis=0)], axis=0)
|
| 71 |
+
left = np.searchsorted(sorted_h, kept_h, side="left")
|
| 72 |
+
right = np.searchsorted(sorted_h, kept_h, side="right")
|
| 73 |
+
counts = (right - left).clip(min=1).astype(np.float64)
|
| 74 |
+
out = ((cumsum[right] - cumsum[left]) / counts[:, None]).astype(np.float32)
|
| 75 |
+
out = np.clip(out, -1.0, 1.0)
|
| 76 |
+
out[right <= left] = 0.0
|
| 77 |
+
return out
|
| 78 |
+
|
| 79 |
+
|
| 80 |
+
def dedup_quantized_mesh(
|
| 81 |
+
v_int: np.ndarray, offsets: np.ndarray, faces: np.ndarray, resolution: int
|
| 82 |
+
):
|
| 83 |
+
tmesh = trimesh.Trimesh(vertices=v_int, faces=faces, process=False)
|
| 84 |
+
tmesh.merge_vertices()
|
| 85 |
+
tmesh.update_faces(tmesh.nondegenerate_faces())
|
| 86 |
+
tmesh.update_faces(tmesh.unique_faces())
|
| 87 |
+
tmesh.remove_unreferenced_vertices()
|
| 88 |
+
|
| 89 |
+
gt_int = np.asarray(tmesh.vertices).astype(np.int32)
|
| 90 |
+
gt_faces = np.asarray(tmesh.faces, dtype=np.int64)
|
| 91 |
+
gt_offsets = realign_offsets(v_int, offsets, gt_int, resolution)
|
| 92 |
+
return gt_int, gt_offsets, gt_faces
|
| 93 |
+
|
| 94 |
+
|
| 95 |
+
def extract_active_voxels(
|
| 96 |
+
vertices: np.ndarray, faces: np.ndarray, resolution: int
|
| 97 |
+
) -> torch.Tensor:
|
| 98 |
+
coords, *_ = mesh_to_flexible_dual_grid(
|
| 99 |
+
vertices=torch.as_tensor(vertices * 0.99999, dtype=torch.float32).contiguous(),
|
| 100 |
+
faces=torch.as_tensor(faces, dtype=torch.int32).contiguous(),
|
| 101 |
+
grid_size=resolution,
|
| 102 |
+
aabb=torch.tensor([[-0.5, -0.5, -0.5], [0.5, 0.5, 0.5]], dtype=torch.float32),
|
| 103 |
+
)
|
| 104 |
+
coords = coords.cpu().long()
|
| 105 |
+
coords_1d = (
|
| 106 |
+
coords[:, 0] * resolution * resolution
|
| 107 |
+
+ coords[:, 1] * resolution
|
| 108 |
+
+ coords[:, 2]
|
| 109 |
+
)
|
| 110 |
+
return coords[torch.argsort(coords_1d)].int()
|
| 111 |
+
|
| 112 |
+
|
| 113 |
+
def union_voxels(a: torch.Tensor, b: torch.Tensor, resolution: int) -> torch.Tensor:
|
| 114 |
+
if a.numel() == 0:
|
| 115 |
+
return b.int().clone()
|
| 116 |
+
if b.numel() == 0:
|
| 117 |
+
return a.int().clone()
|
| 118 |
+
combined = torch.cat([a.reshape(-1, 3).long(), b.reshape(-1, 3).long()], dim=0)
|
| 119 |
+
combined = combined.clamp(0, resolution - 1)
|
| 120 |
+
hashes = (
|
| 121 |
+
combined[:, 0] * resolution * resolution
|
| 122 |
+
+ combined[:, 1] * resolution
|
| 123 |
+
+ combined[:, 2]
|
| 124 |
+
)
|
| 125 |
+
sorted_hashes, sort_idx = torch.sort(hashes)
|
| 126 |
+
keep = torch.ones_like(sorted_hashes, dtype=torch.bool)
|
| 127 |
+
keep[1:] = sorted_hashes[1:] != sorted_hashes[:-1]
|
| 128 |
+
return combined[sort_idx[keep]].int()
|
| 129 |
+
|
| 130 |
+
|
| 131 |
+
def _sample_surface_uniform(tm_mesh: trimesh.Trimesh, n_samples: int):
|
| 132 |
+
face_idx = np.random.choice(len(tm_mesh.faces), size=n_samples, replace=True)
|
| 133 |
+
tri = tm_mesh.vertices[tm_mesh.faces[face_idx]] # (N, 3, 3)
|
| 134 |
+
u = np.random.rand(n_samples, 1)
|
| 135 |
+
v = np.random.rand(n_samples, 1)
|
| 136 |
+
sqrt_u = np.sqrt(u)
|
| 137 |
+
points = (
|
| 138 |
+
(1 - sqrt_u) * tri[:, 0]
|
| 139 |
+
+ (sqrt_u * (1 - v)) * tri[:, 1]
|
| 140 |
+
+ (sqrt_u * v) * tri[:, 2]
|
| 141 |
+
)
|
| 142 |
+
normals = tm_mesh.face_normals[face_idx]
|
| 143 |
+
return points.astype(np.float32), normals.astype(np.float32), face_idx
|
| 144 |
+
|
| 145 |
+
|
| 146 |
+
def _sample_edges_dora(
|
| 147 |
+
tm_mesh: trimesh.Trimesh, n_len_samples: int, n_uniform_samples: int
|
| 148 |
+
):
|
| 149 |
+
parts_start, parts_end, parts_norm, parts_virt = [], [], [], []
|
| 150 |
+
|
| 151 |
+
adj_faces = tm_mesh.face_adjacency
|
| 152 |
+
adj_edges = tm_mesh.face_adjacency_edges
|
| 153 |
+
if len(adj_faces) > 0:
|
| 154 |
+
n0 = tm_mesh.face_normals[adj_faces[:, 0]]
|
| 155 |
+
n1 = tm_mesh.face_normals[adj_faces[:, 1]]
|
| 156 |
+
sum_normals = n0 + n1
|
| 157 |
+
norms = np.linalg.norm(sum_normals, axis=1, keepdims=True)
|
| 158 |
+
norms[norms < 1e-6] = 1.0
|
| 159 |
+
|
| 160 |
+
faces_pair = tm_mesh.faces[adj_faces]
|
| 161 |
+
unique_idx_0 = np.sum(faces_pair, axis=2)[:, 0] - np.sum(adj_edges, axis=1)
|
| 162 |
+
unique_idx_1 = np.sum(faces_pair, axis=2)[:, 1] - np.sum(adj_edges, axis=1)
|
| 163 |
+
virtual = (
|
| 164 |
+
tm_mesh.vertices[unique_idx_0] + tm_mesh.vertices[unique_idx_1]
|
| 165 |
+
) * 0.5
|
| 166 |
+
|
| 167 |
+
parts_start.append(tm_mesh.vertices[adj_edges[:, 0]])
|
| 168 |
+
parts_end.append(tm_mesh.vertices[adj_edges[:, 1]])
|
| 169 |
+
parts_norm.append(sum_normals / norms)
|
| 170 |
+
parts_virt.append(virtual)
|
| 171 |
+
|
| 172 |
+
edges_sorted = tm_mesh.edges_sorted
|
| 173 |
+
if len(edges_sorted) > 0:
|
| 174 |
+
boundary_group = grouping.group_rows(edges_sorted, require_count=1)
|
| 175 |
+
if len(boundary_group) > 0:
|
| 176 |
+
boundary_indices = np.concatenate(
|
| 177 |
+
[np.atleast_1d(g) for g in boundary_group]
|
| 178 |
+
)
|
| 179 |
+
face_indices = boundary_indices // 3
|
| 180 |
+
edge_v = edges_sorted[boundary_indices]
|
| 181 |
+
unique_idx = np.sum(tm_mesh.faces[face_indices], axis=1) - np.sum(
|
| 182 |
+
edge_v, axis=1
|
| 183 |
+
)
|
| 184 |
+
|
| 185 |
+
parts_start.append(tm_mesh.vertices[edge_v[:, 0]])
|
| 186 |
+
parts_end.append(tm_mesh.vertices[edge_v[:, 1]])
|
| 187 |
+
parts_norm.append(tm_mesh.face_normals[face_indices])
|
| 188 |
+
parts_virt.append(tm_mesh.vertices[unique_idx])
|
| 189 |
+
|
| 190 |
+
if not parts_start:
|
| 191 |
+
return None, None, None
|
| 192 |
+
|
| 193 |
+
v_start = np.concatenate(parts_start, axis=0)
|
| 194 |
+
v_end = np.concatenate(parts_end, axis=0)
|
| 195 |
+
normals = np.concatenate(parts_norm, axis=0)
|
| 196 |
+
v_virtual = np.concatenate(parts_virt, axis=0)
|
| 197 |
+
|
| 198 |
+
lengths = np.linalg.norm(v_end - v_start, axis=1)
|
| 199 |
+
total = lengths.sum()
|
| 200 |
+
num_edges = len(lengths)
|
| 201 |
+
probs_len = (
|
| 202 |
+
lengths / total if total >= 1e-9 else np.full(num_edges, 1.0 / num_edges)
|
| 203 |
+
)
|
| 204 |
+
probs_len = probs_len / probs_len.sum()
|
| 205 |
+
|
| 206 |
+
chosen = np.concatenate(
|
| 207 |
+
[
|
| 208 |
+
np.random.choice(num_edges, size=n_len_samples, p=probs_len),
|
| 209 |
+
np.random.choice(num_edges, size=n_uniform_samples),
|
| 210 |
+
]
|
| 211 |
+
)
|
| 212 |
+
t = np.random.rand(len(chosen), 1)
|
| 213 |
+
points = v_start[chosen] + (v_end[chosen] - v_start[chosen]) * t
|
| 214 |
+
triplets = np.stack([v_start[chosen], v_end[chosen], v_virtual[chosen]], axis=1)
|
| 215 |
+
return (
|
| 216 |
+
points.astype(np.float32),
|
| 217 |
+
normals[chosen].astype(np.float32),
|
| 218 |
+
triplets.astype(np.float32),
|
| 219 |
+
)
|
| 220 |
+
|
| 221 |
+
|
| 222 |
+
def _vdf_from_triplets(
|
| 223 |
+
points: np.ndarray, triplets: np.ndarray, normalize: bool
|
| 224 |
+
) -> np.ndarray:
|
| 225 |
+
view_dtype = np.dtype((np.void, triplets.dtype.itemsize * triplets.shape[-1]))
|
| 226 |
+
v_view = triplets.view(view_dtype).squeeze(-1)
|
| 227 |
+
sort_idx = np.argsort(v_view, axis=1)
|
| 228 |
+
v_sorted = triplets[np.arange(triplets.shape[0])[:, None], sort_idx]
|
| 229 |
+
|
| 230 |
+
dirs = v_sorted - points[:, None, :] # (N, 3, 3)
|
| 231 |
+
if normalize:
|
| 232 |
+
dirs = dirs / (np.linalg.norm(dirs, axis=-1, keepdims=True) + 1e-8)
|
| 233 |
+
return dirs.reshape(len(points), 9).astype(np.float32)
|
| 234 |
+
|
| 235 |
+
|
| 236 |
+
def sample_point_features(
|
| 237 |
+
tm_mesh: trimesh.Trimesh,
|
| 238 |
+
n_samples: int,
|
| 239 |
+
sample_type: str = "dora",
|
| 240 |
+
normalize_vdf: bool = True,
|
| 241 |
+
) -> torch.Tensor:
|
| 242 |
+
# (N, 15) float32 point features: [xyz(3), normal(3), vdf(9)].
|
| 243 |
+
vertices = np.asarray(tm_mesh.vertices, dtype=np.float64)
|
| 244 |
+
faces = np.asarray(tm_mesh.faces)
|
| 245 |
+
|
| 246 |
+
if sample_type == "dora":
|
| 247 |
+
n_surf_area = n_samples // 4
|
| 248 |
+
n_surf_uniform = n_samples // 4
|
| 249 |
+
n_edge_len = n_samples // 4
|
| 250 |
+
n_edge_uniform = n_samples - n_surf_area - n_surf_uniform - n_edge_len
|
| 251 |
+
|
| 252 |
+
p_edge, n_edge, triplets_edge = _sample_edges_dora(
|
| 253 |
+
tm_mesh, n_edge_len, n_edge_uniform
|
| 254 |
+
)
|
| 255 |
+
if p_edge is None:
|
| 256 |
+
n_surf_area += n_edge_len
|
| 257 |
+
n_surf_uniform += n_edge_uniform
|
| 258 |
+
elif sample_type == "uniform":
|
| 259 |
+
n_surf_area, n_surf_uniform = n_samples, 0
|
| 260 |
+
p_edge = None
|
| 261 |
+
else:
|
| 262 |
+
raise ValueError(f"unknown sample_type: {sample_type!r}")
|
| 263 |
+
|
| 264 |
+
p_area, idx_area = tm_mesh.sample(n_surf_area, return_index=True)
|
| 265 |
+
n_area = tm_mesh.face_normals[idx_area]
|
| 266 |
+
if n_surf_uniform > 0:
|
| 267 |
+
p_unif, n_unif, idx_unif = _sample_surface_uniform(tm_mesh, n_surf_uniform)
|
| 268 |
+
points = np.concatenate([p_area, p_unif], axis=0).astype(np.float32)
|
| 269 |
+
normals = np.concatenate([n_area, n_unif], axis=0).astype(np.float32)
|
| 270 |
+
idx_surf = np.concatenate([idx_area, idx_unif], axis=0)
|
| 271 |
+
else:
|
| 272 |
+
points = p_area.astype(np.float32)
|
| 273 |
+
normals = n_area.astype(np.float32)
|
| 274 |
+
idx_surf = idx_area
|
| 275 |
+
triplets = vertices[faces[idx_surf]].astype(np.float32)
|
| 276 |
+
|
| 277 |
+
if p_edge is not None:
|
| 278 |
+
points = np.concatenate([points, p_edge], axis=0)
|
| 279 |
+
normals = np.concatenate([normals, n_edge], axis=0)
|
| 280 |
+
triplets = np.concatenate([triplets, triplets_edge], axis=0)
|
| 281 |
+
|
| 282 |
+
vdf = _vdf_from_triplets(points, triplets, normalize=normalize_vdf)
|
| 283 |
+
return torch.from_numpy(np.concatenate([points, normals, vdf], axis=-1))
|
dataset/voxel_dataset.py
ADDED
|
@@ -0,0 +1,211 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import os
|
| 2 |
+
from typing import Dict, List, Optional
|
| 3 |
+
import numpy as np
|
| 4 |
+
import torch
|
| 5 |
+
import trimesh
|
| 6 |
+
from torch.utils.data import Dataset
|
| 7 |
+
|
| 8 |
+
from dataset.utils import (
|
| 9 |
+
MESH_EXTENSIONS,
|
| 10 |
+
dedup_quantized_mesh,
|
| 11 |
+
extract_active_voxels,
|
| 12 |
+
quantize_mesh_clustering,
|
| 13 |
+
sample_point_features,
|
| 14 |
+
union_voxels,
|
| 15 |
+
)
|
| 16 |
+
|
| 17 |
+
|
| 18 |
+
class VoxelVertexDataset(Dataset):
|
| 19 |
+
def __init__(
|
| 20 |
+
self,
|
| 21 |
+
root_dir: str,
|
| 22 |
+
resolution: int = 1024,
|
| 23 |
+
min_resolution: int = 64,
|
| 24 |
+
pc_sample_number: int = 819200,
|
| 25 |
+
sample_type: str = "dora",
|
| 26 |
+
normalize_vdf: bool = True,
|
| 27 |
+
need_encoder_inputs: bool = True,
|
| 28 |
+
min_vertices: int = 0,
|
| 29 |
+
max_vertices: Optional[int] = None,
|
| 30 |
+
num_samples: Optional[int] = None,
|
| 31 |
+
render: bool = False,
|
| 32 |
+
img_res: int = 518,
|
| 33 |
+
render_azimuth: float = 45.0,
|
| 34 |
+
render_elevation: float = 30.0,
|
| 35 |
+
):
|
| 36 |
+
self.root_dir = root_dir
|
| 37 |
+
self.resolution = resolution
|
| 38 |
+
self.min_resolution = min_resolution
|
| 39 |
+
self.pc_sample_number = pc_sample_number
|
| 40 |
+
self.sample_type = sample_type
|
| 41 |
+
self.normalize_vdf = normalize_vdf
|
| 42 |
+
self.need_encoder_inputs = need_encoder_inputs
|
| 43 |
+
self.min_vertices = min_vertices
|
| 44 |
+
self.max_vertices = max_vertices
|
| 45 |
+
self.render = render
|
| 46 |
+
self.img_res = img_res
|
| 47 |
+
self.render_azimuth = render_azimuth
|
| 48 |
+
self.render_elevation = render_elevation
|
| 49 |
+
|
| 50 |
+
self.files = sorted(
|
| 51 |
+
f
|
| 52 |
+
for f in os.listdir(root_dir)
|
| 53 |
+
if os.path.splitext(f)[1].lower() in MESH_EXTENSIONS
|
| 54 |
+
)
|
| 55 |
+
if num_samples is not None:
|
| 56 |
+
self.files = self.files[:num_samples]
|
| 57 |
+
if not self.files:
|
| 58 |
+
raise ValueError(f"no mesh files ({MESH_EXTENSIONS}) under {root_dir}")
|
| 59 |
+
|
| 60 |
+
self._renderer = None # lazy: one EGL context per DataLoader worker
|
| 61 |
+
|
| 62 |
+
def __len__(self) -> int:
|
| 63 |
+
return len(self.files)
|
| 64 |
+
|
| 65 |
+
def _render_image(self, vertices: np.ndarray, faces: np.ndarray) -> np.ndarray:
|
| 66 |
+
if self._renderer is None:
|
| 67 |
+
from dataset.mesh_render import WhiteModelRenderer
|
| 68 |
+
|
| 69 |
+
self._renderer = WhiteModelRenderer(
|
| 70 |
+
img_res=self.img_res,
|
| 71 |
+
mesh_color=(0.78, 0.78, 0.82),
|
| 72 |
+
bg_color=(0.0, 0.0, 0.0),
|
| 73 |
+
up_axis="y",
|
| 74 |
+
add_ground=False,
|
| 75 |
+
shadow=True,
|
| 76 |
+
crop_to_object=True,
|
| 77 |
+
crop_padding=1.2,
|
| 78 |
+
)
|
| 79 |
+
imgs, _ = self._renderer.render(
|
| 80 |
+
np.asarray(vertices, dtype=np.float64),
|
| 81 |
+
np.asarray(faces, dtype=np.int64),
|
| 82 |
+
num_views=1,
|
| 83 |
+
azimuths=[self.render_azimuth],
|
| 84 |
+
elevations=[self.render_elevation],
|
| 85 |
+
)
|
| 86 |
+
return imgs[0] # (img_res, img_res, 3) uint8
|
| 87 |
+
|
| 88 |
+
def __getitem__(self, idx: int) -> Dict:
|
| 89 |
+
name = os.path.splitext(self.files[idx])[0]
|
| 90 |
+
path = os.path.join(self.root_dir, self.files[idx])
|
| 91 |
+
res = self.resolution
|
| 92 |
+
min_res = self.min_resolution
|
| 93 |
+
try:
|
| 94 |
+
quantized = quantize_mesh_clustering(path, resolution=res)
|
| 95 |
+
if quantized is None:
|
| 96 |
+
return {"name": name, "error": "empty mesh"}
|
| 97 |
+
v_int, offsets, faces = quantized
|
| 98 |
+
if len(faces) < 1 or len(v_int) < 3:
|
| 99 |
+
return {"name": name, "error": "degenerate mesh after quantization"}
|
| 100 |
+
|
| 101 |
+
gt_int, gt_offsets, gt_faces = dedup_quantized_mesh(
|
| 102 |
+
v_int, offsets, faces, res
|
| 103 |
+
)
|
| 104 |
+
num_gt = len(gt_int)
|
| 105 |
+
if num_gt < max(self.min_vertices, 3) or len(gt_faces) < 1:
|
| 106 |
+
return {
|
| 107 |
+
"name": name,
|
| 108 |
+
"error": f"too few vertices/faces after dedup ({num_gt})",
|
| 109 |
+
}
|
| 110 |
+
if self.max_vertices is not None and num_gt > self.max_vertices:
|
| 111 |
+
return {
|
| 112 |
+
"name": name,
|
| 113 |
+
"error": f"vertex count {num_gt} exceeds max_vertices={self.max_vertices}",
|
| 114 |
+
}
|
| 115 |
+
|
| 116 |
+
quant_v = gt_int.astype(np.float64) / (res - 1.0) - 0.5
|
| 117 |
+
quant_v = np.clip(quant_v, -0.5 + 1e-6, 0.5 - 1e-6).astype(np.float32)
|
| 118 |
+
tmesh = trimesh.Trimesh(vertices=quant_v, faces=gt_faces, process=False)
|
| 119 |
+
|
| 120 |
+
if self.need_encoder_inputs:
|
| 121 |
+
vertex_added_active = extract_active_voxels(quant_v, gt_faces, res)
|
| 122 |
+
vertex_added_active = union_voxels(
|
| 123 |
+
vertex_added_active, torch.from_numpy(gt_int), res
|
| 124 |
+
)
|
| 125 |
+
point_cloud = sample_point_features(
|
| 126 |
+
tmesh,
|
| 127 |
+
self.pc_sample_number,
|
| 128 |
+
sample_type=self.sample_type,
|
| 129 |
+
normalize_vdf=self.normalize_vdf,
|
| 130 |
+
)
|
| 131 |
+
else:
|
| 132 |
+
vertex_added_active = torch.zeros((0, 3), dtype=torch.int32)
|
| 133 |
+
point_cloud = torch.zeros((0, 15), dtype=torch.float32)
|
| 134 |
+
|
| 135 |
+
min_active = extract_active_voxels(quant_v, gt_faces, self.min_resolution)
|
| 136 |
+
|
| 137 |
+
data = {
|
| 138 |
+
"name": name,
|
| 139 |
+
f"gt_vertex_voxels_{res}": torch.from_numpy(gt_int),
|
| 140 |
+
f"gt_vertex_offsets_{res}": torch.from_numpy(gt_offsets),
|
| 141 |
+
"quantized_vertices": torch.from_numpy(quant_v),
|
| 142 |
+
"quantized_faces": torch.from_numpy(gt_faces),
|
| 143 |
+
f"vertex_added_active_voxels_{res}": vertex_added_active,
|
| 144 |
+
f"point_cloud_{res}": point_cloud,
|
| 145 |
+
f"active_voxels_{min_res}": min_active,
|
| 146 |
+
}
|
| 147 |
+
|
| 148 |
+
# Raw mesh, bbox-normalized into the same [-0.5, 0.5] frame.
|
| 149 |
+
raw = trimesh.load(path, process=False, force="mesh")
|
| 150 |
+
raw_v = np.asarray(raw.vertices, dtype=np.float64)
|
| 151 |
+
center = (raw_v.min(axis=0) + raw_v.max(axis=0)) / 2.0
|
| 152 |
+
extent = max(float((raw_v.max(axis=0) - raw_v.min(axis=0)).max()), 1e-7)
|
| 153 |
+
data["original_vertices"] = torch.from_numpy(
|
| 154 |
+
((raw_v - center) / extent).astype(np.float32)
|
| 155 |
+
)
|
| 156 |
+
data["original_faces"] = torch.from_numpy(
|
| 157 |
+
np.asarray(raw.faces, dtype=np.int64)
|
| 158 |
+
)
|
| 159 |
+
|
| 160 |
+
if self.render:
|
| 161 |
+
render_v = v_int.astype(np.float64) / res - 0.5
|
| 162 |
+
data["image"] = self._render_image(render_v, faces)
|
| 163 |
+
|
| 164 |
+
return data
|
| 165 |
+
except Exception as e:
|
| 166 |
+
import traceback
|
| 167 |
+
|
| 168 |
+
return {"name": name, "error": f"{e}\n{traceback.format_exc()}"}
|
| 169 |
+
|
| 170 |
+
|
| 171 |
+
def collate_fn(
|
| 172 |
+
batch: List[Dict], resolution: int = 1024, min_resolution: int = 64
|
| 173 |
+
) -> Dict:
|
| 174 |
+
res = resolution
|
| 175 |
+
min_res = min_resolution
|
| 176 |
+
errors = [b for b in batch if "error" in b]
|
| 177 |
+
batch = [b for b in batch if "error" not in b]
|
| 178 |
+
collated: Dict = {"errors": errors}
|
| 179 |
+
if not batch:
|
| 180 |
+
return collated
|
| 181 |
+
|
| 182 |
+
collated["name"] = [b["name"] for b in batch]
|
| 183 |
+
for key in (
|
| 184 |
+
"quantized_vertices",
|
| 185 |
+
"quantized_faces",
|
| 186 |
+
"original_vertices",
|
| 187 |
+
"original_faces",
|
| 188 |
+
):
|
| 189 |
+
collated[key] = [b[key] for b in batch]
|
| 190 |
+
if "image" in batch[0]:
|
| 191 |
+
collated["image"] = [b["image"] for b in batch]
|
| 192 |
+
|
| 193 |
+
for key in (
|
| 194 |
+
f"gt_vertex_voxels_{res}",
|
| 195 |
+
f"vertex_added_active_voxels_{res}",
|
| 196 |
+
f"active_voxels_{min_res}",
|
| 197 |
+
):
|
| 198 |
+
rows = []
|
| 199 |
+
for i, b in enumerate(batch):
|
| 200 |
+
coords = b[key]
|
| 201 |
+
batch_idx = torch.full((coords.shape[0], 1), i, dtype=torch.int32)
|
| 202 |
+
rows.append(torch.cat([batch_idx, coords], dim=1))
|
| 203 |
+
collated[key] = torch.cat(rows, dim=0)
|
| 204 |
+
|
| 205 |
+
collated[f"gt_vertex_offsets_{res}"] = torch.cat(
|
| 206 |
+
[b[f"gt_vertex_offsets_{res}"] for b in batch], dim=0
|
| 207 |
+
)
|
| 208 |
+
collated[f"point_cloud_{res}"] = torch.stack(
|
| 209 |
+
[b[f"point_cloud_{res}"] for b in batch], dim=0
|
| 210 |
+
)
|
| 211 |
+
return collated
|
models/__init__.py
ADDED
|
@@ -0,0 +1,9 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from models.offset_head import OffsetHead
|
| 2 |
+
from models.vdf_encoder import VDFEncoder
|
| 3 |
+
from models.vertex_autoencoder import VertexVAE
|
| 4 |
+
from models.dino_encoder import DinoV2Encoder
|
| 5 |
+
from models.vertex_structured_flow import VertexSLatFlowModel
|
| 6 |
+
from models.flow_sampler import VertFlowEulerCfgSampler, TopoFlowEulerSampler
|
| 7 |
+
from models.topo_autoencoder import TopologyVAE
|
| 8 |
+
from models.topo_flow import TopologySiTFlow
|
| 9 |
+
from models.voxel_encoder import VoxelFieldConditioner
|
models/dino_encoder.py
ADDED
|
@@ -0,0 +1,73 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import os
|
| 2 |
+
from typing import Union
|
| 3 |
+
|
| 4 |
+
import numpy as np
|
| 5 |
+
import torch
|
| 6 |
+
import torch.nn as nn
|
| 7 |
+
import torch.nn.functional as F
|
| 8 |
+
|
| 9 |
+
DINO_GITHUB_REPO = "facebookresearch/dinov2"
|
| 10 |
+
DINO_LOCAL_REPO_DIRNAME = "facebookresearch_dinov2_main"
|
| 11 |
+
|
| 12 |
+
|
| 13 |
+
class DinoV2Encoder(nn.Module):
|
| 14 |
+
def __init__(
|
| 15 |
+
self,
|
| 16 |
+
model_name: str,
|
| 17 |
+
hub_dir: str,
|
| 18 |
+
img_res: int,
|
| 19 |
+
):
|
| 20 |
+
super().__init__()
|
| 21 |
+
self.img_res = int(img_res)
|
| 22 |
+
|
| 23 |
+
hub_dir = os.path.abspath(os.path.expanduser(hub_dir))
|
| 24 |
+
os.makedirs(hub_dir, exist_ok=True)
|
| 25 |
+
torch.hub.set_dir(hub_dir)
|
| 26 |
+
local_repo = os.path.join(hub_dir, DINO_LOCAL_REPO_DIRNAME)
|
| 27 |
+
if os.path.isdir(local_repo):
|
| 28 |
+
self.backbone = torch.hub.load(
|
| 29 |
+
local_repo, model_name, source="local", pretrained=True
|
| 30 |
+
)
|
| 31 |
+
else:
|
| 32 |
+
self.backbone = torch.hub.load(
|
| 33 |
+
DINO_GITHUB_REPO, model_name, source="github", pretrained=True
|
| 34 |
+
)
|
| 35 |
+
self.backbone.eval()
|
| 36 |
+
for p in self.backbone.parameters():
|
| 37 |
+
p.requires_grad_(False)
|
| 38 |
+
|
| 39 |
+
self.register_buffer(
|
| 40 |
+
"img_mean",
|
| 41 |
+
torch.tensor([0.485, 0.456, 0.406]).view(1, 3, 1, 1),
|
| 42 |
+
persistent=False,
|
| 43 |
+
)
|
| 44 |
+
self.register_buffer(
|
| 45 |
+
"img_std",
|
| 46 |
+
torch.tensor([0.229, 0.224, 0.225]).view(1, 3, 1, 1),
|
| 47 |
+
persistent=False,
|
| 48 |
+
)
|
| 49 |
+
|
| 50 |
+
@torch.no_grad()
|
| 51 |
+
def forward(self, images: Union[np.ndarray, torch.Tensor]) -> torch.Tensor:
|
| 52 |
+
"""Images -> layer-normed DINO-v2 patch tokens (B, L, C).
|
| 53 |
+
|
| 54 |
+
Accepts uint8 channels-last images — (H, W, 3) or (B, H, W, 3) in
|
| 55 |
+
[0, 255], the dataset render format — or float channels-first
|
| 56 |
+
(B, 3, H, W) in [0, 1]. Any resolution; resized to ``img_res``.
|
| 57 |
+
"""
|
| 58 |
+
if isinstance(images, np.ndarray):
|
| 59 |
+
images = torch.from_numpy(np.ascontiguousarray(images))
|
| 60 |
+
if images.dim() == 3:
|
| 61 |
+
images = images[None]
|
| 62 |
+
if images.dtype == torch.uint8:
|
| 63 |
+
images = images.permute(0, 3, 1, 2).float() / 255.0
|
| 64 |
+
|
| 65 |
+
x = images.to(device=self.img_mean.device, dtype=torch.float32)
|
| 66 |
+
if x.shape[-2:] != (self.img_res, self.img_res):
|
| 67 |
+
x = F.interpolate(
|
| 68 |
+
x, (self.img_res, self.img_res), mode="bicubic", align_corners=False
|
| 69 |
+
)
|
| 70 |
+
x = (x - self.img_mean) / self.img_std
|
| 71 |
+
|
| 72 |
+
feats = self.backbone(x, is_training=True)["x_prenorm"]
|
| 73 |
+
return F.layer_norm(feats, feats.shape[-1:])
|
models/flow_sampler.py
ADDED
|
@@ -0,0 +1,58 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import numpy as np
|
| 2 |
+
import torch
|
| 3 |
+
|
| 4 |
+
|
| 5 |
+
class VertFlowEulerCfgSampler:
|
| 6 |
+
def _pred_v(self, model, x_t, t, cond, **kwargs):
|
| 7 |
+
t_vec = torch.full(
|
| 8 |
+
(x_t.shape[0],), 1000.0 * t, device=x_t.device, dtype=torch.float32
|
| 9 |
+
)
|
| 10 |
+
return model(x_t, t_vec, cond, **kwargs)
|
| 11 |
+
|
| 12 |
+
@torch.no_grad()
|
| 13 |
+
def sample(
|
| 14 |
+
self,
|
| 15 |
+
model,
|
| 16 |
+
noise,
|
| 17 |
+
cond,
|
| 18 |
+
neg_cond,
|
| 19 |
+
steps=12,
|
| 20 |
+
cfg_strength=3.0,
|
| 21 |
+
rescale_t=1.0,
|
| 22 |
+
**kwargs,
|
| 23 |
+
):
|
| 24 |
+
x = noise
|
| 25 |
+
t_seq = np.linspace(1.0, 0.0, steps + 1)
|
| 26 |
+
t_seq = rescale_t * t_seq / (1 + (rescale_t - 1) * t_seq)
|
| 27 |
+
for i in range(steps):
|
| 28 |
+
t, t_prev = float(t_seq[i]), float(t_seq[i + 1])
|
| 29 |
+
v_cond = self._pred_v(model, x, t, cond, **kwargs)
|
| 30 |
+
v_uncond = self._pred_v(model, x, t, neg_cond, **kwargs)
|
| 31 |
+
v = (1 + cfg_strength) * v_cond - cfg_strength * v_uncond
|
| 32 |
+
x = x - (t - t_prev) * v
|
| 33 |
+
return x
|
| 34 |
+
|
| 35 |
+
|
| 36 |
+
class TopoFlowEulerSampler:
|
| 37 |
+
def _pred_v(self, model, x_t, t, verts, mask, cond, cond_mask):
|
| 38 |
+
t_vec = torch.full((x_t.shape[0],), t, device=x_t.device, dtype=torch.float32)
|
| 39 |
+
return model(x_t, t_vec, verts=verts, mask=mask, cond=cond, cond_mask=cond_mask)
|
| 40 |
+
|
| 41 |
+
@torch.no_grad()
|
| 42 |
+
def sample(
|
| 43 |
+
self,
|
| 44 |
+
model,
|
| 45 |
+
noise,
|
| 46 |
+
verts,
|
| 47 |
+
mask,
|
| 48 |
+
cond=None,
|
| 49 |
+
cond_mask=None,
|
| 50 |
+
steps=50,
|
| 51 |
+
):
|
| 52 |
+
x = noise
|
| 53 |
+
t_seq = np.linspace(0.0, 1.0, steps + 1)
|
| 54 |
+
for i in range(steps):
|
| 55 |
+
t, t_next = float(t_seq[i]), float(t_seq[i + 1])
|
| 56 |
+
v = self._pred_v(model, x, t, verts, mask, cond, cond_mask)
|
| 57 |
+
x = x + (t_next - t) * v
|
| 58 |
+
return x
|
models/offset_head.py
ADDED
|
@@ -0,0 +1,22 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import torch.nn as nn
|
| 2 |
+
|
| 3 |
+
|
| 4 |
+
class OffsetHead(nn.Module):
|
| 5 |
+
def __init__(self, feat_dim: int, mlp_ratio: float = 4.0):
|
| 6 |
+
super().__init__()
|
| 7 |
+
self.mlp = nn.Sequential(
|
| 8 |
+
nn.Linear(feat_dim, int(feat_dim * mlp_ratio)),
|
| 9 |
+
nn.GELU(approximate="tanh"),
|
| 10 |
+
nn.Linear(int(feat_dim * mlp_ratio), 3),
|
| 11 |
+
nn.Tanh(),
|
| 12 |
+
)
|
| 13 |
+
|
| 14 |
+
def forward(self, vtx_feats):
|
| 15 |
+
"""
|
| 16 |
+
Input:
|
| 17 |
+
vtx_feats: [N, feat_dim]
|
| 18 |
+
Output:
|
| 19 |
+
offsets: [N, 3], in range (-1, 1)
|
| 20 |
+
"""
|
| 21 |
+
offsets = self.mlp(vtx_feats) # [N, 3], (-1, 1)
|
| 22 |
+
return offsets
|
models/topo_autoencoder.py
ADDED
|
@@ -0,0 +1,297 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from __future__ import annotations
|
| 2 |
+
|
| 3 |
+
from typing import List, Optional
|
| 4 |
+
|
| 5 |
+
import numpy as np
|
| 6 |
+
import torch
|
| 7 |
+
from torch import nn
|
| 8 |
+
|
| 9 |
+
from modules.pointnet import Pointnet
|
| 10 |
+
from modules.transformer.hybrid import (
|
| 11 |
+
HybridGraphFlashStack,
|
| 12 |
+
FlashVarlenTransformerBlock,
|
| 13 |
+
)
|
| 14 |
+
from modules.utils import manual_cast, str_to_dtype
|
| 15 |
+
from modules.transformer.blocks import (
|
| 16 |
+
PointEmbed,
|
| 17 |
+
RotaryPositionPhasesEmbedder,
|
| 18 |
+
MaskedTransformerCrossAttnBlock,
|
| 19 |
+
)
|
| 20 |
+
|
| 21 |
+
|
| 22 |
+
class TopologyEncoderHybrid(nn.Module):
|
| 23 |
+
def __init__(
|
| 24 |
+
self,
|
| 25 |
+
z_dim: int = 32,
|
| 26 |
+
hidden_dim: int = 384,
|
| 27 |
+
num_heads: int = 6,
|
| 28 |
+
num_discrete: int = 256,
|
| 29 |
+
dtype: str = "float32",
|
| 30 |
+
num_hybrid_stages: int = 2,
|
| 31 |
+
num_flash_per_stage: int = 1,
|
| 32 |
+
use_gradient_checkpointing: bool = False,
|
| 33 |
+
pc_cross_attn: bool = False,
|
| 34 |
+
):
|
| 35 |
+
super().__init__()
|
| 36 |
+
self.dtype = str_to_dtype(dtype)
|
| 37 |
+
self.num_discrete = num_discrete
|
| 38 |
+
self.pc_cross_attn = bool(pc_cross_attn)
|
| 39 |
+
|
| 40 |
+
head_dim = hidden_dim // num_heads
|
| 41 |
+
self.rope = RotaryPositionPhasesEmbedder(head_dim=head_dim, dim=3)
|
| 42 |
+
|
| 43 |
+
self.backbone = HybridGraphFlashStack(
|
| 44 |
+
hidden_size=hidden_dim,
|
| 45 |
+
num_heads=num_heads,
|
| 46 |
+
num_stages=num_hybrid_stages,
|
| 47 |
+
num_flash_per_stage=num_flash_per_stage,
|
| 48 |
+
gradient_checkpointing=use_gradient_checkpointing,
|
| 49 |
+
)
|
| 50 |
+
if self.pc_cross_attn:
|
| 51 |
+
self.pc_cross_blocks = nn.ModuleList(
|
| 52 |
+
[
|
| 53 |
+
MaskedTransformerCrossAttnBlock(
|
| 54 |
+
hidden_dim, num_heads, cond_dim=hidden_dim
|
| 55 |
+
)
|
| 56 |
+
for _ in range(num_hybrid_stages)
|
| 57 |
+
]
|
| 58 |
+
)
|
| 59 |
+
self.z_proj = nn.Linear(hidden_dim, z_dim * 2)
|
| 60 |
+
nn.init.zeros_(self.z_proj.weight)
|
| 61 |
+
nn.init.zeros_(self.z_proj.bias)
|
| 62 |
+
|
| 63 |
+
def forward(
|
| 64 |
+
self,
|
| 65 |
+
verts: torch.Tensor,
|
| 66 |
+
pc_tokens: Optional[torch.Tensor],
|
| 67 |
+
point_embedder: PointEmbed,
|
| 68 |
+
verts_mask: Optional[torch.Tensor] = None,
|
| 69 |
+
adj_matrix: Optional[torch.Tensor] = None,
|
| 70 |
+
):
|
| 71 |
+
rope_phases = self.rope(verts.long())
|
| 72 |
+
coords = (verts + 0.5) / self.num_discrete * 2 - 1
|
| 73 |
+
vert_tokens = point_embedder(coords)
|
| 74 |
+
vert_tokens = manual_cast(vert_tokens, self.dtype)
|
| 75 |
+
|
| 76 |
+
adj_mask = None
|
| 77 |
+
if adj_matrix is not None:
|
| 78 |
+
b, n, _ = adj_matrix.shape
|
| 79 |
+
eye = torch.eye(n, device=adj_matrix.device, dtype=torch.bool).unsqueeze(0)
|
| 80 |
+
adj_mask = adj_matrix.bool() | eye
|
| 81 |
+
|
| 82 |
+
if self.pc_cross_attn:
|
| 83 |
+
if pc_tokens is None:
|
| 84 |
+
raise ValueError(
|
| 85 |
+
"pc_tokens required when encoder pc_cross_attn is enabled"
|
| 86 |
+
)
|
| 87 |
+
for stage, ca_block in zip(self.backbone.stages, self.pc_cross_blocks):
|
| 88 |
+
vert_tokens = stage(
|
| 89 |
+
vert_tokens,
|
| 90 |
+
x_mask=verts_mask,
|
| 91 |
+
adj_matrix=adj_mask,
|
| 92 |
+
rope_phases=rope_phases,
|
| 93 |
+
)
|
| 94 |
+
vert_tokens = ca_block(
|
| 95 |
+
vert_tokens,
|
| 96 |
+
pc_tokens,
|
| 97 |
+
x_mask=verts_mask,
|
| 98 |
+
c_mask=None,
|
| 99 |
+
)
|
| 100 |
+
else:
|
| 101 |
+
vert_tokens = self.backbone(
|
| 102 |
+
vert_tokens,
|
| 103 |
+
x_mask=verts_mask,
|
| 104 |
+
adj_matrix=adj_mask,
|
| 105 |
+
rope_phases=rope_phases,
|
| 106 |
+
)
|
| 107 |
+
z = self.z_proj(vert_tokens)
|
| 108 |
+
z = manual_cast(z, self.dtype)
|
| 109 |
+
return z
|
| 110 |
+
|
| 111 |
+
|
| 112 |
+
class TopologyDecoderHybrid(nn.Module):
|
| 113 |
+
def __init__(
|
| 114 |
+
self,
|
| 115 |
+
z_dim: int = 32,
|
| 116 |
+
hidden_dim: int = 384,
|
| 117 |
+
num_heads: int = 6,
|
| 118 |
+
num_discrete: int = 256,
|
| 119 |
+
dtype: str = "float32",
|
| 120 |
+
num_hybrid_stages: int = 2,
|
| 121 |
+
num_flash_per_stage: int = 1,
|
| 122 |
+
use_gradient_checkpointing: bool = False,
|
| 123 |
+
):
|
| 124 |
+
super().__init__()
|
| 125 |
+
self.dtype = str_to_dtype(dtype)
|
| 126 |
+
self.num_discrete = num_discrete
|
| 127 |
+
self.input_proj = nn.Linear(z_dim, hidden_dim)
|
| 128 |
+
self.backbone = HybridGraphFlashStack(
|
| 129 |
+
hidden_size=hidden_dim,
|
| 130 |
+
num_heads=num_heads,
|
| 131 |
+
num_stages=num_hybrid_stages,
|
| 132 |
+
num_flash_per_stage=num_flash_per_stage,
|
| 133 |
+
gradient_checkpointing=use_gradient_checkpointing,
|
| 134 |
+
)
|
| 135 |
+
|
| 136 |
+
def forward(self, z: torch.Tensor, verts_mask: Optional[torch.Tensor] = None):
|
| 137 |
+
h = self.input_proj(z)
|
| 138 |
+
h = manual_cast(h, self.dtype)
|
| 139 |
+
h = self.backbone(
|
| 140 |
+
h,
|
| 141 |
+
x_mask=verts_mask,
|
| 142 |
+
adj_matrix=None,
|
| 143 |
+
rope_phases=None,
|
| 144 |
+
)
|
| 145 |
+
return h
|
| 146 |
+
|
| 147 |
+
|
| 148 |
+
class TopologyConnectionPredictor(nn.Module):
|
| 149 |
+
def __init__(self, hidden_dim: int = 384):
|
| 150 |
+
super().__init__()
|
| 151 |
+
self.mlp = nn.Sequential(
|
| 152 |
+
nn.Linear(hidden_dim * 2, 256),
|
| 153 |
+
nn.GELU(),
|
| 154 |
+
nn.Linear(256, 1),
|
| 155 |
+
)
|
| 156 |
+
|
| 157 |
+
def forward(self, vert_feat_u: torch.Tensor, vert_feat_v: torch.Tensor):
|
| 158 |
+
pair_feat_0 = torch.cat([vert_feat_u, vert_feat_v], dim=-1)
|
| 159 |
+
pair_feat_1 = torch.cat([vert_feat_v, vert_feat_u], dim=-1)
|
| 160 |
+
h = (self.mlp(pair_feat_0) + self.mlp(pair_feat_1)) / 2.0
|
| 161 |
+
return h.squeeze(-1)
|
| 162 |
+
|
| 163 |
+
|
| 164 |
+
class TopologyVAE(nn.Module):
|
| 165 |
+
def __init__(
|
| 166 |
+
self,
|
| 167 |
+
z_dim: int = 32,
|
| 168 |
+
hidden_dim: int = 384,
|
| 169 |
+
pc_dim: int = 15,
|
| 170 |
+
inner_pc_dim: int = 256,
|
| 171 |
+
num_heads: int = 6,
|
| 172 |
+
num_discrete: int = 256,
|
| 173 |
+
dtype: str = "float32",
|
| 174 |
+
num_hybrid_stages: int = 2,
|
| 175 |
+
num_flash_per_stage: int = 1,
|
| 176 |
+
num_connection_blocks: Optional[int] = None,
|
| 177 |
+
use_gradient_checkpointing: bool = False,
|
| 178 |
+
encoder_pc_cross_attn: bool = False,
|
| 179 |
+
):
|
| 180 |
+
super().__init__()
|
| 181 |
+
self.dtype = str_to_dtype(dtype)
|
| 182 |
+
self.num_discrete = num_discrete
|
| 183 |
+
self.encoder_pc_cross_attn = bool(encoder_pc_cross_attn)
|
| 184 |
+
|
| 185 |
+
self.point_embed = PointEmbed(hidden_dim=hidden_dim, dim=hidden_dim)
|
| 186 |
+
self.point_net = Pointnet(
|
| 187 |
+
in_channels=pc_dim,
|
| 188 |
+
out_channels=inner_pc_dim,
|
| 189 |
+
hidden_dim=256,
|
| 190 |
+
n_blocks=5,
|
| 191 |
+
)
|
| 192 |
+
self.point_fusion = nn.Linear(hidden_dim + inner_pc_dim, hidden_dim)
|
| 193 |
+
|
| 194 |
+
self.encoder = TopologyEncoderHybrid(
|
| 195 |
+
z_dim=z_dim,
|
| 196 |
+
hidden_dim=hidden_dim,
|
| 197 |
+
num_heads=num_heads,
|
| 198 |
+
num_discrete=num_discrete,
|
| 199 |
+
dtype=dtype,
|
| 200 |
+
num_hybrid_stages=num_hybrid_stages,
|
| 201 |
+
num_flash_per_stage=num_flash_per_stage,
|
| 202 |
+
use_gradient_checkpointing=use_gradient_checkpointing,
|
| 203 |
+
pc_cross_attn=self.encoder_pc_cross_attn,
|
| 204 |
+
)
|
| 205 |
+
self.decoder = TopologyDecoderHybrid(
|
| 206 |
+
z_dim=z_dim,
|
| 207 |
+
hidden_dim=hidden_dim,
|
| 208 |
+
num_heads=num_heads,
|
| 209 |
+
num_discrete=num_discrete,
|
| 210 |
+
dtype=dtype,
|
| 211 |
+
num_hybrid_stages=num_hybrid_stages,
|
| 212 |
+
num_flash_per_stage=num_flash_per_stage,
|
| 213 |
+
use_gradient_checkpointing=use_gradient_checkpointing,
|
| 214 |
+
)
|
| 215 |
+
self.connection_predictor = TopologyConnectionPredictor(hidden_dim=hidden_dim)
|
| 216 |
+
|
| 217 |
+
head_dim = hidden_dim // num_heads
|
| 218 |
+
self.connection_rope = RotaryPositionPhasesEmbedder(head_dim=head_dim, dim=3)
|
| 219 |
+
|
| 220 |
+
n_conn = (
|
| 221 |
+
num_connection_blocks
|
| 222 |
+
if num_connection_blocks is not None
|
| 223 |
+
else (num_hybrid_stages * (1 + num_flash_per_stage))
|
| 224 |
+
)
|
| 225 |
+
self.connection_transformer_blocks = nn.ModuleList(
|
| 226 |
+
[
|
| 227 |
+
FlashVarlenTransformerBlock(
|
| 228 |
+
hidden_dim,
|
| 229 |
+
num_heads,
|
| 230 |
+
gradient_checkpointing=use_gradient_checkpointing,
|
| 231 |
+
)
|
| 232 |
+
for _ in range(n_conn)
|
| 233 |
+
]
|
| 234 |
+
)
|
| 235 |
+
|
| 236 |
+
def encode(
|
| 237 |
+
self,
|
| 238 |
+
verts: torch.Tensor,
|
| 239 |
+
pc_tokens: torch.Tensor,
|
| 240 |
+
verts_mask: Optional[torch.Tensor] = None,
|
| 241 |
+
adj_matrix: Optional[torch.Tensor] = None,
|
| 242 |
+
):
|
| 243 |
+
moments = self.encoder(
|
| 244 |
+
verts=verts,
|
| 245 |
+
pc_tokens=pc_tokens,
|
| 246 |
+
point_embedder=self.point_embed,
|
| 247 |
+
verts_mask=verts_mask,
|
| 248 |
+
adj_matrix=adj_matrix,
|
| 249 |
+
)
|
| 250 |
+
mean, logvar = moments.chunk(2, dim=-1)
|
| 251 |
+
return mean, logvar
|
| 252 |
+
|
| 253 |
+
def decode(
|
| 254 |
+
self,
|
| 255 |
+
z: torch.Tensor,
|
| 256 |
+
verts: torch.Tensor,
|
| 257 |
+
verts_mask: Optional[torch.Tensor] = None,
|
| 258 |
+
chunk_size: int = 20000,
|
| 259 |
+
threshold: float = 0.0,
|
| 260 |
+
) -> List[np.ndarray]:
|
| 261 |
+
# return: list of [N, 2] numpy arrays of predicted edges for each batch item
|
| 262 |
+
verts_feat = self.decoder(z=z, verts_mask=verts_mask)
|
| 263 |
+
|
| 264 |
+
rope = self.connection_rope(verts.long())
|
| 265 |
+
for block in self.connection_transformer_blocks:
|
| 266 |
+
verts_feat = block(verts_feat, verts_mask, rope_phases=rope)
|
| 267 |
+
|
| 268 |
+
all_pred_edges_list = []
|
| 269 |
+
for i in range(verts_mask.shape[0]):
|
| 270 |
+
valid = verts_mask[i]
|
| 271 |
+
valid_verts = verts[i][valid]
|
| 272 |
+
valid_feats = verts_feat[i][valid]
|
| 273 |
+
num_valid = int(valid_verts.shape[0])
|
| 274 |
+
|
| 275 |
+
u_idx, v_idx = torch.triu_indices(
|
| 276 |
+
num_valid, num_valid, offset=1, device=z.device
|
| 277 |
+
)
|
| 278 |
+
pred_edges_list = []
|
| 279 |
+
for i in range(0, u_idx.numel(), chunk_size):
|
| 280 |
+
cu = u_idx[i : i + chunk_size]
|
| 281 |
+
cv = v_idx[i : i + chunk_size]
|
| 282 |
+
logits = self.connection_predictor(
|
| 283 |
+
valid_feats[cu].unsqueeze(0),
|
| 284 |
+
valid_feats[cv].unsqueeze(0),
|
| 285 |
+
).squeeze(0)
|
| 286 |
+
take = logits > threshold
|
| 287 |
+
if bool(take.any()):
|
| 288 |
+
pred_edges_list.append(torch.stack([cu[take], cv[take]], dim=-1))
|
| 289 |
+
|
| 290 |
+
pred_edges = (
|
| 291 |
+
torch.cat(pred_edges_list, dim=0).cpu().numpy()
|
| 292 |
+
if pred_edges_list
|
| 293 |
+
else np.empty((0, 2), dtype=np.int64)
|
| 294 |
+
)
|
| 295 |
+
all_pred_edges_list.append(pred_edges)
|
| 296 |
+
|
| 297 |
+
return all_pred_edges_list
|
models/topo_flow.py
ADDED
|
@@ -0,0 +1,400 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from __future__ import annotations
|
| 2 |
+
|
| 3 |
+
import torch
|
| 4 |
+
import torch.nn as nn
|
| 5 |
+
import torch.nn.functional as F
|
| 6 |
+
from torch.utils.checkpoint import checkpoint
|
| 7 |
+
|
| 8 |
+
from modules.transformer.blocks import RotaryPositionPhasesEmbedder, TimestepEmbedder
|
| 9 |
+
from modules.attention import (
|
| 10 |
+
can_flash_varlen,
|
| 11 |
+
flash_varlen_self_attention,
|
| 12 |
+
flash_varlen_cross_attention,
|
| 13 |
+
sdpa_padding_mask,
|
| 14 |
+
)
|
| 15 |
+
from modules.norm import RMSNorm
|
| 16 |
+
from modules.utils import modulate
|
| 17 |
+
|
| 18 |
+
|
| 19 |
+
class TopologySiTBlockFlashVarlen(nn.Module):
|
| 20 |
+
def __init__(
|
| 21 |
+
self,
|
| 22 |
+
hidden_size: int,
|
| 23 |
+
num_heads: int,
|
| 24 |
+
mlp_ratio: float = 4.0,
|
| 25 |
+
dropout: float = 0.0,
|
| 26 |
+
gradient_checkpointing: bool = False,
|
| 27 |
+
qk_norm_eps: float = 1e-5,
|
| 28 |
+
qk_norm_variance_in_fp32: bool = True,
|
| 29 |
+
with_cross_attn: bool = False,
|
| 30 |
+
):
|
| 31 |
+
super().__init__()
|
| 32 |
+
if hidden_size % num_heads != 0:
|
| 33 |
+
raise ValueError(
|
| 34 |
+
f"hidden_size {hidden_size} not divisible by num_heads {num_heads}"
|
| 35 |
+
)
|
| 36 |
+
self.hidden_size = hidden_size
|
| 37 |
+
self.num_heads = num_heads
|
| 38 |
+
self.head_dim = hidden_size // num_heads
|
| 39 |
+
self.with_cross_attn = bool(with_cross_attn)
|
| 40 |
+
self.norm1 = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6)
|
| 41 |
+
self.norm2 = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6)
|
| 42 |
+
self.qkv = nn.Linear(hidden_size, hidden_size * 3, bias=True)
|
| 43 |
+
self.proj_out = nn.Linear(hidden_size, hidden_size, bias=True)
|
| 44 |
+
mlp_hidden = int(hidden_size * mlp_ratio)
|
| 45 |
+
self.mlp = nn.Sequential(
|
| 46 |
+
nn.Linear(hidden_size, mlp_hidden, bias=True),
|
| 47 |
+
nn.GELU(approximate="tanh"),
|
| 48 |
+
nn.Linear(mlp_hidden, hidden_size, bias=True),
|
| 49 |
+
)
|
| 50 |
+
self._n_adaln_chunks = 7 if self.with_cross_attn else 6
|
| 51 |
+
self.adaLN_modulation = nn.Sequential(
|
| 52 |
+
nn.SiLU(),
|
| 53 |
+
nn.Linear(hidden_size, self._n_adaln_chunks * hidden_size, bias=True),
|
| 54 |
+
)
|
| 55 |
+
self.dropout = dropout
|
| 56 |
+
self.gradient_checkpointing = bool(gradient_checkpointing)
|
| 57 |
+
|
| 58 |
+
self.norm_q, self.norm_k = (
|
| 59 |
+
RMSNorm(
|
| 60 |
+
self.head_dim,
|
| 61 |
+
eps=qk_norm_eps,
|
| 62 |
+
elementwise_affine=True,
|
| 63 |
+
variance_in_fp32=qk_norm_variance_in_fp32,
|
| 64 |
+
),
|
| 65 |
+
RMSNorm(
|
| 66 |
+
self.head_dim,
|
| 67 |
+
eps=qk_norm_eps,
|
| 68 |
+
elementwise_affine=True,
|
| 69 |
+
variance_in_fp32=qk_norm_variance_in_fp32,
|
| 70 |
+
),
|
| 71 |
+
)
|
| 72 |
+
|
| 73 |
+
if self.with_cross_attn:
|
| 74 |
+
self.norm_ca = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6)
|
| 75 |
+
self.q_ca = nn.Linear(hidden_size, hidden_size, bias=True)
|
| 76 |
+
self.kv_ca = nn.Linear(hidden_size, hidden_size * 2, bias=True)
|
| 77 |
+
self.proj_ca_out = nn.Linear(hidden_size, hidden_size, bias=True)
|
| 78 |
+
self.norm_q_ca, self.norm_k_ca = (
|
| 79 |
+
RMSNorm(
|
| 80 |
+
self.head_dim,
|
| 81 |
+
eps=qk_norm_eps,
|
| 82 |
+
elementwise_affine=True,
|
| 83 |
+
variance_in_fp32=qk_norm_variance_in_fp32,
|
| 84 |
+
),
|
| 85 |
+
RMSNorm(
|
| 86 |
+
self.head_dim,
|
| 87 |
+
eps=qk_norm_eps,
|
| 88 |
+
elementwise_affine=True,
|
| 89 |
+
variance_in_fp32=qk_norm_variance_in_fp32,
|
| 90 |
+
),
|
| 91 |
+
)
|
| 92 |
+
|
| 93 |
+
def _forward_once(
|
| 94 |
+
self,
|
| 95 |
+
x: torch.Tensor,
|
| 96 |
+
c: torch.Tensor,
|
| 97 |
+
key_padding_mask: torch.Tensor | None,
|
| 98 |
+
rope_phases: torch.Tensor,
|
| 99 |
+
cond_emb: torch.Tensor | None,
|
| 100 |
+
cond_mask: torch.Tensor | None,
|
| 101 |
+
) -> torch.Tensor:
|
| 102 |
+
chunks = self.adaLN_modulation(c).chunk(self._n_adaln_chunks, dim=1)
|
| 103 |
+
if self.with_cross_attn:
|
| 104 |
+
shift_msa, scale_msa, gate_msa, gate_mca, shift_mlp, scale_mlp, gate_mlp = (
|
| 105 |
+
chunks
|
| 106 |
+
)
|
| 107 |
+
else:
|
| 108 |
+
shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = chunks
|
| 109 |
+
|
| 110 |
+
h = modulate(self.norm1(x), shift_msa, scale_msa)
|
| 111 |
+
b, n, d = h.shape
|
| 112 |
+
qkv = (
|
| 113 |
+
self.qkv(h)
|
| 114 |
+
.view(b, n, 3, self.num_heads, self.head_dim)
|
| 115 |
+
.permute(2, 0, 3, 1, 4)
|
| 116 |
+
)
|
| 117 |
+
q, k, v = qkv[0], qkv[1], qkv[2]
|
| 118 |
+
if self.norm_q is not None:
|
| 119 |
+
q = self.norm_q(q)
|
| 120 |
+
if self.norm_k is not None:
|
| 121 |
+
k = self.norm_k(k)
|
| 122 |
+
q = RotaryPositionPhasesEmbedder.apply_rotary_embedding(q, rope_phases)
|
| 123 |
+
k = RotaryPositionPhasesEmbedder.apply_rotary_embedding(k, rope_phases)
|
| 124 |
+
|
| 125 |
+
x_mask = None if key_padding_mask is None else ~key_padding_mask.bool()
|
| 126 |
+
if can_flash_varlen(q, x_mask):
|
| 127 |
+
attn_out = flash_varlen_self_attention(q, k, v, x_mask)
|
| 128 |
+
elif x_mask is not None:
|
| 129 |
+
attn_out = F.scaled_dot_product_attention(
|
| 130 |
+
q,
|
| 131 |
+
k,
|
| 132 |
+
v,
|
| 133 |
+
attn_mask=sdpa_padding_mask(x_mask),
|
| 134 |
+
dropout_p=self.dropout if self.training else 0.0,
|
| 135 |
+
)
|
| 136 |
+
else:
|
| 137 |
+
attn_out = F.scaled_dot_product_attention(
|
| 138 |
+
q,
|
| 139 |
+
k,
|
| 140 |
+
v,
|
| 141 |
+
attn_mask=None,
|
| 142 |
+
dropout_p=self.dropout if self.training else 0.0,
|
| 143 |
+
)
|
| 144 |
+
|
| 145 |
+
attn_out = attn_out.transpose(1, 2).reshape(b, n, d)
|
| 146 |
+
x = x + gate_msa.unsqueeze(1) * self.proj_out(attn_out)
|
| 147 |
+
|
| 148 |
+
if self.with_cross_attn and cond_emb is not None:
|
| 149 |
+
h_ca = self.norm_ca(x)
|
| 150 |
+
nk = cond_emb.shape[1]
|
| 151 |
+
q_ca = (
|
| 152 |
+
self.q_ca(h_ca)
|
| 153 |
+
.view(b, n, self.num_heads, self.head_dim)
|
| 154 |
+
.transpose(1, 2)
|
| 155 |
+
)
|
| 156 |
+
kv_ca = (
|
| 157 |
+
self.kv_ca(cond_emb)
|
| 158 |
+
.view(b, nk, 2, self.num_heads, self.head_dim)
|
| 159 |
+
.permute(2, 0, 3, 1, 4)
|
| 160 |
+
)
|
| 161 |
+
k_ca, v_ca = kv_ca[0], kv_ca[1]
|
| 162 |
+
if self.norm_q_ca is not None:
|
| 163 |
+
q_ca = self.norm_q_ca(q_ca)
|
| 164 |
+
if self.norm_k_ca is not None:
|
| 165 |
+
k_ca = self.norm_k_ca(k_ca)
|
| 166 |
+
|
| 167 |
+
q_mask_bool = (
|
| 168 |
+
torch.ones(b, n, dtype=torch.bool, device=q_ca.device)
|
| 169 |
+
if key_padding_mask is None
|
| 170 |
+
else ~key_padding_mask.bool()
|
| 171 |
+
)
|
| 172 |
+
k_mask_bool = (
|
| 173 |
+
torch.ones(b, nk, dtype=torch.bool, device=q_ca.device)
|
| 174 |
+
if cond_mask is None
|
| 175 |
+
else cond_mask.bool()
|
| 176 |
+
)
|
| 177 |
+
|
| 178 |
+
if can_flash_varlen(q_ca, q_mask_bool):
|
| 179 |
+
ca_out = flash_varlen_cross_attention(
|
| 180 |
+
q_ca, k_ca, v_ca, q_mask_bool, k_mask_bool
|
| 181 |
+
)
|
| 182 |
+
else:
|
| 183 |
+
k_attn_mask = k_mask_bool.view(b, 1, 1, nk)
|
| 184 |
+
ca_out = F.scaled_dot_product_attention(
|
| 185 |
+
q_ca, k_ca, v_ca, attn_mask=k_attn_mask, dropout_p=0.0
|
| 186 |
+
)
|
| 187 |
+
ca_out = ca_out.transpose(1, 2).reshape(b, n, d)
|
| 188 |
+
x = x + gate_mca.unsqueeze(1) * self.proj_ca_out(ca_out)
|
| 189 |
+
|
| 190 |
+
h2 = modulate(self.norm2(x), shift_mlp, scale_mlp)
|
| 191 |
+
x = x + gate_mlp.unsqueeze(1) * self.mlp(h2)
|
| 192 |
+
if key_padding_mask is not None:
|
| 193 |
+
valid = ~key_padding_mask.bool()
|
| 194 |
+
x = torch.where(valid.unsqueeze(-1), x, torch.zeros_like(x))
|
| 195 |
+
return x
|
| 196 |
+
|
| 197 |
+
def forward(
|
| 198 |
+
self,
|
| 199 |
+
x: torch.Tensor,
|
| 200 |
+
c: torch.Tensor,
|
| 201 |
+
key_padding_mask: torch.Tensor | None,
|
| 202 |
+
rope_phases: torch.Tensor,
|
| 203 |
+
cond_emb: torch.Tensor | None = None,
|
| 204 |
+
cond_mask: torch.Tensor | None = None,
|
| 205 |
+
) -> torch.Tensor:
|
| 206 |
+
if self.training and self.gradient_checkpointing:
|
| 207 |
+
return checkpoint(
|
| 208 |
+
self._forward_once,
|
| 209 |
+
x,
|
| 210 |
+
c,
|
| 211 |
+
key_padding_mask,
|
| 212 |
+
rope_phases,
|
| 213 |
+
cond_emb,
|
| 214 |
+
cond_mask,
|
| 215 |
+
use_reentrant=False,
|
| 216 |
+
)
|
| 217 |
+
return self._forward_once(
|
| 218 |
+
x, c, key_padding_mask, rope_phases, cond_emb, cond_mask
|
| 219 |
+
)
|
| 220 |
+
|
| 221 |
+
|
| 222 |
+
class TopologyFinalLayer(nn.Module):
|
| 223 |
+
def __init__(self, hidden_size: int, out_channels: int):
|
| 224 |
+
super().__init__()
|
| 225 |
+
self.norm_final = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6)
|
| 226 |
+
self.linear = nn.Linear(hidden_size, out_channels, bias=True)
|
| 227 |
+
self.adaLN_modulation = nn.Sequential(
|
| 228 |
+
nn.SiLU(),
|
| 229 |
+
nn.Linear(hidden_size, 2 * hidden_size, bias=True),
|
| 230 |
+
)
|
| 231 |
+
|
| 232 |
+
def forward(self, x: torch.Tensor, c: torch.Tensor) -> torch.Tensor:
|
| 233 |
+
shift, scale = self.adaLN_modulation(c).chunk(2, dim=1)
|
| 234 |
+
x = modulate(self.norm_final(x), shift, scale)
|
| 235 |
+
return self.linear(x)
|
| 236 |
+
|
| 237 |
+
|
| 238 |
+
def _get_1d_sincos_embed(n: int, dim: int) -> torch.Tensor:
|
| 239 |
+
assert dim % 2 == 0
|
| 240 |
+
pos = torch.arange(n, dtype=torch.float32)
|
| 241 |
+
omega = torch.arange(dim // 2, dtype=torch.float32) / (dim // 2)
|
| 242 |
+
omega = 1.0 / (10000**omega)
|
| 243 |
+
out = pos[:, None] * omega[None, :]
|
| 244 |
+
return torch.cat([torch.sin(out), torch.cos(out)], dim=-1)
|
| 245 |
+
|
| 246 |
+
|
| 247 |
+
class TopologySiTFlow(nn.Module):
|
| 248 |
+
def __init__(
|
| 249 |
+
self,
|
| 250 |
+
z_dim: int,
|
| 251 |
+
hidden_size: int = 768,
|
| 252 |
+
depth: int = 12,
|
| 253 |
+
num_heads: int = 12,
|
| 254 |
+
mlp_ratio: float = 4.0,
|
| 255 |
+
max_vertices: int = 8192,
|
| 256 |
+
num_discrete: int = 1024,
|
| 257 |
+
dropout: float = 0.0,
|
| 258 |
+
gradient_checkpointing: bool = False,
|
| 259 |
+
cond_in_dim: int = 0,
|
| 260 |
+
cond_dropout_prob: float = 0.0,
|
| 261 |
+
qk_norm_eps: float = 1e-5,
|
| 262 |
+
qk_norm_variance_in_fp32: bool = True,
|
| 263 |
+
):
|
| 264 |
+
super().__init__()
|
| 265 |
+
self.z_dim = z_dim
|
| 266 |
+
self.hidden_size = hidden_size
|
| 267 |
+
self.max_vertices = max_vertices
|
| 268 |
+
self.num_discrete = int(num_discrete)
|
| 269 |
+
self.gradient_checkpointing = bool(gradient_checkpointing)
|
| 270 |
+
self.cond_in_dim = int(cond_in_dim)
|
| 271 |
+
self.cond_dropout_prob = float(cond_dropout_prob)
|
| 272 |
+
self.qk_norm_eps = float(qk_norm_eps)
|
| 273 |
+
self.qk_norm_variance_in_fp32 = bool(qk_norm_variance_in_fp32)
|
| 274 |
+
|
| 275 |
+
self.input_proj = nn.Linear(z_dim, hidden_size, bias=True)
|
| 276 |
+
self.coord_embed = nn.Sequential(
|
| 277 |
+
nn.Linear(3, hidden_size, bias=True),
|
| 278 |
+
nn.SiLU(),
|
| 279 |
+
nn.Linear(hidden_size, hidden_size, bias=True),
|
| 280 |
+
)
|
| 281 |
+
self.t_embedder = TimestepEmbedder(hidden_size)
|
| 282 |
+
self.rope = RotaryPositionPhasesEmbedder(
|
| 283 |
+
head_dim=hidden_size // num_heads, dim=3
|
| 284 |
+
)
|
| 285 |
+
|
| 286 |
+
if self.cond_in_dim > 0:
|
| 287 |
+
self.cond_proj = nn.Sequential(
|
| 288 |
+
nn.Linear(self.cond_in_dim, hidden_size, bias=True),
|
| 289 |
+
nn.SiLU(),
|
| 290 |
+
nn.Linear(hidden_size, hidden_size, bias=True),
|
| 291 |
+
)
|
| 292 |
+
self.null_token = nn.Parameter(torch.zeros(hidden_size))
|
| 293 |
+
else:
|
| 294 |
+
self.cond_proj = None
|
| 295 |
+
self.null_token = None
|
| 296 |
+
|
| 297 |
+
pe = _get_1d_sincos_embed(max_vertices, hidden_size)
|
| 298 |
+
self.register_buffer("pos_embed", pe.unsqueeze(0), persistent=False)
|
| 299 |
+
|
| 300 |
+
with_cross_attn = self.cond_in_dim > 0
|
| 301 |
+
self.blocks = nn.ModuleList(
|
| 302 |
+
[
|
| 303 |
+
TopologySiTBlockFlashVarlen(
|
| 304 |
+
hidden_size=hidden_size,
|
| 305 |
+
num_heads=num_heads,
|
| 306 |
+
mlp_ratio=mlp_ratio,
|
| 307 |
+
dropout=dropout,
|
| 308 |
+
gradient_checkpointing=self.gradient_checkpointing,
|
| 309 |
+
with_cross_attn=with_cross_attn,
|
| 310 |
+
qk_norm_eps=self.qk_norm_eps,
|
| 311 |
+
qk_norm_variance_in_fp32=self.qk_norm_variance_in_fp32,
|
| 312 |
+
)
|
| 313 |
+
for _ in range(depth)
|
| 314 |
+
]
|
| 315 |
+
)
|
| 316 |
+
self.final_layer = TopologyFinalLayer(hidden_size, z_dim)
|
| 317 |
+
|
| 318 |
+
def _prepare_cond(
|
| 319 |
+
self,
|
| 320 |
+
b: int,
|
| 321 |
+
cond: torch.Tensor | None,
|
| 322 |
+
cond_mask: torch.Tensor | None,
|
| 323 |
+
cond_drop_override: torch.Tensor | None,
|
| 324 |
+
device: torch.device,
|
| 325 |
+
) -> tuple[torch.Tensor | None, torch.Tensor | None]:
|
| 326 |
+
if self.cond_proj is None:
|
| 327 |
+
return None, None
|
| 328 |
+
|
| 329 |
+
if cond is None:
|
| 330 |
+
null_emb = self.null_token.view(1, 1, -1).expand(b, 1, -1).contiguous()
|
| 331 |
+
mask_out = torch.ones(b, 1, dtype=torch.bool, device=device)
|
| 332 |
+
return null_emb, mask_out
|
| 333 |
+
|
| 334 |
+
if cond.dim() != 3 or cond.shape[0] != b or cond.shape[-1] != self.cond_in_dim:
|
| 335 |
+
raise ValueError(
|
| 336 |
+
f"cond shape {tuple(cond.shape)} expected ({b}, K, {self.cond_in_dim})"
|
| 337 |
+
)
|
| 338 |
+
k = cond.shape[1]
|
| 339 |
+
cond_emb = self.cond_proj(cond)
|
| 340 |
+
null = self.null_token.view(1, 1, -1).to(dtype=cond_emb.dtype)
|
| 341 |
+
|
| 342 |
+
if cond_mask is None:
|
| 343 |
+
mask_out = torch.ones(b, k, dtype=torch.bool, device=device)
|
| 344 |
+
else:
|
| 345 |
+
mask_out = cond_mask.to(device=device, dtype=torch.bool)
|
| 346 |
+
if mask_out.shape != (b, k):
|
| 347 |
+
raise ValueError(
|
| 348 |
+
f"cond_mask shape {tuple(mask_out.shape)} expected ({b}, {k})"
|
| 349 |
+
)
|
| 350 |
+
|
| 351 |
+
drop: torch.Tensor | None = None
|
| 352 |
+
if cond_drop_override is not None:
|
| 353 |
+
drop = cond_drop_override.to(device=device, dtype=torch.bool).reshape(b)
|
| 354 |
+
elif self.training and self.cond_dropout_prob > 0.0:
|
| 355 |
+
drop = torch.rand(b, device=device) < self.cond_dropout_prob
|
| 356 |
+
|
| 357 |
+
if drop is not None:
|
| 358 |
+
null_emb = null.expand(b, k, -1)
|
| 359 |
+
cond_emb = torch.where(drop.view(b, 1, 1), null_emb, cond_emb)
|
| 360 |
+
mask_out = torch.where(drop.view(b, 1), torch.ones_like(mask_out), mask_out)
|
| 361 |
+
return cond_emb, mask_out
|
| 362 |
+
|
| 363 |
+
def forward(
|
| 364 |
+
self,
|
| 365 |
+
x: torch.Tensor,
|
| 366 |
+
t: torch.Tensor,
|
| 367 |
+
verts: torch.Tensor,
|
| 368 |
+
mask: torch.Tensor,
|
| 369 |
+
cond: torch.Tensor | None = None,
|
| 370 |
+
cond_mask: torch.Tensor | None = None,
|
| 371 |
+
cond_drop_override: torch.Tensor | None = None,
|
| 372 |
+
) -> torch.Tensor:
|
| 373 |
+
b, n, _ = x.shape
|
| 374 |
+
if n > self.max_vertices:
|
| 375 |
+
raise ValueError(f"Sequence length {n} > max_vertices {self.max_vertices}")
|
| 376 |
+
if self.cond_in_dim == 0:
|
| 377 |
+
if cond is not None:
|
| 378 |
+
raise ValueError("TopologySiT(cond_in_dim=0): pass cond=None")
|
| 379 |
+
if cond_mask is not None:
|
| 380 |
+
raise ValueError("TopologySiT(cond_in_dim=0): cond_mask is unused")
|
| 381 |
+
if cond_drop_override is not None:
|
| 382 |
+
raise ValueError(
|
| 383 |
+
"TopologySiT(cond_in_dim=0): cond_drop_override is unused"
|
| 384 |
+
)
|
| 385 |
+
|
| 386 |
+
coords = ((verts.float() + 0.5) / self.num_discrete) * 2.0 - 1.0
|
| 387 |
+
h = self.input_proj(x) + self.pos_embed[:, :n, :] + self.coord_embed(coords)
|
| 388 |
+
c = self.t_embedder(t)
|
| 389 |
+
|
| 390 |
+
cond_emb, cond_mask_eff = self._prepare_cond(
|
| 391 |
+
b, cond, cond_mask, cond_drop_override, x.device
|
| 392 |
+
)
|
| 393 |
+
|
| 394 |
+
key_padding_mask = ~mask
|
| 395 |
+
rope_phases = self.rope(verts.long())
|
| 396 |
+
for block in self.blocks:
|
| 397 |
+
h = block(h, c, key_padding_mask, rope_phases, cond_emb, cond_mask_eff)
|
| 398 |
+
out = self.final_layer(h, c)
|
| 399 |
+
out = torch.where(mask.unsqueeze(-1), out, torch.zeros_like(out))
|
| 400 |
+
return out
|
models/vdf_encoder.py
ADDED
|
@@ -0,0 +1,55 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import torch
|
| 2 |
+
import torch.nn as nn
|
| 3 |
+
from torch.utils.checkpoint import checkpoint
|
| 4 |
+
|
| 5 |
+
from modules.pointnet import LocalPoolPointnet
|
| 6 |
+
|
| 7 |
+
|
| 8 |
+
class VDFEncoder(nn.Module):
|
| 9 |
+
def __init__(
|
| 10 |
+
self,
|
| 11 |
+
in_channels,
|
| 12 |
+
hidden_dim,
|
| 13 |
+
out_channels,
|
| 14 |
+
scatter_type,
|
| 15 |
+
n_blocks,
|
| 16 |
+
resolution=64,
|
| 17 |
+
use_checkpoint=False,
|
| 18 |
+
):
|
| 19 |
+
super().__init__()
|
| 20 |
+
self.pointnet = LocalPoolPointnet(
|
| 21 |
+
in_channels=in_channels,
|
| 22 |
+
out_channels=out_channels,
|
| 23 |
+
hidden_dim=hidden_dim,
|
| 24 |
+
n_blocks=n_blocks,
|
| 25 |
+
scatter_type=scatter_type,
|
| 26 |
+
)
|
| 27 |
+
|
| 28 |
+
self.resolution = resolution
|
| 29 |
+
self.use_checkpoint = use_checkpoint
|
| 30 |
+
|
| 31 |
+
def forward(
|
| 32 |
+
self,
|
| 33 |
+
p,
|
| 34 |
+
sparse_coords,
|
| 35 |
+
res=None,
|
| 36 |
+
bbox_size=(-0.5, 0.5),
|
| 37 |
+
):
|
| 38 |
+
"""
|
| 39 |
+
Input:
|
| 40 |
+
p: [N, in_channels]
|
| 41 |
+
sparse_coords: [M, 4], (b, z, y, x)
|
| 42 |
+
Output:
|
| 43 |
+
geo_feats: [N, out_channels]
|
| 44 |
+
"""
|
| 45 |
+
if res is None:
|
| 46 |
+
res = self.resolution
|
| 47 |
+
|
| 48 |
+
if self.use_checkpoint and self.training:
|
| 49 |
+
geo_feats = checkpoint(
|
| 50 |
+
self.pointnet, p, sparse_coords, res, bbox_size, use_reentrant=False
|
| 51 |
+
)
|
| 52 |
+
else:
|
| 53 |
+
geo_feats = self.pointnet(p, sparse_coords, res=res, bbox_size=bbox_size)
|
| 54 |
+
|
| 55 |
+
return geo_feats
|
models/vertex_autoencoder.py
ADDED
|
@@ -0,0 +1,595 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import torch
|
| 2 |
+
import torch.nn as nn
|
| 3 |
+
from typing import *
|
| 4 |
+
import torch.nn.functional as F
|
| 5 |
+
|
| 6 |
+
from modules import sparse as sp
|
| 7 |
+
from modules.sparse import SparseTensor
|
| 8 |
+
from modules.sparse.linear import SparseLinear
|
| 9 |
+
from modules.sparse.nonlinearity import SparseGELU
|
| 10 |
+
from modules.utils import (
|
| 11 |
+
zero_module,
|
| 12 |
+
convert_module_to_f16,
|
| 13 |
+
convert_module_to_f32,
|
| 14 |
+
flatten_coords,
|
| 15 |
+
per_batch_counts,
|
| 16 |
+
)
|
| 17 |
+
from modules.sparse.transformer import SparseTransformerBase, SparseTransformerCrossBase
|
| 18 |
+
from modules.sparse.blocks import SparseResBlock3d
|
| 19 |
+
from modules.utils import DiagonalGaussianDistribution
|
| 20 |
+
|
| 21 |
+
|
| 22 |
+
class SparseOccHead(nn.Module):
|
| 23 |
+
def __init__(self, channels: int, out_channels: int, mlp_ratio: float = 4.0):
|
| 24 |
+
super().__init__()
|
| 25 |
+
self.mlp = nn.Sequential(
|
| 26 |
+
SparseLinear(channels, int(channels * mlp_ratio)),
|
| 27 |
+
SparseGELU(approximate="tanh"),
|
| 28 |
+
SparseLinear(int(channels * mlp_ratio), out_channels),
|
| 29 |
+
)
|
| 30 |
+
|
| 31 |
+
def forward(self, x: sp.SparseTensor) -> sp.SparseTensor:
|
| 32 |
+
return self.mlp(x)
|
| 33 |
+
|
| 34 |
+
|
| 35 |
+
class SparseEncoderBlock(nn.Module):
|
| 36 |
+
def __init__(
|
| 37 |
+
self,
|
| 38 |
+
resolution: int,
|
| 39 |
+
in_channels: int,
|
| 40 |
+
model_channels: int,
|
| 41 |
+
num_blocks: int,
|
| 42 |
+
num_downsample: int = 4,
|
| 43 |
+
num_heads: Optional[int] = None,
|
| 44 |
+
num_head_channels: Optional[int] = 64,
|
| 45 |
+
mlp_ratio: float = 4,
|
| 46 |
+
attn_mode: Literal[
|
| 47 |
+
"full", "shift_window", "shift_sequence", "shift_order", "swin"
|
| 48 |
+
] = "swin",
|
| 49 |
+
window_size: int = 8,
|
| 50 |
+
pe_mode: Literal["ape", "rope"] = "ape",
|
| 51 |
+
use_fp16: bool = False,
|
| 52 |
+
use_checkpoint: bool = False,
|
| 53 |
+
qk_rms_norm: bool = False,
|
| 54 |
+
):
|
| 55 |
+
super().__init__()
|
| 56 |
+
self.resolution = resolution
|
| 57 |
+
|
| 58 |
+
self.self_attn = SparseTransformerBase(
|
| 59 |
+
in_channels=model_channels,
|
| 60 |
+
model_channels=model_channels,
|
| 61 |
+
num_blocks=num_blocks,
|
| 62 |
+
num_heads=num_heads,
|
| 63 |
+
num_head_channels=num_head_channels,
|
| 64 |
+
attn_mode=attn_mode,
|
| 65 |
+
window_size=window_size,
|
| 66 |
+
pe_mode=pe_mode,
|
| 67 |
+
mlp_ratio=mlp_ratio,
|
| 68 |
+
use_fp16=use_fp16,
|
| 69 |
+
use_checkpoint=use_checkpoint,
|
| 70 |
+
qk_rms_norm=qk_rms_norm,
|
| 71 |
+
)
|
| 72 |
+
|
| 73 |
+
self.input_layer1 = sp.SparseLinear(
|
| 74 |
+
in_channels, model_channels >> num_downsample
|
| 75 |
+
)
|
| 76 |
+
|
| 77 |
+
self.downsample = nn.ModuleList(
|
| 78 |
+
[
|
| 79 |
+
SparseResBlock3d(
|
| 80 |
+
channels=model_channels >> (i + 1),
|
| 81 |
+
out_channels=model_channels >> i,
|
| 82 |
+
downsample=True,
|
| 83 |
+
upsample=False,
|
| 84 |
+
use_checkpoint=use_checkpoint,
|
| 85 |
+
)
|
| 86 |
+
for i in range(num_downsample - 1, -1, -1)
|
| 87 |
+
]
|
| 88 |
+
)
|
| 89 |
+
|
| 90 |
+
def forward(
|
| 91 |
+
self,
|
| 92 |
+
x: SparseTensor,
|
| 93 |
+
):
|
| 94 |
+
"""
|
| 95 |
+
Input:
|
| 96 |
+
x: SparseTensor in N resolution, with feats of in_channels
|
| 97 |
+
Output:
|
| 98 |
+
h: SparseTensor in N>>num_downsample resolution, with feats of model_channels
|
| 99 |
+
"""
|
| 100 |
+
x = self.input_layer1(x)
|
| 101 |
+
for block in self.downsample:
|
| 102 |
+
x = block(x)
|
| 103 |
+
h = self.self_attn(x)
|
| 104 |
+
return h
|
| 105 |
+
|
| 106 |
+
|
| 107 |
+
class SparseDecoderUpsampleBlock(nn.Module):
|
| 108 |
+
def __init__(
|
| 109 |
+
self,
|
| 110 |
+
channels: int,
|
| 111 |
+
resolution: int,
|
| 112 |
+
out_channels: int,
|
| 113 |
+
model_channels: int = 512,
|
| 114 |
+
num_blocks: int = 4,
|
| 115 |
+
num_heads: int = 8,
|
| 116 |
+
mlp_ratio: float = 4.0,
|
| 117 |
+
num_groups: int = 32,
|
| 118 |
+
):
|
| 119 |
+
super().__init__()
|
| 120 |
+
self.channels = channels
|
| 121 |
+
self.resolution = resolution
|
| 122 |
+
self.out_resolution = resolution * 2
|
| 123 |
+
self.model_channels = model_channels
|
| 124 |
+
self.out_channels = out_channels
|
| 125 |
+
|
| 126 |
+
self.act_layers = nn.Sequential(
|
| 127 |
+
sp.SparseGroupNorm32(num_groups, channels), sp.SparseSiLU()
|
| 128 |
+
)
|
| 129 |
+
|
| 130 |
+
self.sub = sp.SparseSubdivide()
|
| 131 |
+
|
| 132 |
+
self.out_layers = nn.Sequential(
|
| 133 |
+
sp.SparseConv3d(
|
| 134 |
+
channels, self.out_channels, 3, indice_key=f"res_{self.out_resolution}"
|
| 135 |
+
),
|
| 136 |
+
sp.SparseGroupNorm32(num_groups, self.out_channels),
|
| 137 |
+
sp.SparseSiLU(),
|
| 138 |
+
zero_module(
|
| 139 |
+
sp.SparseConv3d(
|
| 140 |
+
self.out_channels,
|
| 141 |
+
self.out_channels,
|
| 142 |
+
3,
|
| 143 |
+
indice_key=f"res_{self.out_resolution}",
|
| 144 |
+
)
|
| 145 |
+
),
|
| 146 |
+
)
|
| 147 |
+
|
| 148 |
+
if self.out_channels == channels:
|
| 149 |
+
self.skip_connection = nn.Identity()
|
| 150 |
+
else:
|
| 151 |
+
self.skip_connection = sp.SparseConv3d(
|
| 152 |
+
channels, self.out_channels, 1, indice_key=f"res_{self.out_resolution}"
|
| 153 |
+
)
|
| 154 |
+
|
| 155 |
+
self.pruning_head = SparseOccHead(self.out_channels, out_channels=1)
|
| 156 |
+
|
| 157 |
+
self.ca = SparseTransformerCrossBase(
|
| 158 |
+
in_channels=self.out_channels,
|
| 159 |
+
model_channels=self.model_channels,
|
| 160 |
+
context_channels=self.model_channels,
|
| 161 |
+
num_blocks=num_blocks,
|
| 162 |
+
num_heads=num_heads,
|
| 163 |
+
mlp_ratio=mlp_ratio,
|
| 164 |
+
attn_mode="full",
|
| 165 |
+
pe_mode="ape",
|
| 166 |
+
use_checkpoint=True,
|
| 167 |
+
qk_rms_norm=False,
|
| 168 |
+
)
|
| 169 |
+
|
| 170 |
+
self.proj_ctx = sp.SparseLinear(self.out_channels, self.model_channels)
|
| 171 |
+
self.proj_out = sp.SparseLinear(self.model_channels, self.out_channels)
|
| 172 |
+
|
| 173 |
+
def forward(
|
| 174 |
+
self,
|
| 175 |
+
x: sp.SparseTensor,
|
| 176 |
+
training=False,
|
| 177 |
+
threshold=0.5,
|
| 178 |
+
) -> sp.SparseTensor:
|
| 179 |
+
h = self.act_layers(x)
|
| 180 |
+
h = self.sub(h)
|
| 181 |
+
x_sub = self.sub(x)
|
| 182 |
+
h = self.out_layers(h)
|
| 183 |
+
h = h + self.skip_connection(x_sub)
|
| 184 |
+
h = self.proj_out(self.ca(x=h, context=self.proj_ctx(h)))
|
| 185 |
+
|
| 186 |
+
occ_prob_q = self.pruning_head(h)
|
| 187 |
+
|
| 188 |
+
if training:
|
| 189 |
+
return h, occ_prob_q, [0]
|
| 190 |
+
|
| 191 |
+
scores_q = torch.sigmoid(occ_prob_q.feats).squeeze(-1)
|
| 192 |
+
N_full = h.feats.shape[0]
|
| 193 |
+
if N_full % 8 != 0:
|
| 194 |
+
raise ValueError(f"Number of nodes({N_full}) is not divisible by 8.")
|
| 195 |
+
|
| 196 |
+
# ensure at least one point is kept in each group of 8
|
| 197 |
+
n_parents = N_full // 8
|
| 198 |
+
|
| 199 |
+
scores_q_grouped = scores_q.view(n_parents, 8)
|
| 200 |
+
|
| 201 |
+
mask_grouped = scores_q_grouped >= threshold
|
| 202 |
+
|
| 203 |
+
none_survived = mask_grouped.sum(dim=1) == 0
|
| 204 |
+
|
| 205 |
+
# per-batch rescue counts; all 8 children of a parent share one batch index
|
| 206 |
+
if n_parents > 0:
|
| 207 |
+
parent_batch = h.coords[:, 0].view(n_parents, 8)[:, 0]
|
| 208 |
+
num_rescue = per_batch_counts(
|
| 209 |
+
parent_batch[none_survived], int(parent_batch.max().item()) + 1
|
| 210 |
+
)
|
| 211 |
+
else:
|
| 212 |
+
num_rescue = [0]
|
| 213 |
+
if none_survived.any():
|
| 214 |
+
failed_scores = scores_q_grouped[none_survived]
|
| 215 |
+
_, topk_indices = torch.topk(failed_scores, k=1, dim=1)
|
| 216 |
+
|
| 217 |
+
failed_row_idxs = torch.nonzero(none_survived, as_tuple=True)[0]
|
| 218 |
+
rows_expanded = failed_row_idxs.unsqueeze(1).expand(-1, 1)
|
| 219 |
+
|
| 220 |
+
mask_grouped[rows_expanded, topk_indices] = True
|
| 221 |
+
|
| 222 |
+
sub_mask = mask_grouped.view(-1)
|
| 223 |
+
|
| 224 |
+
h = sp.SparseTensor(feats=h.feats[sub_mask], coords=h.coords[sub_mask])
|
| 225 |
+
occ_prob_final = sp.SparseTensor(
|
| 226 |
+
feats=occ_prob_q.feats[sub_mask], coords=occ_prob_q.coords[sub_mask]
|
| 227 |
+
)
|
| 228 |
+
|
| 229 |
+
return h, occ_prob_final, num_rescue
|
| 230 |
+
|
| 231 |
+
|
| 232 |
+
class SparseDecoderBlock(nn.Module):
|
| 233 |
+
def __init__(
|
| 234 |
+
self,
|
| 235 |
+
resolution: int,
|
| 236 |
+
in_channels: int,
|
| 237 |
+
out_channels: int,
|
| 238 |
+
model_channels: int = 512,
|
| 239 |
+
num_blocks: int = 4,
|
| 240 |
+
num_heads: int = 8,
|
| 241 |
+
mlp_ratio: float = 4.0,
|
| 242 |
+
use_fp16: bool = False,
|
| 243 |
+
):
|
| 244 |
+
super().__init__()
|
| 245 |
+
self.resolution = resolution
|
| 246 |
+
|
| 247 |
+
self.upsample = SparseDecoderUpsampleBlock(
|
| 248 |
+
channels=in_channels,
|
| 249 |
+
resolution=resolution,
|
| 250 |
+
out_channels=out_channels,
|
| 251 |
+
num_blocks=num_blocks,
|
| 252 |
+
num_heads=num_heads,
|
| 253 |
+
mlp_ratio=mlp_ratio,
|
| 254 |
+
model_channels=model_channels,
|
| 255 |
+
num_groups=32,
|
| 256 |
+
)
|
| 257 |
+
|
| 258 |
+
if use_fp16:
|
| 259 |
+
self.convert_to_fp16()
|
| 260 |
+
|
| 261 |
+
def forward(
|
| 262 |
+
self,
|
| 263 |
+
x: sp.SparseTensor,
|
| 264 |
+
training: bool = False,
|
| 265 |
+
threshold: float = 0.5,
|
| 266 |
+
):
|
| 267 |
+
h = x
|
| 268 |
+
h = h.type(x.dtype)
|
| 269 |
+
h = h.replace(F.layer_norm(h.feats, h.feats.shape[-1:]))
|
| 270 |
+
h, occ_prob, num_rescue = self.upsample(
|
| 271 |
+
h,
|
| 272 |
+
training=training,
|
| 273 |
+
threshold=threshold,
|
| 274 |
+
)
|
| 275 |
+
return h, occ_prob, num_rescue
|
| 276 |
+
|
| 277 |
+
def convert_to_fp16(self):
|
| 278 |
+
"""Convert all components to float16"""
|
| 279 |
+
convert_module_to_f16(self.upsample)
|
| 280 |
+
|
| 281 |
+
def convert_to_fp32(self):
|
| 282 |
+
"""Convert all components to float32"""
|
| 283 |
+
convert_module_to_f32(self.upsample)
|
| 284 |
+
|
| 285 |
+
|
| 286 |
+
class VertexVAE(nn.Module):
|
| 287 |
+
def __init__(
|
| 288 |
+
self,
|
| 289 |
+
# Core architecture parameters
|
| 290 |
+
encoder_cfg: Dict = {},
|
| 291 |
+
expander_cfg: Dict = {},
|
| 292 |
+
decoder_cfg: List[Dict] = [],
|
| 293 |
+
# Shared transformer parameters
|
| 294 |
+
resolution: int = 1024,
|
| 295 |
+
num_head_channels: Optional[int] = 64,
|
| 296 |
+
mlp_ratio: float = 4.0,
|
| 297 |
+
attn_mode: str = "swin",
|
| 298 |
+
window_size: int = 8,
|
| 299 |
+
pe_mode: str = "ape",
|
| 300 |
+
use_fp16: bool = False,
|
| 301 |
+
use_checkpoint: bool = True,
|
| 302 |
+
qk_rms_norm: bool = False,
|
| 303 |
+
latent_dim: int = 8,
|
| 304 |
+
):
|
| 305 |
+
super().__init__()
|
| 306 |
+
self.latent_dim = latent_dim
|
| 307 |
+
self.decoder_cfg = decoder_cfg
|
| 308 |
+
|
| 309 |
+
self.encoder = SparseEncoderBlock(
|
| 310 |
+
resolution=resolution,
|
| 311 |
+
in_channels=encoder_cfg["in_channels"],
|
| 312 |
+
model_channels=encoder_cfg["model_channels"],
|
| 313 |
+
num_blocks=encoder_cfg["num_blocks"],
|
| 314 |
+
num_heads=encoder_cfg["num_heads"],
|
| 315 |
+
num_downsample=len(decoder_cfg),
|
| 316 |
+
num_head_channels=num_head_channels,
|
| 317 |
+
attn_mode=attn_mode,
|
| 318 |
+
window_size=window_size,
|
| 319 |
+
pe_mode=pe_mode,
|
| 320 |
+
mlp_ratio=mlp_ratio,
|
| 321 |
+
use_fp16=use_fp16,
|
| 322 |
+
use_checkpoint=use_checkpoint,
|
| 323 |
+
qk_rms_norm=qk_rms_norm,
|
| 324 |
+
)
|
| 325 |
+
|
| 326 |
+
self.latent_expander = SparseTransformerBase(
|
| 327 |
+
in_channels=latent_dim,
|
| 328 |
+
model_channels=expander_cfg["model_channels"],
|
| 329 |
+
num_blocks=expander_cfg["num_blocks"],
|
| 330 |
+
num_heads=expander_cfg["num_heads"],
|
| 331 |
+
num_head_channels=num_head_channels,
|
| 332 |
+
attn_mode=attn_mode,
|
| 333 |
+
window_size=window_size,
|
| 334 |
+
pe_mode=pe_mode,
|
| 335 |
+
mlp_ratio=mlp_ratio,
|
| 336 |
+
use_fp16=use_fp16,
|
| 337 |
+
use_checkpoint=use_checkpoint,
|
| 338 |
+
qk_rms_norm=qk_rms_norm,
|
| 339 |
+
)
|
| 340 |
+
|
| 341 |
+
self.vtx_proj = sp.SparseLinear(
|
| 342 |
+
expander_cfg["model_channels"], decoder_cfg[0]["in_channels"]
|
| 343 |
+
)
|
| 344 |
+
|
| 345 |
+
self.vtx_pruning_head = SparseOccHead(
|
| 346 |
+
expander_cfg["model_channels"], out_channels=1
|
| 347 |
+
)
|
| 348 |
+
|
| 349 |
+
self.out_layer = sp.SparseLinear(expander_cfg["model_channels"], latent_dim * 2)
|
| 350 |
+
|
| 351 |
+
self.decoder_vtx = nn.ModuleList()
|
| 352 |
+
self.decoder_vtx_ca = nn.ModuleList()
|
| 353 |
+
self.latent_proj = nn.ModuleList()
|
| 354 |
+
for config in decoder_cfg:
|
| 355 |
+
self.decoder_vtx.append(
|
| 356 |
+
# using default parameters to init the upsample block
|
| 357 |
+
SparseDecoderBlock(
|
| 358 |
+
resolution=config["resolution"],
|
| 359 |
+
in_channels=config["in_channels"],
|
| 360 |
+
out_channels=config["out_channels"],
|
| 361 |
+
num_blocks=config["num_blocks"],
|
| 362 |
+
num_heads=config["num_heads"],
|
| 363 |
+
use_fp16=use_fp16,
|
| 364 |
+
)
|
| 365 |
+
)
|
| 366 |
+
self.latent_proj.append(
|
| 367 |
+
sp.SparseLinear(latent_dim, config["context_channels"])
|
| 368 |
+
)
|
| 369 |
+
self.decoder_vtx_ca.append(
|
| 370 |
+
SparseTransformerCrossBase(
|
| 371 |
+
in_channels=config["out_channels"],
|
| 372 |
+
model_channels=config["model_channels"],
|
| 373 |
+
context_channels=config["context_channels"],
|
| 374 |
+
num_blocks=config["num_blocks"],
|
| 375 |
+
num_heads=config["num_heads"],
|
| 376 |
+
num_head_channels=num_head_channels,
|
| 377 |
+
mlp_ratio=mlp_ratio,
|
| 378 |
+
attn_mode="full",
|
| 379 |
+
window_size=window_size,
|
| 380 |
+
pe_mode=pe_mode,
|
| 381 |
+
use_fp16=use_fp16,
|
| 382 |
+
use_checkpoint=use_checkpoint,
|
| 383 |
+
qk_rms_norm=qk_rms_norm,
|
| 384 |
+
)
|
| 385 |
+
)
|
| 386 |
+
|
| 387 |
+
if use_fp16:
|
| 388 |
+
self.convert_to_fp16()
|
| 389 |
+
|
| 390 |
+
def encode(
|
| 391 |
+
self,
|
| 392 |
+
x: sp.SparseTensor,
|
| 393 |
+
sample_posterior=True,
|
| 394 |
+
) -> Tuple[torch.Tensor, torch.Tensor]:
|
| 395 |
+
h = self.encoder(x)
|
| 396 |
+
h = h.type(x.dtype)
|
| 397 |
+
h = h.replace(F.layer_norm(h.feats, h.feats.shape[-1:]))
|
| 398 |
+
h = self.out_layer(h)
|
| 399 |
+
|
| 400 |
+
posterior = DiagonalGaussianDistribution(h.feats, feat_dim=-1)
|
| 401 |
+
if sample_posterior:
|
| 402 |
+
z = posterior.sample()
|
| 403 |
+
else:
|
| 404 |
+
z = posterior.mode()
|
| 405 |
+
z = h.replace(z)
|
| 406 |
+
return z, posterior
|
| 407 |
+
|
| 408 |
+
def decode(
|
| 409 |
+
self,
|
| 410 |
+
latent_: sp.SparseTensor,
|
| 411 |
+
gt_vertex_voxels_list: List[sp.SparseTensor],
|
| 412 |
+
training=True,
|
| 413 |
+
inference_threshold=0.5,
|
| 414 |
+
verbose=False,
|
| 415 |
+
) -> List[Dict]:
|
| 416 |
+
"""
|
| 417 |
+
Args:
|
| 418 |
+
latent: Initial SparseTensor from encoder at 64-resolution.
|
| 419 |
+
gt_vertex_voxels_list: Ground-truth vertex SparseTensors at [64, 128, 256, 512, 1024]
|
| 420 |
+
training: Whether to apply pruning during training
|
| 421 |
+
|
| 422 |
+
Returns:
|
| 423 |
+
List[Dict] with separate vertex and edge predictions at each level
|
| 424 |
+
"""
|
| 425 |
+
latent = self.latent_expander(latent_)
|
| 426 |
+
|
| 427 |
+
results = []
|
| 428 |
+
|
| 429 |
+
# step0: shell voxels to vertex voxels
|
| 430 |
+
vtx_probs = self.vtx_pruning_head(latent) # (N, 1)
|
| 431 |
+
if not training:
|
| 432 |
+
# Inference path: use predicted vertex mask to split vertex
|
| 433 |
+
|
| 434 |
+
scores = torch.sigmoid(vtx_probs.feats).squeeze(-1) # (N,)
|
| 435 |
+
|
| 436 |
+
vertex_mask = scores >= inference_threshold # (N,)
|
| 437 |
+
batch_indices = latent.coords[:, 0]
|
| 438 |
+
for b in batch_indices.unique():
|
| 439 |
+
batch_sel = batch_indices == b
|
| 440 |
+
if vertex_mask[batch_sel].any():
|
| 441 |
+
continue
|
| 442 |
+
batch_scores = scores[batch_sel]
|
| 443 |
+
k = min(2, batch_scores.numel())
|
| 444 |
+
print(
|
| 445 |
+
f"[VertexVAE] Warning: No points passed threshold {inference_threshold} in batch {b.item()}. Forcing top {k} points."
|
| 446 |
+
)
|
| 447 |
+
|
| 448 |
+
_, top_local = torch.topk(batch_scores, k=k)
|
| 449 |
+
|
| 450 |
+
vertex_mask[batch_sel.nonzero(as_tuple=True)[0][top_local]] = True
|
| 451 |
+
|
| 452 |
+
vertex_x = sp.SparseTensor(
|
| 453 |
+
feats=latent.feats[vertex_mask],
|
| 454 |
+
coords=latent.coords[vertex_mask],
|
| 455 |
+
)
|
| 456 |
+
|
| 457 |
+
if verbose:
|
| 458 |
+
num_batches = int(latent.coords[:, 0].max().item()) + 1
|
| 459 |
+
print(
|
| 460 |
+
f"[VertexVAE] Shell2Vertex: "
|
| 461 |
+
f"num_vertex={per_batch_counts(vertex_x.coords[:, 0], num_batches)}, "
|
| 462 |
+
f"num_shell={per_batch_counts(latent.coords[:, 0], num_batches)}"
|
| 463 |
+
)
|
| 464 |
+
|
| 465 |
+
results.append(
|
| 466 |
+
{
|
| 467 |
+
"coords": vtx_probs.coords,
|
| 468 |
+
"occ_probs": vtx_probs.feats,
|
| 469 |
+
"vertex_mask": vertex_mask,
|
| 470 |
+
}
|
| 471 |
+
)
|
| 472 |
+
else:
|
| 473 |
+
# Training path: using gt voxels to split vertex
|
| 474 |
+
gt_vertex_coords = gt_vertex_voxels_list[0].coords
|
| 475 |
+
|
| 476 |
+
pred_flat = flatten_coords(latent.coords)
|
| 477 |
+
vertex_gt_flat = flatten_coords(gt_vertex_coords)
|
| 478 |
+
|
| 479 |
+
vertex_mask = torch.isin(pred_flat, vertex_gt_flat)
|
| 480 |
+
|
| 481 |
+
vertex_x = sp.SparseTensor(
|
| 482 |
+
feats=latent.feats[vertex_mask],
|
| 483 |
+
coords=latent.coords[vertex_mask],
|
| 484 |
+
)
|
| 485 |
+
|
| 486 |
+
results.append(
|
| 487 |
+
{
|
| 488 |
+
"coords": vtx_probs.coords,
|
| 489 |
+
"occ_probs": vtx_probs.feats,
|
| 490 |
+
"vertex_mask": vertex_mask,
|
| 491 |
+
"vertex_gt_coords": gt_vertex_coords,
|
| 492 |
+
}
|
| 493 |
+
)
|
| 494 |
+
|
| 495 |
+
vertex_x = self.vtx_proj(vertex_x)
|
| 496 |
+
|
| 497 |
+
# step1: upsample
|
| 498 |
+
for i, _ in enumerate(self.decoder_vtx):
|
| 499 |
+
vertex_x, vertex_occ_probs, num_rescue = self.decoder_vtx[i](
|
| 500 |
+
vertex_x,
|
| 501 |
+
training=training,
|
| 502 |
+
threshold=inference_threshold,
|
| 503 |
+
)
|
| 504 |
+
vertex_x = self.decoder_vtx_ca[i](
|
| 505 |
+
x=vertex_x,
|
| 506 |
+
context=self.latent_proj[i](latent_),
|
| 507 |
+
)
|
| 508 |
+
|
| 509 |
+
if not training:
|
| 510 |
+
# Inference path
|
| 511 |
+
if verbose:
|
| 512 |
+
num_batches = int(latent_.coords[:, 0].max().item()) + 1
|
| 513 |
+
print(
|
| 514 |
+
f"[VertexVAE] Layer{i}: "
|
| 515 |
+
f"num_vertex={per_batch_counts(vertex_x.coords[:, 0], num_batches)}, "
|
| 516 |
+
f"num_rescue={num_rescue}"
|
| 517 |
+
)
|
| 518 |
+
|
| 519 |
+
results.append(
|
| 520 |
+
{
|
| 521 |
+
"coords": vertex_x.coords,
|
| 522 |
+
"feats": vertex_x.feats,
|
| 523 |
+
"occ_probs": vertex_occ_probs.feats,
|
| 524 |
+
"occ_coords": vertex_occ_probs.coords,
|
| 525 |
+
}
|
| 526 |
+
)
|
| 527 |
+
else:
|
| 528 |
+
# Training path
|
| 529 |
+
vertex_pred_coords = vertex_x.coords
|
| 530 |
+
gt_vertex_coords = gt_vertex_voxels_list[i + 1].coords
|
| 531 |
+
|
| 532 |
+
vertex_pred_flat = flatten_coords(vertex_pred_coords)
|
| 533 |
+
vertex_gt_flat = flatten_coords(gt_vertex_coords)
|
| 534 |
+
vertex_mask = torch.isin(vertex_pred_flat, vertex_gt_flat)
|
| 535 |
+
vertex_prune_labels = vertex_mask.float()
|
| 536 |
+
|
| 537 |
+
vertex_x = sp.SparseTensor(
|
| 538 |
+
feats=vertex_x.feats[vertex_mask],
|
| 539 |
+
coords=vertex_x.coords[vertex_mask],
|
| 540 |
+
)
|
| 541 |
+
|
| 542 |
+
results.append(
|
| 543 |
+
{
|
| 544 |
+
"coords": vertex_x.coords,
|
| 545 |
+
"feats": vertex_x.feats,
|
| 546 |
+
"occ_probs": vertex_occ_probs.feats,
|
| 547 |
+
"occ_coords": vertex_occ_probs.coords,
|
| 548 |
+
"prune_labels": vertex_prune_labels,
|
| 549 |
+
"sp_tensor": vertex_x,
|
| 550 |
+
"gt_coords": gt_vertex_coords,
|
| 551 |
+
"pred_mask": vertex_mask,
|
| 552 |
+
},
|
| 553 |
+
)
|
| 554 |
+
|
| 555 |
+
return results
|
| 556 |
+
|
| 557 |
+
def forward(
|
| 558 |
+
self,
|
| 559 |
+
sparse_input,
|
| 560 |
+
gt_vertex_voxels_list=None,
|
| 561 |
+
training=True,
|
| 562 |
+
sample_posterior=True,
|
| 563 |
+
):
|
| 564 |
+
latent_64, posterior = self.encode(sparse_input, sample_posterior)
|
| 565 |
+
results = self.decode(
|
| 566 |
+
latent_64,
|
| 567 |
+
gt_vertex_voxels_list=gt_vertex_voxels_list,
|
| 568 |
+
training=training,
|
| 569 |
+
)
|
| 570 |
+
|
| 571 |
+
return results, posterior, latent_64
|
| 572 |
+
|
| 573 |
+
def convert_to_fp16(self):
|
| 574 |
+
"""Convert all components to float16"""
|
| 575 |
+
self.encoder.apply(
|
| 576 |
+
lambda m: m.convert_to_fp16() if hasattr(m, "convert_to_fp16") else None
|
| 577 |
+
)
|
| 578 |
+
self.decoder_vtx.apply(
|
| 579 |
+
lambda m: m.convert_to_fp16() if hasattr(m, "convert_to_fp16") else None
|
| 580 |
+
)
|
| 581 |
+
self.decoder_vtx_ca.apply(
|
| 582 |
+
lambda m: m.convert_to_fp16() if hasattr(m, "convert_to_fp16") else None
|
| 583 |
+
)
|
| 584 |
+
|
| 585 |
+
def convert_to_fp32(self):
|
| 586 |
+
"""Convert all components to float32"""
|
| 587 |
+
self.encoder.apply(
|
| 588 |
+
lambda m: m.convert_to_fp32() if hasattr(m, "convert_to_fp32") else None
|
| 589 |
+
)
|
| 590 |
+
self.decoder_vtx.apply(
|
| 591 |
+
lambda m: m.convert_to_fp32() if hasattr(m, "convert_to_fp32") else None
|
| 592 |
+
)
|
| 593 |
+
self.decoder_vtx_ca.apply(
|
| 594 |
+
lambda m: m.convert_to_fp32() if hasattr(m, "convert_to_fp32") else None
|
| 595 |
+
)
|
models/vertex_structured_flow.py
ADDED
|
@@ -0,0 +1,147 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from typing import *
|
| 2 |
+
import torch
|
| 3 |
+
import torch.nn as nn
|
| 4 |
+
import torch.nn.functional as F
|
| 5 |
+
|
| 6 |
+
from modules.utils import convert_module_to_f16, convert_module_to_f32
|
| 7 |
+
from modules.transformer import (
|
| 8 |
+
AbsolutePositionEmbedder,
|
| 9 |
+
TimestepEmbedder,
|
| 10 |
+
)
|
| 11 |
+
from modules import sparse as sp
|
| 12 |
+
from modules.sparse.transformer import ModulatedSparseTransformerCrossBlock
|
| 13 |
+
|
| 14 |
+
|
| 15 |
+
class VertexSLatFlowModel(nn.Module):
|
| 16 |
+
def __init__(
|
| 17 |
+
self,
|
| 18 |
+
resolution: int,
|
| 19 |
+
in_channels: int,
|
| 20 |
+
model_channels: int,
|
| 21 |
+
cond_channels: int,
|
| 22 |
+
out_channels: int,
|
| 23 |
+
num_blocks: int,
|
| 24 |
+
num_heads: Optional[int] = None,
|
| 25 |
+
num_head_channels: Optional[int] = 64,
|
| 26 |
+
mlp_ratio: float = 4,
|
| 27 |
+
pe_mode: Literal["ape", "rope"] = "ape",
|
| 28 |
+
use_fp16: bool = False,
|
| 29 |
+
use_checkpoint: bool = False,
|
| 30 |
+
share_mod: bool = False,
|
| 31 |
+
qk_rms_norm: bool = False,
|
| 32 |
+
qk_rms_norm_cross: bool = False,
|
| 33 |
+
use_density: bool = False,
|
| 34 |
+
**kwargs
|
| 35 |
+
):
|
| 36 |
+
if kwargs:
|
| 37 |
+
print(f"[SLatFlowModel] Found unused arguments: {kwargs}")
|
| 38 |
+
super().__init__()
|
| 39 |
+
self.resolution = resolution
|
| 40 |
+
self.in_channels = in_channels
|
| 41 |
+
self.model_channels = model_channels
|
| 42 |
+
self.cond_channels = cond_channels
|
| 43 |
+
self.out_channels = out_channels
|
| 44 |
+
self.num_blocks = num_blocks
|
| 45 |
+
self.num_heads = num_heads or model_channels // num_head_channels
|
| 46 |
+
self.mlp_ratio = mlp_ratio
|
| 47 |
+
self.pe_mode = pe_mode
|
| 48 |
+
self.use_fp16 = use_fp16
|
| 49 |
+
self.use_checkpoint = use_checkpoint
|
| 50 |
+
self.share_mod = share_mod
|
| 51 |
+
self.qk_rms_norm = qk_rms_norm
|
| 52 |
+
self.qk_rms_norm_cross = qk_rms_norm_cross
|
| 53 |
+
self.use_density = use_density
|
| 54 |
+
self.dtype = torch.float16 if use_fp16 else torch.float32
|
| 55 |
+
|
| 56 |
+
self.t_embedder = TimestepEmbedder(model_channels)
|
| 57 |
+
|
| 58 |
+
if self.use_density:
|
| 59 |
+
self.density_embedder = TimestepEmbedder(model_channels)
|
| 60 |
+
|
| 61 |
+
if share_mod:
|
| 62 |
+
self.adaLN_modulation = nn.Sequential(
|
| 63 |
+
nn.SiLU(), nn.Linear(model_channels, 6 * model_channels, bias=True)
|
| 64 |
+
)
|
| 65 |
+
|
| 66 |
+
if pe_mode == "ape":
|
| 67 |
+
self.pos_embedder = AbsolutePositionEmbedder(model_channels)
|
| 68 |
+
|
| 69 |
+
self.input_layer = sp.SparseLinear(
|
| 70 |
+
in_channels,
|
| 71 |
+
model_channels,
|
| 72 |
+
)
|
| 73 |
+
|
| 74 |
+
self.blocks = nn.ModuleList(
|
| 75 |
+
[
|
| 76 |
+
ModulatedSparseTransformerCrossBlock(
|
| 77 |
+
model_channels,
|
| 78 |
+
cond_channels,
|
| 79 |
+
num_heads=self.num_heads,
|
| 80 |
+
mlp_ratio=self.mlp_ratio,
|
| 81 |
+
attn_mode="full",
|
| 82 |
+
use_checkpoint=self.use_checkpoint,
|
| 83 |
+
use_rope=(pe_mode == "rope"),
|
| 84 |
+
share_mod=self.share_mod,
|
| 85 |
+
qk_rms_norm=self.qk_rms_norm,
|
| 86 |
+
qk_rms_norm_cross=self.qk_rms_norm_cross,
|
| 87 |
+
)
|
| 88 |
+
for _ in range(self.num_blocks)
|
| 89 |
+
]
|
| 90 |
+
)
|
| 91 |
+
|
| 92 |
+
self.out_layer = sp.SparseLinear(
|
| 93 |
+
model_channels,
|
| 94 |
+
out_channels,
|
| 95 |
+
)
|
| 96 |
+
|
| 97 |
+
if use_fp16:
|
| 98 |
+
self.convert_to_fp16()
|
| 99 |
+
else:
|
| 100 |
+
self.convert_to_fp32()
|
| 101 |
+
|
| 102 |
+
@property
|
| 103 |
+
def device(self) -> torch.device:
|
| 104 |
+
"""
|
| 105 |
+
Return the device of the model.
|
| 106 |
+
"""
|
| 107 |
+
return next(self.parameters()).device
|
| 108 |
+
|
| 109 |
+
def convert_to_fp16(self) -> None:
|
| 110 |
+
"""
|
| 111 |
+
Convert the torso of the model to float16.
|
| 112 |
+
"""
|
| 113 |
+
self.blocks.apply(convert_module_to_f16)
|
| 114 |
+
|
| 115 |
+
def convert_to_fp32(self) -> None:
|
| 116 |
+
"""
|
| 117 |
+
Convert the torso of the model to float32.
|
| 118 |
+
"""
|
| 119 |
+
self.blocks.apply(convert_module_to_f32)
|
| 120 |
+
|
| 121 |
+
def forward(
|
| 122 |
+
self,
|
| 123 |
+
x: sp.SparseTensor,
|
| 124 |
+
t: torch.Tensor,
|
| 125 |
+
cond: torch.Tensor,
|
| 126 |
+
density: Optional[torch.Tensor] = None,
|
| 127 |
+
) -> sp.SparseTensor:
|
| 128 |
+
h = self.input_layer(x).type(self.dtype)
|
| 129 |
+
t_emb = self.t_embedder(t)
|
| 130 |
+
if self.use_density:
|
| 131 |
+
assert (
|
| 132 |
+
density is not None
|
| 133 |
+
), "Density tensor must be provided when use_density is True"
|
| 134 |
+
t_emb = t_emb + self.density_embedder(density.reshape(-1).float())
|
| 135 |
+
if self.share_mod:
|
| 136 |
+
t_emb = self.adaLN_modulation(t_emb)
|
| 137 |
+
t_emb = t_emb.type(self.dtype)
|
| 138 |
+
cond = cond.type(self.dtype)
|
| 139 |
+
|
| 140 |
+
if self.pe_mode == "ape":
|
| 141 |
+
h = h + self.pos_embedder(h.coords[:, 1:]).type(self.dtype)
|
| 142 |
+
for block in self.blocks:
|
| 143 |
+
h = block(h, t_emb, cond)
|
| 144 |
+
|
| 145 |
+
h = h.replace(F.layer_norm(h.feats, h.feats.shape[-1:]))
|
| 146 |
+
h = self.out_layer(h.type(x.dtype))
|
| 147 |
+
return h
|
models/voxel_encoder.py
ADDED
|
@@ -0,0 +1,183 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from __future__ import annotations
|
| 2 |
+
|
| 3 |
+
import torch
|
| 4 |
+
import torch.nn as nn
|
| 5 |
+
|
| 6 |
+
|
| 7 |
+
def _safe_group_norm(num_channels: int, max_groups: int = 8) -> nn.GroupNorm:
|
| 8 |
+
g = min(max_groups, num_channels)
|
| 9 |
+
while g > 1 and num_channels % g != 0:
|
| 10 |
+
g -= 1
|
| 11 |
+
return nn.GroupNorm(g, num_channels)
|
| 12 |
+
|
| 13 |
+
|
| 14 |
+
def _sincos_1d(n: int, dim: int) -> torch.Tensor:
|
| 15 |
+
assert dim % 2 == 0 and dim > 0, f"sincos dim must be positive even, got {dim}"
|
| 16 |
+
pos = torch.arange(n, dtype=torch.float32)
|
| 17 |
+
omega = torch.arange(dim // 2, dtype=torch.float32) / (dim // 2)
|
| 18 |
+
omega = 1.0 / (10000**omega)
|
| 19 |
+
out = pos[:, None] * omega[None, :]
|
| 20 |
+
return torch.cat([torch.sin(out), torch.cos(out)], dim=-1)
|
| 21 |
+
|
| 22 |
+
|
| 23 |
+
def _get_3d_sincos_embed(n: int, dim: int) -> torch.Tensor:
|
| 24 |
+
axis_dim = (dim // 3) // 2 * 2 # split across 3 axes, round to even
|
| 25 |
+
if axis_dim <= 0:
|
| 26 |
+
raise ValueError(
|
| 27 |
+
f"cond_in_dim={dim} too small for 3D sincos PE (need >= 6 so each axis gets a positive even slice)"
|
| 28 |
+
)
|
| 29 |
+
e = _sincos_1d(n, axis_dim) # (n, axis_dim)ß
|
| 30 |
+
pe_d = e[:, None, None, :].expand(n, n, n, axis_dim)
|
| 31 |
+
pe_h = e[None, :, None, :].expand(n, n, n, axis_dim)
|
| 32 |
+
pe_w = e[None, None, :, :].expand(n, n, n, axis_dim)
|
| 33 |
+
pe = torch.cat([pe_d, pe_h, pe_w], dim=-1).reshape(n * n * n, 3 * axis_dim)
|
| 34 |
+
if pe.shape[-1] < dim:
|
| 35 |
+
pad = torch.zeros(pe.shape[0], dim - pe.shape[-1])
|
| 36 |
+
pe = torch.cat([pe, pad], dim=-1)
|
| 37 |
+
return pe
|
| 38 |
+
|
| 39 |
+
|
| 40 |
+
class _ResBlock3d(nn.Module):
|
| 41 |
+
def __init__(self, channels: int) -> None:
|
| 42 |
+
super().__init__()
|
| 43 |
+
self.block = nn.Sequential(
|
| 44 |
+
_safe_group_norm(channels),
|
| 45 |
+
nn.SiLU(),
|
| 46 |
+
nn.Conv3d(channels, channels, kernel_size=3, padding=1, bias=False),
|
| 47 |
+
_safe_group_norm(channels),
|
| 48 |
+
nn.SiLU(),
|
| 49 |
+
nn.Conv3d(channels, channels, kernel_size=3, padding=1, bias=False),
|
| 50 |
+
)
|
| 51 |
+
|
| 52 |
+
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
| 53 |
+
return x + self.block(x)
|
| 54 |
+
|
| 55 |
+
|
| 56 |
+
class _DownBlock3d(nn.Module):
|
| 57 |
+
def __init__(self, in_ch: int, out_ch: int, blocks_per_level: int) -> None:
|
| 58 |
+
super().__init__()
|
| 59 |
+
res_blocks: list[nn.Module] = [
|
| 60 |
+
_ResBlock3d(in_ch) for _ in range(blocks_per_level)
|
| 61 |
+
]
|
| 62 |
+
res_blocks.append(
|
| 63 |
+
nn.Conv3d(in_ch, out_ch, kernel_size=3, stride=2, padding=1, bias=False)
|
| 64 |
+
)
|
| 65 |
+
res_blocks.append(_safe_group_norm(out_ch))
|
| 66 |
+
res_blocks.append(nn.SiLU())
|
| 67 |
+
self.net = nn.Sequential(*res_blocks)
|
| 68 |
+
|
| 69 |
+
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
| 70 |
+
return self.net(x)
|
| 71 |
+
|
| 72 |
+
|
| 73 |
+
class VoxelFieldConditioner(nn.Module):
|
| 74 |
+
def __init__(
|
| 75 |
+
self,
|
| 76 |
+
in_channels: int,
|
| 77 |
+
cond_in_dim: int,
|
| 78 |
+
*,
|
| 79 |
+
num_downsamples: int = 2,
|
| 80 |
+
base_channels: int = 32,
|
| 81 |
+
channel_mult: int = 2,
|
| 82 |
+
blocks_per_level: int = 2,
|
| 83 |
+
pos_embed: str = "sincos",
|
| 84 |
+
) -> None:
|
| 85 |
+
super().__init__()
|
| 86 |
+
if num_downsamples < 0:
|
| 87 |
+
raise ValueError(f"num_downsamples must be >= 0, got {num_downsamples}")
|
| 88 |
+
if in_channels <= 0:
|
| 89 |
+
raise ValueError(f"in_channels must be > 0, got {in_channels}")
|
| 90 |
+
if cond_in_dim <= 0:
|
| 91 |
+
raise ValueError(f"cond_in_dim must be > 0, got {cond_in_dim}")
|
| 92 |
+
if base_channels <= 0:
|
| 93 |
+
raise ValueError(f"base_channels must be > 0, got {base_channels}")
|
| 94 |
+
if channel_mult < 1:
|
| 95 |
+
raise ValueError(f"channel_mult must be >= 1, got {channel_mult}")
|
| 96 |
+
if blocks_per_level < 0:
|
| 97 |
+
raise ValueError(f"blocks_per_level must be >= 0, got {blocks_per_level}")
|
| 98 |
+
pos_embed = str(pos_embed).lower()
|
| 99 |
+
if pos_embed not in ("sincos", "none"):
|
| 100 |
+
raise ValueError(
|
| 101 |
+
f"pos_embed={pos_embed!r} unsupported (use 'sincos' or 'none')"
|
| 102 |
+
)
|
| 103 |
+
|
| 104 |
+
self.in_channels = int(in_channels)
|
| 105 |
+
self.cond_in_dim = int(cond_in_dim)
|
| 106 |
+
self.num_downsamples = int(num_downsamples)
|
| 107 |
+
self.base_channels = int(base_channels)
|
| 108 |
+
self.channel_mult = int(channel_mult)
|
| 109 |
+
self.blocks_per_level = int(blocks_per_level)
|
| 110 |
+
self.pos_embed = pos_embed
|
| 111 |
+
|
| 112 |
+
self.stem = nn.Sequential(
|
| 113 |
+
nn.Conv3d(
|
| 114 |
+
self.in_channels,
|
| 115 |
+
self.base_channels,
|
| 116 |
+
kernel_size=3,
|
| 117 |
+
padding=1,
|
| 118 |
+
bias=False,
|
| 119 |
+
),
|
| 120 |
+
_safe_group_norm(self.base_channels),
|
| 121 |
+
nn.SiLU(),
|
| 122 |
+
)
|
| 123 |
+
|
| 124 |
+
down_blocks: list[nn.Module] = []
|
| 125 |
+
ch = self.base_channels
|
| 126 |
+
for _ in range(self.num_downsamples):
|
| 127 |
+
out_ch = ch * self.channel_mult
|
| 128 |
+
down_blocks.append(_DownBlock3d(ch, out_ch, self.blocks_per_level))
|
| 129 |
+
ch = out_ch
|
| 130 |
+
self.down_blocks = nn.ModuleList(down_blocks)
|
| 131 |
+
self._final_channels = ch # base_channels * channel_mult ** num_downsamples
|
| 132 |
+
|
| 133 |
+
self.tail_blocks = nn.Sequential(
|
| 134 |
+
*[_ResBlock3d(ch) for _ in range(blocks_per_level)]
|
| 135 |
+
)
|
| 136 |
+
|
| 137 |
+
self.proj = nn.Conv3d(
|
| 138 |
+
self._final_channels, self.cond_in_dim, kernel_size=1, bias=True
|
| 139 |
+
)
|
| 140 |
+
|
| 141 |
+
def _get_pe(
|
| 142 |
+
self, n_out: int, device: torch.device, dtype: torch.dtype
|
| 143 |
+
) -> torch.Tensor:
|
| 144 |
+
buf_name = f"_sincos_pe_{n_out}"
|
| 145 |
+
if not hasattr(self, buf_name):
|
| 146 |
+
pe = _get_3d_sincos_embed(n_out, self.cond_in_dim)
|
| 147 |
+
self.register_buffer(buf_name, pe, persistent=False)
|
| 148 |
+
return getattr(self, buf_name).to(device=device, dtype=dtype)
|
| 149 |
+
|
| 150 |
+
def forward(self, field: torch.Tensor) -> torch.Tensor:
|
| 151 |
+
"""
|
| 152 |
+
Input:
|
| 153 |
+
field: (B, R, R, R) or (B, C_in, R, R, R)
|
| 154 |
+
Output:
|
| 155 |
+
(B, R'^3, cond_in_dim) token sequence with 3D PE added.
|
| 156 |
+
"""
|
| 157 |
+
if field.dim() == 4:
|
| 158 |
+
field = field.unsqueeze(1)
|
| 159 |
+
elif field.dim() != 5:
|
| 160 |
+
raise ValueError(
|
| 161 |
+
f"field must be 4D (B,R,R,R) or 5D (B,C,R,R,R), got {tuple(field.shape)}"
|
| 162 |
+
)
|
| 163 |
+
if field.shape[1] != self.in_channels:
|
| 164 |
+
raise ValueError(
|
| 165 |
+
f"field channel dim {field.shape[1]} != in_channels {self.in_channels}"
|
| 166 |
+
)
|
| 167 |
+
if not (field.shape[2] == field.shape[3] == field.shape[4]):
|
| 168 |
+
raise ValueError(
|
| 169 |
+
f"field must be cubic (R,R,R), got spatial {tuple(field.shape[2:])}"
|
| 170 |
+
)
|
| 171 |
+
|
| 172 |
+
x = self.stem(field)
|
| 173 |
+
for blk in self.down_blocks:
|
| 174 |
+
x = blk(x)
|
| 175 |
+
x = self.tail_blocks(x)
|
| 176 |
+
feat = self.proj(x)
|
| 177 |
+
|
| 178 |
+
n_out = feat.shape[-1]
|
| 179 |
+
tokens = feat.flatten(2).transpose(1, 2).contiguous()
|
| 180 |
+
if self.pos_embed == "sincos":
|
| 181 |
+
pe = self._get_pe(n_out, tokens.device, tokens.dtype)
|
| 182 |
+
tokens = tokens + pe.unsqueeze(0)
|
| 183 |
+
return tokens
|
modules/attention.py
ADDED
|
@@ -0,0 +1,160 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from __future__ import annotations
|
| 2 |
+
|
| 3 |
+
from typing import Optional
|
| 4 |
+
import torch
|
| 5 |
+
import torch.nn.functional as F
|
| 6 |
+
|
| 7 |
+
try:
|
| 8 |
+
from flash_attn import flash_attn_varlen_func
|
| 9 |
+
|
| 10 |
+
_FLASH_ATTN_AVAILABLE = True
|
| 11 |
+
except Exception:
|
| 12 |
+
flash_attn_varlen_func = None
|
| 13 |
+
_FLASH_ATTN_AVAILABLE = False
|
| 14 |
+
|
| 15 |
+
|
| 16 |
+
def flash_varlen_self_attention(
|
| 17 |
+
q: torch.Tensor,
|
| 18 |
+
k: torch.Tensor,
|
| 19 |
+
v: torch.Tensor,
|
| 20 |
+
x_mask: torch.Tensor,
|
| 21 |
+
) -> torch.Tensor:
|
| 22 |
+
"""q,k,v: (B, H, N, Dh); x_mask: (B, N) bool."""
|
| 23 |
+
bsz, nheads, seqlen, head_dim = q.shape
|
| 24 |
+
mask = x_mask.bool()
|
| 25 |
+
lengths = mask.sum(dim=-1, dtype=torch.int32)
|
| 26 |
+
if int(lengths.max().item()) <= 0:
|
| 27 |
+
return torch.zeros_like(q)
|
| 28 |
+
|
| 29 |
+
cu_seqlens = torch.zeros((bsz + 1,), dtype=torch.int32, device=q.device)
|
| 30 |
+
cu_seqlens[1:] = torch.cumsum(lengths, dim=0)
|
| 31 |
+
max_seqlen = int(lengths.max().item())
|
| 32 |
+
|
| 33 |
+
q_flat = q.permute(0, 2, 1, 3).reshape(bsz * seqlen, nheads, head_dim)
|
| 34 |
+
k_flat = k.permute(0, 2, 1, 3).reshape(bsz * seqlen, nheads, head_dim)
|
| 35 |
+
v_flat = v.permute(0, 2, 1, 3).reshape(bsz * seqlen, nheads, head_dim)
|
| 36 |
+
valid_token_indices = torch.nonzero(mask.reshape(-1), as_tuple=False).squeeze(-1)
|
| 37 |
+
|
| 38 |
+
q_unpad = q_flat.index_select(0, valid_token_indices)
|
| 39 |
+
k_unpad = k_flat.index_select(0, valid_token_indices)
|
| 40 |
+
v_unpad = v_flat.index_select(0, valid_token_indices)
|
| 41 |
+
|
| 42 |
+
attn_unpad = flash_attn_varlen_func(
|
| 43 |
+
q_unpad,
|
| 44 |
+
k_unpad,
|
| 45 |
+
v_unpad,
|
| 46 |
+
cu_seqlens_q=cu_seqlens,
|
| 47 |
+
cu_seqlens_k=cu_seqlens,
|
| 48 |
+
max_seqlen_q=max_seqlen,
|
| 49 |
+
max_seqlen_k=max_seqlen,
|
| 50 |
+
dropout_p=0.0,
|
| 51 |
+
causal=False,
|
| 52 |
+
)
|
| 53 |
+
out_flat = torch.zeros_like(q_flat)
|
| 54 |
+
out_flat.index_copy_(0, valid_token_indices, attn_unpad)
|
| 55 |
+
out = out_flat.reshape(bsz, seqlen, nheads, head_dim).permute(0, 2, 1, 3)
|
| 56 |
+
return out
|
| 57 |
+
|
| 58 |
+
|
| 59 |
+
def flash_varlen_cross_attention(
|
| 60 |
+
q: torch.Tensor,
|
| 61 |
+
k: torch.Tensor,
|
| 62 |
+
v: torch.Tensor,
|
| 63 |
+
q_mask: torch.Tensor,
|
| 64 |
+
k_mask: torch.Tensor,
|
| 65 |
+
) -> torch.Tensor:
|
| 66 |
+
"""Varlen cross-attn. q: (B,H,Nq,Dh), k/v: (B,H,Nk,Dh), masks (B,Nq)/(B,Nk) bool."""
|
| 67 |
+
bsz, nheads, nq, head_dim = q.shape
|
| 68 |
+
nk = k.shape[2]
|
| 69 |
+
q_mask_b = q_mask.bool()
|
| 70 |
+
k_mask_b = k_mask.bool()
|
| 71 |
+
q_lengths = q_mask_b.sum(dim=-1, dtype=torch.int32)
|
| 72 |
+
k_lengths = k_mask_b.sum(dim=-1, dtype=torch.int32)
|
| 73 |
+
if int(q_lengths.max().item()) <= 0:
|
| 74 |
+
return torch.zeros_like(q)
|
| 75 |
+
|
| 76 |
+
cu_q = torch.zeros((bsz + 1,), dtype=torch.int32, device=q.device)
|
| 77 |
+
cu_q[1:] = torch.cumsum(q_lengths, dim=0)
|
| 78 |
+
cu_k = torch.zeros((bsz + 1,), dtype=torch.int32, device=q.device)
|
| 79 |
+
cu_k[1:] = torch.cumsum(k_lengths, dim=0)
|
| 80 |
+
max_q = int(q_lengths.max().item())
|
| 81 |
+
max_k = int(k_lengths.max().item())
|
| 82 |
+
|
| 83 |
+
q_flat = q.permute(0, 2, 1, 3).reshape(bsz * nq, nheads, head_dim)
|
| 84 |
+
k_flat = k.permute(0, 2, 1, 3).reshape(bsz * nk, nheads, head_dim)
|
| 85 |
+
v_flat = v.permute(0, 2, 1, 3).reshape(bsz * nk, nheads, head_dim)
|
| 86 |
+
|
| 87 |
+
q_idx = torch.nonzero(q_mask_b.reshape(-1), as_tuple=False).squeeze(-1)
|
| 88 |
+
k_idx = torch.nonzero(k_mask_b.reshape(-1), as_tuple=False).squeeze(-1)
|
| 89 |
+
|
| 90 |
+
q_unpad = q_flat.index_select(0, q_idx)
|
| 91 |
+
k_unpad = k_flat.index_select(0, k_idx)
|
| 92 |
+
v_unpad = v_flat.index_select(0, k_idx)
|
| 93 |
+
|
| 94 |
+
attn_unpad = flash_attn_varlen_func(
|
| 95 |
+
q_unpad,
|
| 96 |
+
k_unpad,
|
| 97 |
+
v_unpad,
|
| 98 |
+
cu_seqlens_q=cu_q,
|
| 99 |
+
cu_seqlens_k=cu_k,
|
| 100 |
+
max_seqlen_q=max_q,
|
| 101 |
+
max_seqlen_k=max_k,
|
| 102 |
+
dropout_p=0.0,
|
| 103 |
+
causal=False,
|
| 104 |
+
)
|
| 105 |
+
out_flat = torch.zeros_like(q_flat)
|
| 106 |
+
out_flat.index_copy_(0, q_idx, attn_unpad)
|
| 107 |
+
out = out_flat.reshape(bsz, nq, nheads, head_dim).permute(0, 2, 1, 3)
|
| 108 |
+
return out
|
| 109 |
+
|
| 110 |
+
|
| 111 |
+
def can_flash_varlen(q: torch.Tensor, x_mask: Optional[torch.Tensor]) -> bool:
|
| 112 |
+
if not _FLASH_ATTN_AVAILABLE or x_mask is None:
|
| 113 |
+
return False
|
| 114 |
+
if not q.is_cuda:
|
| 115 |
+
return False
|
| 116 |
+
if q.dtype not in (torch.float16, torch.bfloat16):
|
| 117 |
+
return False
|
| 118 |
+
return True
|
| 119 |
+
|
| 120 |
+
|
| 121 |
+
def sdpa_padding_mask(x_mask: torch.Tensor) -> torch.Tensor:
|
| 122 |
+
"""(B, 1, 1, N) bool: keys valid."""
|
| 123 |
+
return x_mask.bool().view(x_mask.shape[0], 1, 1, x_mask.shape[1])
|
| 124 |
+
|
| 125 |
+
|
| 126 |
+
def graph_adj_varlen_attention(
|
| 127 |
+
q: torch.Tensor,
|
| 128 |
+
k: torch.Tensor,
|
| 129 |
+
v: torch.Tensor,
|
| 130 |
+
x_mask: torch.Tensor,
|
| 131 |
+
adj_matrix: Optional[torch.Tensor],
|
| 132 |
+
) -> torch.Tensor:
|
| 133 |
+
bsz, nheads, seqlen, _ = q.shape
|
| 134 |
+
out = torch.zeros_like(q)
|
| 135 |
+
x_mask = x_mask.bool()
|
| 136 |
+
|
| 137 |
+
for b in range(bsz):
|
| 138 |
+
valid_indices = torch.nonzero(x_mask[b], as_tuple=False).squeeze(-1)
|
| 139 |
+
if valid_indices.numel() == 0:
|
| 140 |
+
continue
|
| 141 |
+
q_b = q[b].index_select(1, valid_indices).unsqueeze(0)
|
| 142 |
+
k_b = k[b].index_select(1, valid_indices).unsqueeze(0)
|
| 143 |
+
v_b = v[b].index_select(1, valid_indices).unsqueeze(0)
|
| 144 |
+
l_now = valid_indices.numel()
|
| 145 |
+
|
| 146 |
+
if adj_matrix is not None:
|
| 147 |
+
sub = (
|
| 148 |
+
adj_matrix[b]
|
| 149 |
+
.bool()
|
| 150 |
+
.index_select(0, valid_indices)
|
| 151 |
+
.index_select(1, valid_indices)
|
| 152 |
+
)
|
| 153 |
+
eye = torch.eye(l_now, dtype=torch.bool, device=q.device)
|
| 154 |
+
attn_mask_b = (sub | eye).view(1, 1, l_now, l_now)
|
| 155 |
+
else:
|
| 156 |
+
attn_mask_b = None
|
| 157 |
+
|
| 158 |
+
out_b = F.scaled_dot_product_attention(q_b, k_b, v_b, attn_mask=attn_mask_b)
|
| 159 |
+
out[b].index_copy_(1, valid_indices, out_b.squeeze(0))
|
| 160 |
+
return out
|
modules/norm.py
ADDED
|
@@ -0,0 +1,41 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import torch
|
| 2 |
+
import torch.nn as nn
|
| 3 |
+
|
| 4 |
+
|
| 5 |
+
class LayerNorm32(nn.LayerNorm):
|
| 6 |
+
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
| 7 |
+
return super().forward(x.float()).type(x.dtype)
|
| 8 |
+
|
| 9 |
+
|
| 10 |
+
class RMSNorm(nn.Module):
|
| 11 |
+
def __init__(
|
| 12 |
+
self,
|
| 13 |
+
dim: int,
|
| 14 |
+
eps: float = 1e-5,
|
| 15 |
+
elementwise_affine: bool = True,
|
| 16 |
+
variance_in_fp32: bool = True,
|
| 17 |
+
):
|
| 18 |
+
super().__init__()
|
| 19 |
+
self.eps = eps
|
| 20 |
+
self.elementwise_affine = elementwise_affine
|
| 21 |
+
self.variance_in_fp32 = bool(variance_in_fp32)
|
| 22 |
+
self.weight = nn.Parameter(torch.ones(dim)) if elementwise_affine else None
|
| 23 |
+
|
| 24 |
+
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
| 25 |
+
input_dtype = x.dtype
|
| 26 |
+
if self.variance_in_fp32:
|
| 27 |
+
variance = x.to(torch.float32).pow(2).mean(-1, keepdim=True)
|
| 28 |
+
inv_rms = torch.rsqrt(variance + self.eps).to(input_dtype)
|
| 29 |
+
else:
|
| 30 |
+
variance = x.pow(2).mean(-1, keepdim=True)
|
| 31 |
+
inv_rms = torch.rsqrt(variance + self.eps)
|
| 32 |
+
x = x * inv_rms
|
| 33 |
+
if self.weight is not None:
|
| 34 |
+
w = self.weight
|
| 35 |
+
if w.dtype in (torch.float16, torch.bfloat16):
|
| 36 |
+
x = (x.to(w.dtype) * w).to(input_dtype)
|
| 37 |
+
else:
|
| 38 |
+
x = (x * w).to(input_dtype)
|
| 39 |
+
else:
|
| 40 |
+
x = x.to(input_dtype)
|
| 41 |
+
return x
|
modules/pointnet.py
ADDED
|
@@ -0,0 +1,330 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# MIT License
|
| 2 |
+
|
| 3 |
+
# Copyright (c) 2020 Songyou Peng, Michael Niemeyer, Lars Mescheder, Marc Pollefeys, Andreas Geiger.
|
| 4 |
+
# Copyright (c) 2025 VAST-AI-Research and contributors.
|
| 5 |
+
|
| 6 |
+
# Permission is hereby granted, free of charge, to any person obtaining a copy
|
| 7 |
+
# of this software and associated documentation files (the "Software"), to deal
|
| 8 |
+
# in the Software without restriction, including without limitation the rights
|
| 9 |
+
# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
|
| 10 |
+
# copies of the Software, and to permit persons to whom the Software is
|
| 11 |
+
# furnished to do so, subject to the following conditions:
|
| 12 |
+
|
| 13 |
+
# The above copyright notice and this permission notice shall be included in all
|
| 14 |
+
# copies or substantial portions of the Software.
|
| 15 |
+
|
| 16 |
+
# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
|
| 17 |
+
# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
|
| 18 |
+
# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
|
| 19 |
+
# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
|
| 20 |
+
# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
|
| 21 |
+
# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
|
| 22 |
+
# SOFTWARE
|
| 23 |
+
|
| 24 |
+
# modified from https://github.com/autonomousvision/convolutional_occupancy_networks/blob/master/src/encoder/pointnet.py
|
| 25 |
+
# modified from https://github.com/VAST-AI-Research/TripoSF/blob/main/triposf/modules/pointclouds/pointnet.py
|
| 26 |
+
|
| 27 |
+
import torch
|
| 28 |
+
import torch.nn as nn
|
| 29 |
+
import copy
|
| 30 |
+
from torch import Tensor
|
| 31 |
+
from torch_scatter import scatter_mean
|
| 32 |
+
from torch.utils.checkpoint import checkpoint
|
| 33 |
+
|
| 34 |
+
|
| 35 |
+
def scale_tensor(dat, inp_scale=None, tgt_scale=None):
|
| 36 |
+
if inp_scale is None:
|
| 37 |
+
inp_scale = (-0.5, 0.5)
|
| 38 |
+
if tgt_scale is None:
|
| 39 |
+
tgt_scale = (0, 1)
|
| 40 |
+
assert tgt_scale[1] > tgt_scale[0] and inp_scale[1] > inp_scale[0]
|
| 41 |
+
if isinstance(tgt_scale, Tensor):
|
| 42 |
+
assert dat.shape[-1] == tgt_scale.shape[-1]
|
| 43 |
+
dat = (dat - inp_scale[0]) / (inp_scale[1] - inp_scale[0])
|
| 44 |
+
dat = dat * (tgt_scale[1] - tgt_scale[0]) + tgt_scale[0]
|
| 45 |
+
return dat.clamp(tgt_scale[0] + 1e-6, tgt_scale[1] - 1e-6)
|
| 46 |
+
|
| 47 |
+
|
| 48 |
+
# Resnet Blocks for pointnet
|
| 49 |
+
class ResnetBlockFC(nn.Module):
|
| 50 |
+
"""Fully connected ResNet Block class.
|
| 51 |
+
|
| 52 |
+
Args:
|
| 53 |
+
size_in (int): input dimension
|
| 54 |
+
size_out (int): output dimension
|
| 55 |
+
size_h (int): hidden dimension
|
| 56 |
+
"""
|
| 57 |
+
|
| 58 |
+
def __init__(self, size_in, size_out=None, size_h=None):
|
| 59 |
+
super().__init__()
|
| 60 |
+
# Attributes
|
| 61 |
+
if size_out is None:
|
| 62 |
+
size_out = size_in
|
| 63 |
+
|
| 64 |
+
if size_h is None:
|
| 65 |
+
size_h = min(size_in, size_out)
|
| 66 |
+
|
| 67 |
+
self.size_in = size_in
|
| 68 |
+
self.size_h = size_h
|
| 69 |
+
self.size_out = size_out
|
| 70 |
+
# Submodules
|
| 71 |
+
self.fc_0 = nn.Linear(size_in, size_h)
|
| 72 |
+
self.fc_1 = nn.Linear(size_h, size_out)
|
| 73 |
+
self.actvn = nn.GELU(approximate="tanh")
|
| 74 |
+
|
| 75 |
+
if size_in == size_out:
|
| 76 |
+
self.shortcut = None
|
| 77 |
+
else:
|
| 78 |
+
self.shortcut = nn.Linear(size_in, size_out, bias=False)
|
| 79 |
+
# Initialization
|
| 80 |
+
nn.init.xavier_uniform_(self.fc_0.weight)
|
| 81 |
+
if self.fc_0.bias is not None:
|
| 82 |
+
nn.init.constant_(self.fc_0.bias, 0)
|
| 83 |
+
if self.shortcut is not None:
|
| 84 |
+
nn.init.xavier_uniform_(self.shortcut.weight)
|
| 85 |
+
if self.shortcut.bias is not None:
|
| 86 |
+
nn.init.constant_(self.shortcut.bias, 0)
|
| 87 |
+
|
| 88 |
+
nn.init.xavier_uniform_(self.fc_1.weight)
|
| 89 |
+
if self.fc_1.bias is not None:
|
| 90 |
+
nn.init.constant_(self.fc_1.bias, 0)
|
| 91 |
+
|
| 92 |
+
def forward(self, x):
|
| 93 |
+
net = self.fc_0(self.actvn(x))
|
| 94 |
+
dx = self.fc_1(self.actvn(net))
|
| 95 |
+
|
| 96 |
+
if self.shortcut is not None:
|
| 97 |
+
x_s = self.shortcut(x)
|
| 98 |
+
else:
|
| 99 |
+
x_s = x
|
| 100 |
+
|
| 101 |
+
return x_s + dx
|
| 102 |
+
|
| 103 |
+
|
| 104 |
+
class LocalPoolPointnet(nn.Module):
|
| 105 |
+
def __init__(
|
| 106 |
+
self,
|
| 107 |
+
in_channels=3,
|
| 108 |
+
out_channels=128,
|
| 109 |
+
hidden_dim=128,
|
| 110 |
+
scatter_type="mean",
|
| 111 |
+
n_blocks=5,
|
| 112 |
+
):
|
| 113 |
+
super().__init__()
|
| 114 |
+
self.scatter_type = scatter_type
|
| 115 |
+
self.in_channels = in_channels
|
| 116 |
+
self.hidden_dim = hidden_dim
|
| 117 |
+
self.out_channels = out_channels
|
| 118 |
+
self.fc_pos = nn.Linear(in_channels, 2 * hidden_dim)
|
| 119 |
+
self.blocks = nn.ModuleList(
|
| 120 |
+
[ResnetBlockFC(2 * hidden_dim, hidden_dim) for i in range(n_blocks)]
|
| 121 |
+
)
|
| 122 |
+
self.fc_c = nn.Linear(hidden_dim, out_channels)
|
| 123 |
+
self.in_channels = in_channels
|
| 124 |
+
if self.scatter_type == "mean":
|
| 125 |
+
self.scatter = scatter_mean
|
| 126 |
+
else:
|
| 127 |
+
raise ValueError("Incorrect scatter type")
|
| 128 |
+
self.initialize_weights()
|
| 129 |
+
|
| 130 |
+
def initialize_weights(self):
|
| 131 |
+
|
| 132 |
+
nn.init.xavier_uniform_(self.fc_pos.weight)
|
| 133 |
+
if self.fc_pos.bias is not None:
|
| 134 |
+
nn.init.constant_(self.fc_pos.bias, 0)
|
| 135 |
+
|
| 136 |
+
nn.init.xavier_uniform_(self.fc_c.weight)
|
| 137 |
+
if self.fc_c.bias is not None:
|
| 138 |
+
nn.init.constant_(self.fc_c.bias, 0)
|
| 139 |
+
|
| 140 |
+
def convert_to_sparse_feats(self, c, sparse_coords):
|
| 141 |
+
"""
|
| 142 |
+
Input:
|
| 143 |
+
sparse_coords: Tensor [Nx, 4], point to sparse indices
|
| 144 |
+
c: Tensor [B, res, C], input feats of each grid
|
| 145 |
+
Output:
|
| 146 |
+
c_out: Tensor [B, Np, C], aggregated grid feats of each point
|
| 147 |
+
"""
|
| 148 |
+
feats_new = torch.zeros(
|
| 149 |
+
(sparse_coords.shape[0], c.shape[-1]), device=c.device, dtype=c.dtype
|
| 150 |
+
)
|
| 151 |
+
offsets = 0
|
| 152 |
+
|
| 153 |
+
batch_nums = copy.deepcopy(sparse_coords[..., 0])
|
| 154 |
+
for i in range(len(c)):
|
| 155 |
+
coords_num_i = (batch_nums == i).sum()
|
| 156 |
+
feats_new[offsets : offsets + coords_num_i] = c[i, :coords_num_i]
|
| 157 |
+
offsets += coords_num_i
|
| 158 |
+
return feats_new
|
| 159 |
+
|
| 160 |
+
def generate_sparse_grid_features(self, index, c, max_coord_num):
|
| 161 |
+
# scatter grid features from points
|
| 162 |
+
bs, fea_dim = c.size(0), c.size(2)
|
| 163 |
+
res = max_coord_num
|
| 164 |
+
c_out = c.new_zeros(bs, self.out_channels, res)
|
| 165 |
+
c_out = scatter_mean(c.permute(0, 2, 1), index, out=c_out).permute(
|
| 166 |
+
0, 2, 1
|
| 167 |
+
) # B x res X C
|
| 168 |
+
return c_out
|
| 169 |
+
|
| 170 |
+
def pool_sparse_local(self, index, c, max_coord_num):
|
| 171 |
+
"""
|
| 172 |
+
Input:
|
| 173 |
+
index: Tensor [B, 1, Np], sparse indices of each point
|
| 174 |
+
c: Tensor [B, Np, C], input feats of each point
|
| 175 |
+
Output:
|
| 176 |
+
c_out: Tensor [B, Np, C], aggregated grid feats of each point
|
| 177 |
+
"""
|
| 178 |
+
|
| 179 |
+
bs, fea_dim = c.size(0), c.size(2)
|
| 180 |
+
res = max_coord_num
|
| 181 |
+
c_out = c.new_zeros(bs, fea_dim, res)
|
| 182 |
+
c_out = self.scatter(c.permute(0, 2, 1), index, out=c_out)
|
| 183 |
+
|
| 184 |
+
# gather feature back to points
|
| 185 |
+
c_out = c_out.gather(dim=2, index=index.expand(-1, fea_dim, -1))
|
| 186 |
+
return c_out.permute(0, 2, 1)
|
| 187 |
+
|
| 188 |
+
@torch.no_grad()
|
| 189 |
+
def coordinate2sparseindex(self, x, sparse_coords, res):
|
| 190 |
+
"""
|
| 191 |
+
Input:
|
| 192 |
+
x: Tensor [B, Np, 3], points scaled at ([0, 1] * res)
|
| 193 |
+
sparse_coords: Tensor [Nx, 4] ([batch_number, x, y, z])
|
| 194 |
+
res: Int, resolution of the grid index
|
| 195 |
+
Output:
|
| 196 |
+
sparse_index: Tensor [B, 1, Np], sparse indices of each point
|
| 197 |
+
"""
|
| 198 |
+
B = x.shape[0]
|
| 199 |
+
sparse_index = torch.zeros((B, x.shape[1]), device=x.device, dtype=torch.int64)
|
| 200 |
+
|
| 201 |
+
index = (x[..., 0] * res + x[..., 1]) * res + x[..., 2]
|
| 202 |
+
sparse_indices = copy.deepcopy(sparse_coords)
|
| 203 |
+
sparse_indices[..., 1] = (
|
| 204 |
+
sparse_indices[..., 1] * res + sparse_indices[..., 2]
|
| 205 |
+
) * res + sparse_indices[..., 3]
|
| 206 |
+
sparse_indices = sparse_indices[..., :2]
|
| 207 |
+
|
| 208 |
+
for i in range(B):
|
| 209 |
+
mask_i = sparse_indices[..., 0] == i
|
| 210 |
+
coords_i = sparse_indices[mask_i, 1]
|
| 211 |
+
coords_num_i = len(coords_i)
|
| 212 |
+
sparse_index[i] = torch.searchsorted(coords_i, index[i])
|
| 213 |
+
|
| 214 |
+
return sparse_index[:, None, :]
|
| 215 |
+
|
| 216 |
+
def forward(self, p, sparse_coords, res=64, bbox_size=(-0.5, 0.5)):
|
| 217 |
+
"""
|
| 218 |
+
Input:
|
| 219 |
+
p : Tensor [B, Np(819_200), 3]
|
| 220 |
+
sparse_coords: Tensor [Nx, 4] ([batch_number, x, y, z])
|
| 221 |
+
|
| 222 |
+
Output:
|
| 223 |
+
sparse_pc_feats: [Nx, self.out_channels]
|
| 224 |
+
"""
|
| 225 |
+
batch_size, T, D = p.size()
|
| 226 |
+
max_coord_num = 0
|
| 227 |
+
for i in range(batch_size):
|
| 228 |
+
max_coord_num = max(
|
| 229 |
+
max_coord_num, (sparse_coords[..., 0] == i).sum().item() + 5
|
| 230 |
+
)
|
| 231 |
+
|
| 232 |
+
if D == self.in_channels:
|
| 233 |
+
p, normals = p[..., :3], p[..., 3:]
|
| 234 |
+
|
| 235 |
+
coord = scale_tensor(p, inp_scale=bbox_size) * res
|
| 236 |
+
p = 2 * (coord - (coord.floor() + 0.5)) # dist to the centrios, [-1., 1.]
|
| 237 |
+
index = self.coordinate2sparseindex(coord.long(), sparse_coords, res)
|
| 238 |
+
|
| 239 |
+
if D == self.in_channels:
|
| 240 |
+
p = torch.cat((p, normals), dim=-1)
|
| 241 |
+
net = self.fc_pos(p)
|
| 242 |
+
net = self.blocks[0](net)
|
| 243 |
+
for block in self.blocks[1:]:
|
| 244 |
+
pooled = self.pool_sparse_local(index, net, max_coord_num=max_coord_num)
|
| 245 |
+
|
| 246 |
+
net = torch.cat([net, pooled], dim=2)
|
| 247 |
+
net = block(net)
|
| 248 |
+
c = self.fc_c(net)
|
| 249 |
+
feats = self.generate_sparse_grid_features(
|
| 250 |
+
index, c, max_coord_num=max_coord_num
|
| 251 |
+
)
|
| 252 |
+
feats = self.convert_to_sparse_feats(feats, sparse_coords)
|
| 253 |
+
|
| 254 |
+
# torch.cuda.empty_cache()
|
| 255 |
+
return feats
|
| 256 |
+
|
| 257 |
+
|
| 258 |
+
class Pointnet(nn.Module):
|
| 259 |
+
def __init__(
|
| 260 |
+
self,
|
| 261 |
+
in_channels=16,
|
| 262 |
+
out_channels=32,
|
| 263 |
+
hidden_dim=32,
|
| 264 |
+
n_blocks=5,
|
| 265 |
+
use_checkpoint=True,
|
| 266 |
+
):
|
| 267 |
+
super().__init__()
|
| 268 |
+
self.in_channels = in_channels
|
| 269 |
+
self.out_channels = out_channels
|
| 270 |
+
self.hidden_dim = hidden_dim
|
| 271 |
+
self.use_checkpoint = use_checkpoint
|
| 272 |
+
|
| 273 |
+
self.fc_pos = nn.Linear(in_channels, 2 * hidden_dim)
|
| 274 |
+
|
| 275 |
+
self.blocks = nn.ModuleList(
|
| 276 |
+
[ResnetBlockFC(2 * hidden_dim, hidden_dim) for i in range(n_blocks)]
|
| 277 |
+
)
|
| 278 |
+
|
| 279 |
+
self.fc_c = nn.Linear(hidden_dim, out_channels)
|
| 280 |
+
|
| 281 |
+
self.initialize_weights()
|
| 282 |
+
|
| 283 |
+
def initialize_weights(self):
|
| 284 |
+
nn.init.xavier_uniform_(self.fc_pos.weight)
|
| 285 |
+
if self.fc_pos.bias is not None:
|
| 286 |
+
nn.init.constant_(self.fc_pos.bias, 0)
|
| 287 |
+
|
| 288 |
+
nn.init.xavier_uniform_(self.fc_c.weight)
|
| 289 |
+
if self.fc_c.bias is not None:
|
| 290 |
+
nn.init.constant_(self.fc_c.bias, 0)
|
| 291 |
+
|
| 292 |
+
@staticmethod
|
| 293 |
+
def _forward_block_concat(module, x):
|
| 294 |
+
return module(torch.cat([x, x], dim=-1))
|
| 295 |
+
|
| 296 |
+
def forward(self, p, res=64, bbox_size=(-0.5, 0.5)):
|
| 297 |
+
"""
|
| 298 |
+
Input:
|
| 299 |
+
p : Tensor [M, in_channels]
|
| 300 |
+
Output:
|
| 301 |
+
feats: Tensor [M, out_channels]
|
| 302 |
+
"""
|
| 303 |
+
|
| 304 |
+
pos_world = p[..., 0:3] # [M, 3]
|
| 305 |
+
other_feats = p[..., 3:] # [M, in_channels - 3]
|
| 306 |
+
|
| 307 |
+
scaled_pos = scale_tensor(pos_world, inp_scale=bbox_size) * res
|
| 308 |
+
local_pos = 2 * (scaled_pos - (scaled_pos.floor() + 0.5))
|
| 309 |
+
|
| 310 |
+
net_input = torch.cat([local_pos, other_feats], dim=-1)
|
| 311 |
+
|
| 312 |
+
net = self.fc_pos(net_input)
|
| 313 |
+
|
| 314 |
+
if self.use_checkpoint and net.requires_grad:
|
| 315 |
+
net = checkpoint(self.blocks[0], net, use_reentrant=False)
|
| 316 |
+
else:
|
| 317 |
+
net = self.blocks[0](net)
|
| 318 |
+
|
| 319 |
+
for block in self.blocks[1:]:
|
| 320 |
+
if self.use_checkpoint and net.requires_grad:
|
| 321 |
+
net = checkpoint(
|
| 322 |
+
self._forward_block_concat, block, net, use_reentrant=False
|
| 323 |
+
)
|
| 324 |
+
else:
|
| 325 |
+
net_concat = torch.cat([net, net], dim=-1)
|
| 326 |
+
net = block(net_concat)
|
| 327 |
+
|
| 328 |
+
feats = self.fc_c(net)
|
| 329 |
+
|
| 330 |
+
return feats
|
modules/sparse/__init__.py
ADDED
|
@@ -0,0 +1,130 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# MIT License
|
| 2 |
+
|
| 3 |
+
# Copyright (c) Microsoft Corporation.
|
| 4 |
+
# Copyright (c) 2025 VAST-AI-Research and contributors.
|
| 5 |
+
|
| 6 |
+
# Permission is hereby granted, free of charge, to any person obtaining a copy
|
| 7 |
+
# of this software and associated documentation files (the "Software"), to deal
|
| 8 |
+
# in the Software without restriction, including without limitation the rights
|
| 9 |
+
# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
|
| 10 |
+
# copies of the Software, and to permit persons to whom the Software is
|
| 11 |
+
# furnished to do so, subject to the following conditions:
|
| 12 |
+
|
| 13 |
+
# The above copyright notice and this permission notice shall be included in all
|
| 14 |
+
# copies or substantial portions of the Software.
|
| 15 |
+
|
| 16 |
+
# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
|
| 17 |
+
# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
|
| 18 |
+
# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
|
| 19 |
+
# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
|
| 20 |
+
# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
|
| 21 |
+
# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
|
| 22 |
+
# SOFTWARE
|
| 23 |
+
|
| 24 |
+
from typing import *
|
| 25 |
+
|
| 26 |
+
BACKEND = 'spconv'
|
| 27 |
+
DEBUG = False
|
| 28 |
+
ATTN = 'flash_attn'
|
| 29 |
+
|
| 30 |
+
def __from_env():
|
| 31 |
+
import os
|
| 32 |
+
|
| 33 |
+
global BACKEND
|
| 34 |
+
global DEBUG
|
| 35 |
+
global ATTN
|
| 36 |
+
|
| 37 |
+
env_sparse_backend = os.environ.get('SPARSE_BACKEND')
|
| 38 |
+
env_sparse_debug = os.environ.get('SPARSE_DEBUG')
|
| 39 |
+
env_sparse_attn = os.environ.get('SPARSE_ATTN_BACKEND')
|
| 40 |
+
if env_sparse_attn is None:
|
| 41 |
+
env_sparse_attn = os.environ.get('ATTN_BACKEND')
|
| 42 |
+
|
| 43 |
+
if env_sparse_backend is not None and env_sparse_backend in ['spconv', 'torchsparse']:
|
| 44 |
+
BACKEND = env_sparse_backend
|
| 45 |
+
if env_sparse_debug is not None:
|
| 46 |
+
DEBUG = env_sparse_debug == '1'
|
| 47 |
+
if env_sparse_attn is not None and env_sparse_attn in ['xformers', 'flash_attn']:
|
| 48 |
+
ATTN = env_sparse_attn
|
| 49 |
+
|
| 50 |
+
print(f"[SPARSE] Backend: {BACKEND}, Attention: {ATTN}")
|
| 51 |
+
|
| 52 |
+
|
| 53 |
+
__from_env()
|
| 54 |
+
|
| 55 |
+
|
| 56 |
+
def set_backend(backend: Literal['spconv', 'torchsparse']):
|
| 57 |
+
global BACKEND
|
| 58 |
+
BACKEND = backend
|
| 59 |
+
|
| 60 |
+
def set_debug(debug: bool):
|
| 61 |
+
global DEBUG
|
| 62 |
+
DEBUG = debug
|
| 63 |
+
|
| 64 |
+
def set_attn(attn: Literal['xformers', 'flash_attn']):
|
| 65 |
+
global ATTN
|
| 66 |
+
ATTN = attn
|
| 67 |
+
|
| 68 |
+
|
| 69 |
+
import importlib
|
| 70 |
+
|
| 71 |
+
__attributes = {
|
| 72 |
+
'SparseTensor': 'basic',
|
| 73 |
+
'sparse_batch_broadcast': 'basic',
|
| 74 |
+
'sparse_batch_op': 'basic',
|
| 75 |
+
'sparse_cat': 'basic',
|
| 76 |
+
'sparse_unbind': 'basic',
|
| 77 |
+
'SparseGroupNorm': 'norm',
|
| 78 |
+
'SparseLayerNorm': 'norm',
|
| 79 |
+
'SparseGroupNorm32': 'norm',
|
| 80 |
+
'SparseLayerNorm32': 'norm',
|
| 81 |
+
'SparseReLU': 'nonlinearity',
|
| 82 |
+
'SparseSiLU': 'nonlinearity',
|
| 83 |
+
'SparseGELU': 'nonlinearity',
|
| 84 |
+
'SparseActivation': 'nonlinearity',
|
| 85 |
+
'SparseLinear': 'linear',
|
| 86 |
+
'sparse_scaled_dot_product_attention': 'attention',
|
| 87 |
+
'SerializeMode': 'attention',
|
| 88 |
+
'SerializeModes': 'attention',
|
| 89 |
+
'sparse_serialized_scaled_dot_product_self_attention': 'attention',
|
| 90 |
+
'sparse_windowed_scaled_dot_product_self_attention': 'attention',
|
| 91 |
+
'SparseMultiHeadAttention': 'attention',
|
| 92 |
+
'SparseConv3d': 'conv',
|
| 93 |
+
'SparseInverseConv3d': 'conv',
|
| 94 |
+
'SparseDownsample': 'spatial',
|
| 95 |
+
'SparseUpsample': 'spatial',
|
| 96 |
+
'SparseSubdivide' : 'spatial',
|
| 97 |
+
|
| 98 |
+
'SparseSubdivide_attn' : 'spatial',
|
| 99 |
+
'SparseSpatial2Channel': 'spatial',
|
| 100 |
+
'SparseChannel2Spatial': 'spatial',
|
| 101 |
+
}
|
| 102 |
+
|
| 103 |
+
__submodules = ['transformer']
|
| 104 |
+
|
| 105 |
+
__all__ = list(__attributes.keys()) + __submodules
|
| 106 |
+
|
| 107 |
+
def __getattr__(name):
|
| 108 |
+
if name not in globals():
|
| 109 |
+
if name in __attributes:
|
| 110 |
+
module_name = __attributes[name]
|
| 111 |
+
module = importlib.import_module(f".{module_name}", __name__)
|
| 112 |
+
globals()[name] = getattr(module, name)
|
| 113 |
+
elif name in __submodules:
|
| 114 |
+
module = importlib.import_module(f".{name}", __name__)
|
| 115 |
+
globals()[name] = module
|
| 116 |
+
else:
|
| 117 |
+
raise AttributeError(f"module {__name__} has no attribute {name}")
|
| 118 |
+
return globals()[name]
|
| 119 |
+
|
| 120 |
+
|
| 121 |
+
# For Pylance
|
| 122 |
+
if __name__ == '__main__':
|
| 123 |
+
from .basic import *
|
| 124 |
+
from .norm import *
|
| 125 |
+
from .nonlinearity import *
|
| 126 |
+
from .linear import *
|
| 127 |
+
from .attention import *
|
| 128 |
+
from .conv import *
|
| 129 |
+
from .spatial import *
|
| 130 |
+
import transformer
|
modules/sparse/attention/__init__.py
ADDED
|
@@ -0,0 +1,27 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# MIT License
|
| 2 |
+
|
| 3 |
+
# Copyright (c) Microsoft Corporation.
|
| 4 |
+
# Copyright (c) 2025 VAST-AI-Research and contributors.
|
| 5 |
+
|
| 6 |
+
# Permission is hereby granted, free of charge, to any person obtaining a copy
|
| 7 |
+
# of this software and associated documentation files (the "Software"), to deal
|
| 8 |
+
# in the Software without restriction, including without limitation the rights
|
| 9 |
+
# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
|
| 10 |
+
# copies of the Software, and to permit persons to whom the Software is
|
| 11 |
+
# furnished to do so, subject to the following conditions:
|
| 12 |
+
|
| 13 |
+
# The above copyright notice and this permission notice shall be included in all
|
| 14 |
+
# copies or substantial portions of the Software.
|
| 15 |
+
|
| 16 |
+
# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
|
| 17 |
+
# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
|
| 18 |
+
# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
|
| 19 |
+
# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
|
| 20 |
+
# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
|
| 21 |
+
# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
|
| 22 |
+
# SOFTWARE
|
| 23 |
+
|
| 24 |
+
from .full_attn import *
|
| 25 |
+
from .serialized_attn import *
|
| 26 |
+
from .windowed_attn import *
|
| 27 |
+
from .modules import *
|
modules/sparse/attention/full_attn.py
ADDED
|
@@ -0,0 +1,238 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# MIT License
|
| 2 |
+
|
| 3 |
+
# Copyright (c) Microsoft Corporation.
|
| 4 |
+
# Copyright (c) 2025 VAST-AI-Research and contributors.
|
| 5 |
+
|
| 6 |
+
# Permission is hereby granted, free of charge, to any person obtaining a copy
|
| 7 |
+
# of this software and associated documentation files (the "Software"), to deal
|
| 8 |
+
# in the Software without restriction, including without limitation the rights
|
| 9 |
+
# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
|
| 10 |
+
# copies of the Software, and to permit persons to whom the Software is
|
| 11 |
+
# furnished to do so, subject to the following conditions:
|
| 12 |
+
|
| 13 |
+
# The above copyright notice and this permission notice shall be included in all
|
| 14 |
+
# copies or substantial portions of the Software.
|
| 15 |
+
|
| 16 |
+
# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
|
| 17 |
+
# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
|
| 18 |
+
# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
|
| 19 |
+
# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
|
| 20 |
+
# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
|
| 21 |
+
# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
|
| 22 |
+
# SOFTWARE
|
| 23 |
+
|
| 24 |
+
from typing import *
|
| 25 |
+
import torch
|
| 26 |
+
from .. import SparseTensor
|
| 27 |
+
from .. import DEBUG, ATTN
|
| 28 |
+
|
| 29 |
+
if ATTN == 'xformers':
|
| 30 |
+
import xformers.ops as xops
|
| 31 |
+
elif ATTN == 'flash_attn':
|
| 32 |
+
import flash_attn
|
| 33 |
+
else:
|
| 34 |
+
raise ValueError(f"Unknown attention module: {ATTN}")
|
| 35 |
+
|
| 36 |
+
|
| 37 |
+
__all__ = [
|
| 38 |
+
'sparse_scaled_dot_product_attention',
|
| 39 |
+
]
|
| 40 |
+
|
| 41 |
+
|
| 42 |
+
@overload
|
| 43 |
+
def sparse_scaled_dot_product_attention(qkv: SparseTensor) -> SparseTensor:
|
| 44 |
+
"""
|
| 45 |
+
Apply scaled dot product attention to a sparse tensor.
|
| 46 |
+
|
| 47 |
+
Args:
|
| 48 |
+
qkv (SparseTensor): A [N, *, 3, H, C] sparse tensor containing Qs, Ks, and Vs.
|
| 49 |
+
"""
|
| 50 |
+
...
|
| 51 |
+
|
| 52 |
+
@overload
|
| 53 |
+
def sparse_scaled_dot_product_attention(q: SparseTensor, kv: Union[SparseTensor, torch.Tensor]) -> SparseTensor:
|
| 54 |
+
"""
|
| 55 |
+
Apply scaled dot product attention to a sparse tensor.
|
| 56 |
+
|
| 57 |
+
Args:
|
| 58 |
+
q (SparseTensor): A [N, *, H, C] sparse tensor containing Qs.
|
| 59 |
+
kv (SparseTensor or torch.Tensor): A [N, *, 2, H, C] sparse tensor or a [N, L, 2, H, C] dense tensor containing Ks and Vs.
|
| 60 |
+
"""
|
| 61 |
+
...
|
| 62 |
+
|
| 63 |
+
@overload
|
| 64 |
+
def sparse_scaled_dot_product_attention(q: torch.Tensor, kv: SparseTensor) -> torch.Tensor:
|
| 65 |
+
"""
|
| 66 |
+
Apply scaled dot product attention to a sparse tensor.
|
| 67 |
+
|
| 68 |
+
Args:
|
| 69 |
+
q (SparseTensor): A [N, L, H, C] dense tensor containing Qs.
|
| 70 |
+
kv (SparseTensor or torch.Tensor): A [N, *, 2, H, C] sparse tensor containing Ks and Vs.
|
| 71 |
+
"""
|
| 72 |
+
...
|
| 73 |
+
|
| 74 |
+
@overload
|
| 75 |
+
def sparse_scaled_dot_product_attention(q: SparseTensor, k: SparseTensor, v: SparseTensor) -> SparseTensor:
|
| 76 |
+
"""
|
| 77 |
+
Apply scaled dot product attention to a sparse tensor.
|
| 78 |
+
|
| 79 |
+
Args:
|
| 80 |
+
q (SparseTensor): A [N, *, H, Ci] sparse tensor containing Qs.
|
| 81 |
+
k (SparseTensor): A [N, *, H, Ci] sparse tensor containing Ks.
|
| 82 |
+
v (SparseTensor): A [N, *, H, Co] sparse tensor containing Vs.
|
| 83 |
+
|
| 84 |
+
Note:
|
| 85 |
+
k and v are assumed to have the same coordinate map.
|
| 86 |
+
"""
|
| 87 |
+
...
|
| 88 |
+
|
| 89 |
+
@overload
|
| 90 |
+
def sparse_scaled_dot_product_attention(q: SparseTensor, k: torch.Tensor, v: torch.Tensor) -> SparseTensor:
|
| 91 |
+
"""
|
| 92 |
+
Apply scaled dot product attention to a sparse tensor.
|
| 93 |
+
|
| 94 |
+
Args:
|
| 95 |
+
q (SparseTensor): A [N, *, H, Ci] sparse tensor containing Qs.
|
| 96 |
+
k (torch.Tensor): A [N, L, H, Ci] dense tensor containing Ks.
|
| 97 |
+
v (torch.Tensor): A [N, L, H, Co] dense tensor containing Vs.
|
| 98 |
+
"""
|
| 99 |
+
...
|
| 100 |
+
|
| 101 |
+
@overload
|
| 102 |
+
def sparse_scaled_dot_product_attention(q: torch.Tensor, k: SparseTensor, v: SparseTensor) -> torch.Tensor:
|
| 103 |
+
"""
|
| 104 |
+
Apply scaled dot product attention to a sparse tensor.
|
| 105 |
+
|
| 106 |
+
Args:
|
| 107 |
+
q (torch.Tensor): A [N, L, H, Ci] dense tensor containing Qs.
|
| 108 |
+
k (SparseTensor): A [N, *, H, Ci] sparse tensor containing Ks.
|
| 109 |
+
v (SparseTensor): A [N, *, H, Co] sparse tensor containing Vs.
|
| 110 |
+
"""
|
| 111 |
+
...
|
| 112 |
+
|
| 113 |
+
def sparse_scaled_dot_product_attention(*args, **kwargs):
|
| 114 |
+
arg_names_dict = {
|
| 115 |
+
1: ['qkv'],
|
| 116 |
+
2: ['q', 'kv'],
|
| 117 |
+
3: ['q', 'k', 'v']
|
| 118 |
+
}
|
| 119 |
+
num_all_args = len(args) + len(kwargs)
|
| 120 |
+
assert num_all_args in arg_names_dict, f"Invalid number of arguments, got {num_all_args}, expected 1, 2, or 3"
|
| 121 |
+
for key in arg_names_dict[num_all_args][len(args):]:
|
| 122 |
+
assert key in kwargs, f"Missing argument {key}"
|
| 123 |
+
|
| 124 |
+
if num_all_args == 1:
|
| 125 |
+
qkv = args[0] if len(args) > 0 else kwargs['qkv']
|
| 126 |
+
assert isinstance(qkv, SparseTensor), f"qkv must be a SparseTensor, got {type(qkv)}"
|
| 127 |
+
assert len(qkv.shape) == 4 and qkv.shape[1] == 3, f"Invalid shape for qkv, got {qkv.shape}, expected [N, *, 3, H, C]"
|
| 128 |
+
device = qkv.device
|
| 129 |
+
|
| 130 |
+
s = qkv
|
| 131 |
+
q_seqlen = [qkv.layout[i].stop - qkv.layout[i].start for i in range(qkv.shape[0])]
|
| 132 |
+
kv_seqlen = q_seqlen
|
| 133 |
+
qkv = qkv.feats # [T, 3, H, C]
|
| 134 |
+
|
| 135 |
+
elif num_all_args == 2:
|
| 136 |
+
q = args[0] if len(args) > 0 else kwargs['q']
|
| 137 |
+
kv = args[1] if len(args) > 1 else kwargs['kv']
|
| 138 |
+
assert isinstance(q, SparseTensor) and isinstance(kv, (SparseTensor, torch.Tensor)) or \
|
| 139 |
+
isinstance(q, torch.Tensor) and isinstance(kv, SparseTensor), \
|
| 140 |
+
f"Invalid types, got {type(q)} and {type(kv)}"
|
| 141 |
+
assert q.shape[0] == kv.shape[0], f"Batch size mismatch, got {q.shape[0]} and {kv.shape[0]}"
|
| 142 |
+
device = q.device
|
| 143 |
+
|
| 144 |
+
if isinstance(q, SparseTensor):
|
| 145 |
+
assert len(q.shape) == 3, f"Invalid shape for q, got {q.shape}, expected [N, *, H, C]"
|
| 146 |
+
s = q
|
| 147 |
+
q_seqlen = [q.layout[i].stop - q.layout[i].start for i in range(q.shape[0])]
|
| 148 |
+
q = q.feats # [T_Q, H, C]
|
| 149 |
+
else:
|
| 150 |
+
assert len(q.shape) == 4, f"Invalid shape for q, got {q.shape}, expected [N, L, H, C]"
|
| 151 |
+
s = None
|
| 152 |
+
N, L, H, C = q.shape
|
| 153 |
+
q_seqlen = [L] * N
|
| 154 |
+
q = q.reshape(N * L, H, C) # [T_Q, H, C]
|
| 155 |
+
|
| 156 |
+
if isinstance(kv, SparseTensor):
|
| 157 |
+
assert len(kv.shape) == 4 and kv.shape[1] == 2, f"Invalid shape for kv, got {kv.shape}, expected [N, *, 2, H, C]"
|
| 158 |
+
kv_seqlen = [kv.layout[i].stop - kv.layout[i].start for i in range(kv.shape[0])]
|
| 159 |
+
kv = kv.feats # [T_KV, 2, H, C]
|
| 160 |
+
else:
|
| 161 |
+
assert len(kv.shape) == 5, f"Invalid shape for kv, got {kv.shape}, expected [N, L, 2, H, C]"
|
| 162 |
+
N, L, _, H, C = kv.shape
|
| 163 |
+
kv_seqlen = [L] * N
|
| 164 |
+
kv = kv.reshape(N * L, 2, H, C) # [T_KV, 2, H, C]
|
| 165 |
+
|
| 166 |
+
elif num_all_args == 3:
|
| 167 |
+
q = args[0] if len(args) > 0 else kwargs['q']
|
| 168 |
+
k = args[1] if len(args) > 1 else kwargs['k']
|
| 169 |
+
v = args[2] if len(args) > 2 else kwargs['v']
|
| 170 |
+
assert isinstance(q, SparseTensor) and isinstance(k, (SparseTensor, torch.Tensor)) and type(k) == type(v) or \
|
| 171 |
+
isinstance(q, torch.Tensor) and isinstance(k, SparseTensor) and isinstance(v, SparseTensor), \
|
| 172 |
+
f"Invalid types, got {type(q)}, {type(k)}, and {type(v)}"
|
| 173 |
+
assert q.shape[0] == k.shape[0] == v.shape[0], f"Batch size mismatch, got {q.shape[0]}, {k.shape[0]}, and {v.shape[0]}"
|
| 174 |
+
device = q.device
|
| 175 |
+
|
| 176 |
+
if isinstance(q, SparseTensor):
|
| 177 |
+
assert len(q.shape) == 3, f"Invalid shape for q, got {q.shape}, expected [N, *, H, Ci]"
|
| 178 |
+
s = q
|
| 179 |
+
q_seqlen = [q.layout[i].stop - q.layout[i].start for i in range(q.shape[0])]
|
| 180 |
+
q = q.feats # [T_Q, H, Ci]
|
| 181 |
+
else:
|
| 182 |
+
assert len(q.shape) == 4, f"Invalid shape for q, got {q.shape}, expected [N, L, H, Ci]"
|
| 183 |
+
s = None
|
| 184 |
+
N, L, H, CI = q.shape
|
| 185 |
+
q_seqlen = [L] * N
|
| 186 |
+
q = q.reshape(N * L, H, CI) # [T_Q, H, Ci]
|
| 187 |
+
|
| 188 |
+
if isinstance(k, SparseTensor):
|
| 189 |
+
assert len(k.shape) == 3, f"Invalid shape for k, got {k.shape}, expected [N, *, H, Ci]"
|
| 190 |
+
assert len(v.shape) == 3, f"Invalid shape for v, got {v.shape}, expected [N, *, H, Co]"
|
| 191 |
+
kv_seqlen = [k.layout[i].stop - k.layout[i].start for i in range(k.shape[0])]
|
| 192 |
+
k = k.feats # [T_KV, H, Ci]
|
| 193 |
+
v = v.feats # [T_KV, H, Co]
|
| 194 |
+
else:
|
| 195 |
+
assert len(k.shape) == 4, f"Invalid shape for k, got {k.shape}, expected [N, L, H, Ci]"
|
| 196 |
+
assert len(v.shape) == 4, f"Invalid shape for v, got {v.shape}, expected [N, L, H, Co]"
|
| 197 |
+
N, L, H, CI, CO = *k.shape, v.shape[-1]
|
| 198 |
+
kv_seqlen = [L] * N
|
| 199 |
+
k = k.reshape(N * L, H, CI) # [T_KV, H, Ci]
|
| 200 |
+
v = v.reshape(N * L, H, CO) # [T_KV, H, Co]
|
| 201 |
+
|
| 202 |
+
if DEBUG:
|
| 203 |
+
if s is not None:
|
| 204 |
+
for i in range(s.shape[0]):
|
| 205 |
+
assert (s.coords[s.layout[i]] == i).all(), f"SparseScaledDotProductSelfAttention: batch index mismatch"
|
| 206 |
+
if num_all_args in [2, 3]:
|
| 207 |
+
assert q.shape[:2] == [1, sum(q_seqlen)], f"SparseScaledDotProductSelfAttention: q shape mismatch"
|
| 208 |
+
if num_all_args == 3:
|
| 209 |
+
assert k.shape[:2] == [1, sum(kv_seqlen)], f"SparseScaledDotProductSelfAttention: k shape mismatch"
|
| 210 |
+
assert v.shape[:2] == [1, sum(kv_seqlen)], f"SparseScaledDotProductSelfAttention: v shape mismatch"
|
| 211 |
+
|
| 212 |
+
if ATTN == 'xformers':
|
| 213 |
+
if num_all_args == 1:
|
| 214 |
+
q, k, v = qkv.unbind(dim=1)
|
| 215 |
+
elif num_all_args == 2:
|
| 216 |
+
k, v = kv.unbind(dim=1)
|
| 217 |
+
q = q.unsqueeze(0)
|
| 218 |
+
k = k.unsqueeze(0)
|
| 219 |
+
v = v.unsqueeze(0)
|
| 220 |
+
mask = xops.fmha.BlockDiagonalMask.from_seqlens(q_seqlen, kv_seqlen)
|
| 221 |
+
out = xops.memory_efficient_attention(q, k, v, mask)[0]
|
| 222 |
+
elif ATTN == 'flash_attn':
|
| 223 |
+
cu_seqlens_q = torch.cat([torch.tensor([0]), torch.cumsum(torch.tensor(q_seqlen), dim=0)]).int().to(device)
|
| 224 |
+
if num_all_args in [2, 3]:
|
| 225 |
+
cu_seqlens_kv = torch.cat([torch.tensor([0]), torch.cumsum(torch.tensor(kv_seqlen), dim=0)]).int().to(device)
|
| 226 |
+
if num_all_args == 1:
|
| 227 |
+
out = flash_attn.flash_attn_varlen_qkvpacked_func(qkv, cu_seqlens_q, max(q_seqlen))
|
| 228 |
+
elif num_all_args == 2:
|
| 229 |
+
out = flash_attn.flash_attn_varlen_kvpacked_func(q, kv, cu_seqlens_q, cu_seqlens_kv, max(q_seqlen), max(kv_seqlen))
|
| 230 |
+
elif num_all_args == 3:
|
| 231 |
+
out = flash_attn.flash_attn_varlen_func(q, k, v, cu_seqlens_q, cu_seqlens_kv, max(q_seqlen), max(kv_seqlen))
|
| 232 |
+
else:
|
| 233 |
+
raise ValueError(f"Unknown attention module: {ATTN}")
|
| 234 |
+
|
| 235 |
+
if s is not None:
|
| 236 |
+
return s.replace(out)
|
| 237 |
+
else:
|
| 238 |
+
return out.reshape(N, L, H, -1)
|
modules/sparse/attention/modules.py
ADDED
|
@@ -0,0 +1,214 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# MIT License
|
| 2 |
+
|
| 3 |
+
# Copyright (c) Microsoft Corporation.
|
| 4 |
+
# Copyright (c) 2025 VAST-AI-Research and contributors.
|
| 5 |
+
|
| 6 |
+
# Permission is hereby granted, free of charge, to any person obtaining a copy
|
| 7 |
+
# of this software and associated documentation files (the "Software"), to deal
|
| 8 |
+
# in the Software without restriction, including without limitation the rights
|
| 9 |
+
# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
|
| 10 |
+
# copies of the Software, and to permit persons to whom the Software is
|
| 11 |
+
# furnished to do so, subject to the following conditions:
|
| 12 |
+
|
| 13 |
+
# The above copyright notice and this permission notice shall be included in all
|
| 14 |
+
# copies or substantial portions of the Software.
|
| 15 |
+
|
| 16 |
+
# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
|
| 17 |
+
# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
|
| 18 |
+
# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
|
| 19 |
+
# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
|
| 20 |
+
# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
|
| 21 |
+
# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
|
| 22 |
+
# SOFTWARE
|
| 23 |
+
|
| 24 |
+
from typing import *
|
| 25 |
+
import torch
|
| 26 |
+
import torch.nn as nn
|
| 27 |
+
import torch.nn.functional as F
|
| 28 |
+
from .. import SparseTensor
|
| 29 |
+
from .full_attn import sparse_scaled_dot_product_attention
|
| 30 |
+
from .serialized_attn import SerializeMode, sparse_serialized_scaled_dot_product_self_attention
|
| 31 |
+
from .windowed_attn import sparse_windowed_scaled_dot_product_self_attention
|
| 32 |
+
|
| 33 |
+
|
| 34 |
+
class RotaryPositionEmbedder(nn.Module):
|
| 35 |
+
def __init__(self, hidden_size: int, in_channels: int = 3):
|
| 36 |
+
super().__init__()
|
| 37 |
+
assert hidden_size % 2 == 0, "Hidden size must be divisible by 2"
|
| 38 |
+
self.hidden_size = hidden_size
|
| 39 |
+
self.in_channels = in_channels
|
| 40 |
+
self.freq_dim = hidden_size // in_channels // 2
|
| 41 |
+
self.freqs = torch.arange(self.freq_dim, dtype=torch.float32) / self.freq_dim
|
| 42 |
+
self.freqs = 1.0 / (10000 ** self.freqs)
|
| 43 |
+
|
| 44 |
+
def _get_phases(self, indices: torch.Tensor) -> torch.Tensor:
|
| 45 |
+
self.freqs = self.freqs.to(indices.device)
|
| 46 |
+
phases = torch.outer(indices, self.freqs)
|
| 47 |
+
phases = torch.polar(torch.ones_like(phases), phases)
|
| 48 |
+
return phases
|
| 49 |
+
|
| 50 |
+
def _rotary_embedding(self, x: torch.Tensor, phases: torch.Tensor) -> torch.Tensor:
|
| 51 |
+
x_complex = torch.view_as_complex(x.float().reshape(*x.shape[:-1], -1, 2))
|
| 52 |
+
|
| 53 |
+
if phases.dim() == x_complex.dim() - 1:
|
| 54 |
+
phases = phases.unsqueeze(-2)
|
| 55 |
+
|
| 56 |
+
x_rotated = x_complex * phases
|
| 57 |
+
x_embed = torch.view_as_real(x_rotated).reshape(*x_rotated.shape[:-1], -1).to(x.dtype)
|
| 58 |
+
return x_embed
|
| 59 |
+
|
| 60 |
+
def forward(self, q: torch.Tensor, k: torch.Tensor, indices: Optional[torch.Tensor] = None) -> Tuple[torch.Tensor, torch.Tensor]:
|
| 61 |
+
"""
|
| 62 |
+
Args:
|
| 63 |
+
q (torch.Tensor): [..., N, D] tensor of queries
|
| 64 |
+
k (torch.Tensor): [..., N, D] tensor of keys
|
| 65 |
+
indices (torch.Tensor): [..., N, C] tensor of spatial positions
|
| 66 |
+
"""
|
| 67 |
+
if indices is None:
|
| 68 |
+
indices = torch.arange(q.shape[-2], device=q.device)
|
| 69 |
+
if len(q.shape) > 2:
|
| 70 |
+
indices = indices.unsqueeze(0).expand(q.shape[:-2] + (-1,))
|
| 71 |
+
|
| 72 |
+
phases = self._get_phases(indices.reshape(-1)).reshape(*indices.shape[:-1], -1)
|
| 73 |
+
if phases.shape[1] < self.hidden_size // 2:
|
| 74 |
+
phases = torch.cat([phases, torch.polar(
|
| 75 |
+
torch.ones(*phases.shape[:-1], self.hidden_size // 2 - phases.shape[1], device=phases.device),
|
| 76 |
+
torch.zeros(*phases.shape[:-1], self.hidden_size // 2 - phases.shape[1], device=phases.device)
|
| 77 |
+
)], dim=-1)
|
| 78 |
+
q_embed = self._rotary_embedding(q, phases)
|
| 79 |
+
k_embed = self._rotary_embedding(k, phases)
|
| 80 |
+
return q_embed, k_embed
|
| 81 |
+
|
| 82 |
+
|
| 83 |
+
class SparseMultiHeadRMSNorm(nn.Module):
|
| 84 |
+
def __init__(self, dim: int, heads: int):
|
| 85 |
+
super().__init__()
|
| 86 |
+
self.scale = dim ** 0.5
|
| 87 |
+
self.gamma = nn.Parameter(torch.ones(heads, dim))
|
| 88 |
+
|
| 89 |
+
def forward(self, x: Union[SparseTensor, torch.Tensor]) -> Union[SparseTensor, torch.Tensor]:
|
| 90 |
+
x_type = x.dtype
|
| 91 |
+
x = x.float()
|
| 92 |
+
if isinstance(x, SparseTensor):
|
| 93 |
+
x = x.replace(F.normalize(x.feats, dim=-1))
|
| 94 |
+
else:
|
| 95 |
+
x = F.normalize(x, dim=-1)
|
| 96 |
+
return (x * self.gamma * self.scale).to(x_type)
|
| 97 |
+
|
| 98 |
+
|
| 99 |
+
class SparseMultiHeadAttention(nn.Module):
|
| 100 |
+
def __init__(
|
| 101 |
+
self,
|
| 102 |
+
channels: int,
|
| 103 |
+
num_heads: int,
|
| 104 |
+
ctx_channels: Optional[int] = None,
|
| 105 |
+
type: Literal["self", "cross"] = "self",
|
| 106 |
+
attn_mode: Literal["full", "serialized", "windowed"] = "full",
|
| 107 |
+
window_size: Optional[int] = None,
|
| 108 |
+
shift_sequence: Optional[int] = None,
|
| 109 |
+
shift_window: Optional[Tuple[int, int, int]] = None,
|
| 110 |
+
serialize_mode: Optional[SerializeMode] = None,
|
| 111 |
+
qkv_bias: bool = True,
|
| 112 |
+
use_rope: bool = False,
|
| 113 |
+
qk_rms_norm: bool = False,
|
| 114 |
+
):
|
| 115 |
+
super().__init__()
|
| 116 |
+
assert channels % num_heads == 0
|
| 117 |
+
assert type in ["self", "cross"], f"Invalid attention type: {type}"
|
| 118 |
+
assert attn_mode in ["full", "serialized", "windowed"], f"Invalid attention mode: {attn_mode}"
|
| 119 |
+
assert type == "self" or attn_mode == "full", "Cross-attention only supports full attention"
|
| 120 |
+
assert type == "self" or use_rope is False, "Rotary position embeddings only supported for self-attention"
|
| 121 |
+
self.channels = channels
|
| 122 |
+
self.ctx_channels = ctx_channels if ctx_channels is not None else channels
|
| 123 |
+
self.num_heads = num_heads
|
| 124 |
+
self._type = type
|
| 125 |
+
self.attn_mode = attn_mode
|
| 126 |
+
self.window_size = window_size
|
| 127 |
+
self.shift_sequence = shift_sequence
|
| 128 |
+
self.shift_window = shift_window
|
| 129 |
+
self.serialize_mode = serialize_mode
|
| 130 |
+
self.use_rope = use_rope
|
| 131 |
+
self.qk_rms_norm = qk_rms_norm
|
| 132 |
+
|
| 133 |
+
if self._type == "self":
|
| 134 |
+
self.to_qkv = nn.Linear(channels, channels * 3, bias=qkv_bias)
|
| 135 |
+
else:
|
| 136 |
+
self.to_q = nn.Linear(channels, channels, bias=qkv_bias)
|
| 137 |
+
self.to_kv = nn.Linear(self.ctx_channels, channels * 2, bias=qkv_bias)
|
| 138 |
+
|
| 139 |
+
if self.qk_rms_norm:
|
| 140 |
+
self.q_rms_norm = SparseMultiHeadRMSNorm(channels // num_heads, num_heads)
|
| 141 |
+
self.k_rms_norm = SparseMultiHeadRMSNorm(channels // num_heads, num_heads)
|
| 142 |
+
|
| 143 |
+
self.to_out = nn.Linear(channels, channels)
|
| 144 |
+
|
| 145 |
+
if use_rope:
|
| 146 |
+
# self.rope = RotaryPositionEmbedder(channels)
|
| 147 |
+
|
| 148 |
+
head_dim = channels // self.num_heads
|
| 149 |
+
self.rope = RotaryPositionEmbedder(head_dim)
|
| 150 |
+
|
| 151 |
+
|
| 152 |
+
@staticmethod
|
| 153 |
+
def _linear(module: nn.Linear, x: Union[SparseTensor, torch.Tensor]) -> Union[SparseTensor, torch.Tensor]:
|
| 154 |
+
if isinstance(x, SparseTensor):
|
| 155 |
+
return x.replace(module(x.feats))
|
| 156 |
+
else:
|
| 157 |
+
return module(x)
|
| 158 |
+
|
| 159 |
+
@staticmethod
|
| 160 |
+
def _reshape_chs(x: Union[SparseTensor, torch.Tensor], shape: Tuple[int, ...]) -> Union[SparseTensor, torch.Tensor]:
|
| 161 |
+
if isinstance(x, SparseTensor):
|
| 162 |
+
return x.reshape(*shape)
|
| 163 |
+
else:
|
| 164 |
+
return x.reshape(*x.shape[:2], *shape)
|
| 165 |
+
|
| 166 |
+
def _fused_pre(self, x: Union[SparseTensor, torch.Tensor], num_fused: int) -> Union[SparseTensor, torch.Tensor]:
|
| 167 |
+
if isinstance(x, SparseTensor):
|
| 168 |
+
x_feats = x.feats.unsqueeze(0)
|
| 169 |
+
else:
|
| 170 |
+
x_feats = x
|
| 171 |
+
x_feats = x_feats.reshape(*x_feats.shape[:2], num_fused, self.num_heads, -1)
|
| 172 |
+
return x.replace(x_feats.squeeze(0)) if isinstance(x, SparseTensor) else x_feats
|
| 173 |
+
|
| 174 |
+
def _rope(self, qkv: SparseTensor) -> SparseTensor:
|
| 175 |
+
q, k, v = qkv.feats.unbind(dim=1) # [T, H, C]
|
| 176 |
+
q, k = self.rope(q, k, qkv.coords[:, 1:])
|
| 177 |
+
qkv = qkv.replace(torch.stack([q, k, v], dim=1))
|
| 178 |
+
return qkv
|
| 179 |
+
|
| 180 |
+
def forward(self, x: Union[SparseTensor, torch.Tensor], context: Optional[Union[SparseTensor, torch.Tensor]] = None) -> Union[SparseTensor, torch.Tensor]:
|
| 181 |
+
if self._type == "self": # self-attn, default
|
| 182 |
+
qkv = self._linear(self.to_qkv, x)
|
| 183 |
+
qkv = self._fused_pre(qkv, num_fused=3) # to reshape
|
| 184 |
+
if self.use_rope: # False, default
|
| 185 |
+
qkv = self._rope(qkv)
|
| 186 |
+
if self.qk_rms_norm:
|
| 187 |
+
q, k, v = qkv.unbind(dim=1)
|
| 188 |
+
q = self.q_rms_norm(q)
|
| 189 |
+
k = self.k_rms_norm(k)
|
| 190 |
+
qkv = qkv.replace(torch.stack([q.feats, k.feats, v.feats], dim=1))
|
| 191 |
+
if self.attn_mode == "full":
|
| 192 |
+
h = sparse_scaled_dot_product_attention(qkv)
|
| 193 |
+
elif self.attn_mode == "serialized":
|
| 194 |
+
h = sparse_serialized_scaled_dot_product_self_attention(
|
| 195 |
+
qkv, self.window_size, serialize_mode=self.serialize_mode, shift_sequence=self.shift_sequence, shift_window=self.shift_window
|
| 196 |
+
)
|
| 197 |
+
elif self.attn_mode == "windowed":
|
| 198 |
+
h = sparse_windowed_scaled_dot_product_self_attention(
|
| 199 |
+
qkv, self.window_size, shift_window=self.shift_window
|
| 200 |
+
)
|
| 201 |
+
else: # cross attn, default False
|
| 202 |
+
q = self._linear(self.to_q, x)
|
| 203 |
+
q = self._reshape_chs(q, (self.num_heads, -1))
|
| 204 |
+
kv = self._linear(self.to_kv, context)
|
| 205 |
+
kv = self._fused_pre(kv, num_fused=2)
|
| 206 |
+
if self.qk_rms_norm:
|
| 207 |
+
q = self.q_rms_norm(q)
|
| 208 |
+
k, v = kv.unbind(dim=1)
|
| 209 |
+
k = self.k_rms_norm(k)
|
| 210 |
+
kv = kv.replace(torch.stack([k.feats, v.feats], dim=1))
|
| 211 |
+
h = sparse_scaled_dot_product_attention(q, kv)
|
| 212 |
+
h = self._reshape_chs(h, (-1,))
|
| 213 |
+
h = self._linear(self.to_out, h)
|
| 214 |
+
return h
|
modules/sparse/attention/serialized_attn.py
ADDED
|
@@ -0,0 +1,217 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# MIT License
|
| 2 |
+
|
| 3 |
+
# Copyright (c) Microsoft Corporation.
|
| 4 |
+
# Copyright (c) 2025 VAST-AI-Research and contributors.
|
| 5 |
+
|
| 6 |
+
# Permission is hereby granted, free of charge, to any person obtaining a copy
|
| 7 |
+
# of this software and associated documentation files (the "Software"), to deal
|
| 8 |
+
# in the Software without restriction, including without limitation the rights
|
| 9 |
+
# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
|
| 10 |
+
# copies of the Software, and to permit persons to whom the Software is
|
| 11 |
+
# furnished to do so, subject to the following conditions:
|
| 12 |
+
|
| 13 |
+
# The above copyright notice and this permission notice shall be included in all
|
| 14 |
+
# copies or substantial portions of the Software.
|
| 15 |
+
|
| 16 |
+
# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
|
| 17 |
+
# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
|
| 18 |
+
# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
|
| 19 |
+
# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
|
| 20 |
+
# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
|
| 21 |
+
# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
|
| 22 |
+
# SOFTWARE
|
| 23 |
+
|
| 24 |
+
from typing import *
|
| 25 |
+
from enum import Enum
|
| 26 |
+
import torch
|
| 27 |
+
import math
|
| 28 |
+
from .. import SparseTensor
|
| 29 |
+
from .. import DEBUG, ATTN
|
| 30 |
+
|
| 31 |
+
if ATTN == 'xformers':
|
| 32 |
+
import xformers.ops as xops
|
| 33 |
+
elif ATTN == 'flash_attn':
|
| 34 |
+
import flash_attn
|
| 35 |
+
else:
|
| 36 |
+
raise ValueError(f"Unknown attention module: {ATTN}")
|
| 37 |
+
|
| 38 |
+
|
| 39 |
+
__all__ = [
|
| 40 |
+
'sparse_serialized_scaled_dot_product_self_attention',
|
| 41 |
+
'SerializeModes',
|
| 42 |
+
]
|
| 43 |
+
|
| 44 |
+
|
| 45 |
+
class SerializeMode(Enum):
|
| 46 |
+
Z_ORDER = 0
|
| 47 |
+
Z_ORDER_TRANSPOSED = 1
|
| 48 |
+
HILBERT = 2
|
| 49 |
+
HILBERT_TRANSPOSED = 3
|
| 50 |
+
|
| 51 |
+
|
| 52 |
+
SerializeModes = [
|
| 53 |
+
SerializeMode.Z_ORDER,
|
| 54 |
+
SerializeMode.Z_ORDER_TRANSPOSED,
|
| 55 |
+
SerializeMode.HILBERT,
|
| 56 |
+
SerializeMode.HILBERT_TRANSPOSED
|
| 57 |
+
]
|
| 58 |
+
|
| 59 |
+
|
| 60 |
+
def calc_serialization(
|
| 61 |
+
tensor: SparseTensor,
|
| 62 |
+
window_size: int,
|
| 63 |
+
serialize_mode: SerializeMode = SerializeMode.Z_ORDER,
|
| 64 |
+
shift_sequence: int = 0,
|
| 65 |
+
shift_window: Tuple[int, int, int] = (0, 0, 0)
|
| 66 |
+
) -> Tuple[torch.Tensor, torch.Tensor, List[int]]:
|
| 67 |
+
"""
|
| 68 |
+
Calculate serialization and partitioning for a set of coordinates.
|
| 69 |
+
|
| 70 |
+
Args:
|
| 71 |
+
tensor (SparseTensor): The input tensor.
|
| 72 |
+
window_size (int): The window size to use.
|
| 73 |
+
serialize_mode (SerializeMode): The serialization mode to use.
|
| 74 |
+
shift_sequence (int): The shift of serialized sequence.
|
| 75 |
+
shift_window (Tuple[int, int, int]): The shift of serialized coordinates.
|
| 76 |
+
|
| 77 |
+
Returns:
|
| 78 |
+
(torch.Tensor, torch.Tensor): Forwards and backwards indices.
|
| 79 |
+
"""
|
| 80 |
+
fwd_indices = []
|
| 81 |
+
bwd_indices = []
|
| 82 |
+
seq_lens = []
|
| 83 |
+
seq_batch_indices = []
|
| 84 |
+
offsets = [0]
|
| 85 |
+
|
| 86 |
+
if 'vox2seq' not in globals():
|
| 87 |
+
import vox2seq
|
| 88 |
+
|
| 89 |
+
# Serialize the input
|
| 90 |
+
serialize_coords = tensor.coords[:, 1:].clone()
|
| 91 |
+
serialize_coords += torch.tensor(shift_window, dtype=torch.int32, device=tensor.device).reshape(1, 3)
|
| 92 |
+
if serialize_mode == SerializeMode.Z_ORDER:
|
| 93 |
+
code = vox2seq.encode(serialize_coords, mode='z_order', permute=[0, 1, 2])
|
| 94 |
+
elif serialize_mode == SerializeMode.Z_ORDER_TRANSPOSED:
|
| 95 |
+
code = vox2seq.encode(serialize_coords, mode='z_order', permute=[1, 0, 2])
|
| 96 |
+
elif serialize_mode == SerializeMode.HILBERT:
|
| 97 |
+
code = vox2seq.encode(serialize_coords, mode='hilbert', permute=[0, 1, 2])
|
| 98 |
+
elif serialize_mode == SerializeMode.HILBERT_TRANSPOSED:
|
| 99 |
+
code = vox2seq.encode(serialize_coords, mode='hilbert', permute=[1, 0, 2])
|
| 100 |
+
else:
|
| 101 |
+
raise ValueError(f"Unknown serialize mode: {serialize_mode}")
|
| 102 |
+
|
| 103 |
+
for bi, s in enumerate(tensor.layout):
|
| 104 |
+
num_points = s.stop - s.start
|
| 105 |
+
num_windows = (num_points + window_size - 1) // window_size
|
| 106 |
+
valid_window_size = num_points / num_windows
|
| 107 |
+
to_ordered = torch.argsort(code[s.start:s.stop])
|
| 108 |
+
if num_windows == 1:
|
| 109 |
+
fwd_indices.append(to_ordered)
|
| 110 |
+
bwd_indices.append(torch.zeros_like(to_ordered).scatter_(0, to_ordered, torch.arange(num_points, device=tensor.device)))
|
| 111 |
+
fwd_indices[-1] += s.start
|
| 112 |
+
bwd_indices[-1] += offsets[-1]
|
| 113 |
+
seq_lens.append(num_points)
|
| 114 |
+
seq_batch_indices.append(bi)
|
| 115 |
+
offsets.append(offsets[-1] + seq_lens[-1])
|
| 116 |
+
else:
|
| 117 |
+
# Partition the input
|
| 118 |
+
offset = 0
|
| 119 |
+
mids = [(i + 0.5) * valid_window_size + shift_sequence for i in range(num_windows)]
|
| 120 |
+
split = [math.floor(i * valid_window_size + shift_sequence) for i in range(num_windows + 1)]
|
| 121 |
+
bwd_index = torch.zeros((num_points,), dtype=torch.int64, device=tensor.device)
|
| 122 |
+
for i in range(num_windows):
|
| 123 |
+
mid = mids[i]
|
| 124 |
+
valid_start = split[i]
|
| 125 |
+
valid_end = split[i + 1]
|
| 126 |
+
padded_start = math.floor(mid - 0.5 * window_size)
|
| 127 |
+
padded_end = padded_start + window_size
|
| 128 |
+
fwd_indices.append(to_ordered[torch.arange(padded_start, padded_end, device=tensor.device) % num_points])
|
| 129 |
+
offset += valid_start - padded_start
|
| 130 |
+
bwd_index.scatter_(0, fwd_indices[-1][valid_start-padded_start:valid_end-padded_start], torch.arange(offset, offset + valid_end - valid_start, device=tensor.device))
|
| 131 |
+
offset += padded_end - valid_start
|
| 132 |
+
fwd_indices[-1] += s.start
|
| 133 |
+
seq_lens.extend([window_size] * num_windows)
|
| 134 |
+
seq_batch_indices.extend([bi] * num_windows)
|
| 135 |
+
bwd_indices.append(bwd_index + offsets[-1])
|
| 136 |
+
offsets.append(offsets[-1] + num_windows * window_size)
|
| 137 |
+
|
| 138 |
+
fwd_indices = torch.cat(fwd_indices)
|
| 139 |
+
bwd_indices = torch.cat(bwd_indices)
|
| 140 |
+
|
| 141 |
+
return fwd_indices, bwd_indices, seq_lens, seq_batch_indices
|
| 142 |
+
|
| 143 |
+
|
| 144 |
+
def sparse_serialized_scaled_dot_product_self_attention(
|
| 145 |
+
qkv: SparseTensor,
|
| 146 |
+
window_size: int,
|
| 147 |
+
serialize_mode: SerializeMode = SerializeMode.Z_ORDER,
|
| 148 |
+
shift_sequence: int = 0,
|
| 149 |
+
shift_window: Tuple[int, int, int] = (0, 0, 0)
|
| 150 |
+
) -> SparseTensor:
|
| 151 |
+
"""
|
| 152 |
+
Apply serialized scaled dot product self attention to a sparse tensor.
|
| 153 |
+
|
| 154 |
+
Args:
|
| 155 |
+
qkv (SparseTensor): [N, *, 3, H, C] sparse tensor containing Qs, Ks, and Vs.
|
| 156 |
+
window_size (int): The window size to use.
|
| 157 |
+
serialize_mode (SerializeMode): The serialization mode to use.
|
| 158 |
+
shift_sequence (int): The shift of serialized sequence.
|
| 159 |
+
shift_window (Tuple[int, int, int]): The shift of serialized coordinates.
|
| 160 |
+
shift (int): The shift to use.
|
| 161 |
+
"""
|
| 162 |
+
assert len(qkv.shape) == 4 and qkv.shape[1] == 3, f"Invalid shape for qkv, got {qkv.shape}, expected [N, *, 3, H, C]"
|
| 163 |
+
|
| 164 |
+
serialization_spatial_cache_name = f'serialization_{serialize_mode}_{window_size}_{shift_sequence}_{shift_window}'
|
| 165 |
+
serialization_spatial_cache = qkv.get_spatial_cache(serialization_spatial_cache_name)
|
| 166 |
+
if serialization_spatial_cache is None:
|
| 167 |
+
fwd_indices, bwd_indices, seq_lens, seq_batch_indices = calc_serialization(qkv, window_size, serialize_mode, shift_sequence, shift_window)
|
| 168 |
+
qkv.register_spatial_cache(serialization_spatial_cache_name, (fwd_indices, bwd_indices, seq_lens, seq_batch_indices))
|
| 169 |
+
else:
|
| 170 |
+
fwd_indices, bwd_indices, seq_lens, seq_batch_indices = serialization_spatial_cache
|
| 171 |
+
|
| 172 |
+
M = fwd_indices.shape[0]
|
| 173 |
+
T = qkv.feats.shape[0]
|
| 174 |
+
H = qkv.feats.shape[2]
|
| 175 |
+
C = qkv.feats.shape[3]
|
| 176 |
+
|
| 177 |
+
qkv_feats = qkv.feats[fwd_indices] # [M, 3, H, C]
|
| 178 |
+
|
| 179 |
+
if DEBUG:
|
| 180 |
+
start = 0
|
| 181 |
+
qkv_coords = qkv.coords[fwd_indices]
|
| 182 |
+
for i in range(len(seq_lens)):
|
| 183 |
+
assert (qkv_coords[start:start+seq_lens[i], 0] == seq_batch_indices[i]).all(), f"SparseWindowedScaledDotProductSelfAttention: batch index mismatch"
|
| 184 |
+
start += seq_lens[i]
|
| 185 |
+
|
| 186 |
+
if all([seq_len == window_size for seq_len in seq_lens]):
|
| 187 |
+
B = len(seq_lens)
|
| 188 |
+
N = window_size
|
| 189 |
+
qkv_feats = qkv_feats.reshape(B, N, 3, H, C)
|
| 190 |
+
if ATTN == 'xformers':
|
| 191 |
+
q, k, v = qkv_feats.unbind(dim=2) # [B, N, H, C]
|
| 192 |
+
out = xops.memory_efficient_attention(q, k, v) # [B, N, H, C]
|
| 193 |
+
elif ATTN == 'flash_attn':
|
| 194 |
+
out = flash_attn.flash_attn_qkvpacked_func(qkv_feats) # [B, N, H, C]
|
| 195 |
+
else:
|
| 196 |
+
raise ValueError(f"Unknown attention module: {ATTN}")
|
| 197 |
+
out = out.reshape(B * N, H, C) # [M, H, C]
|
| 198 |
+
else:
|
| 199 |
+
if ATTN == 'xformers':
|
| 200 |
+
q, k, v = qkv_feats.unbind(dim=1) # [M, H, C]
|
| 201 |
+
q = q.unsqueeze(0) # [1, M, H, C]
|
| 202 |
+
k = k.unsqueeze(0) # [1, M, H, C]
|
| 203 |
+
v = v.unsqueeze(0) # [1, M, H, C]
|
| 204 |
+
mask = xops.fmha.BlockDiagonalMask.from_seqlens(seq_lens)
|
| 205 |
+
out = xops.memory_efficient_attention(q, k, v, mask)[0] # [M, H, C]
|
| 206 |
+
elif ATTN == 'flash_attn':
|
| 207 |
+
cu_seqlens = torch.cat([torch.tensor([0]), torch.cumsum(torch.tensor(seq_lens), dim=0)], dim=0) \
|
| 208 |
+
.to(qkv.device).int()
|
| 209 |
+
out = flash_attn.flash_attn_varlen_qkvpacked_func(qkv_feats, cu_seqlens, max(seq_lens)) # [M, H, C]
|
| 210 |
+
|
| 211 |
+
out = out[bwd_indices] # [T, H, C]
|
| 212 |
+
|
| 213 |
+
if DEBUG:
|
| 214 |
+
qkv_coords = qkv_coords[bwd_indices]
|
| 215 |
+
assert torch.equal(qkv_coords, qkv.coords), "SparseWindowedScaledDotProductSelfAttention: coordinate mismatch"
|
| 216 |
+
|
| 217 |
+
return qkv.replace(out)
|
modules/sparse/attention/windowed_attn.py
ADDED
|
@@ -0,0 +1,158 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# MIT License
|
| 2 |
+
|
| 3 |
+
# Copyright (c) Microsoft Corporation.
|
| 4 |
+
# Copyright (c) 2025 VAST-AI-Research and contributors.
|
| 5 |
+
|
| 6 |
+
# Permission is hereby granted, free of charge, to any person obtaining a copy
|
| 7 |
+
# of this software and associated documentation files (the "Software"), to deal
|
| 8 |
+
# in the Software without restriction, including without limitation the rights
|
| 9 |
+
# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
|
| 10 |
+
# copies of the Software, and to permit persons to whom the Software is
|
| 11 |
+
# furnished to do so, subject to the following conditions:
|
| 12 |
+
|
| 13 |
+
# The above copyright notice and this permission notice shall be included in all
|
| 14 |
+
# copies or substantial portions of the Software.
|
| 15 |
+
|
| 16 |
+
# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
|
| 17 |
+
# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
|
| 18 |
+
# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
|
| 19 |
+
# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
|
| 20 |
+
# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
|
| 21 |
+
# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
|
| 22 |
+
# SOFTWARE
|
| 23 |
+
|
| 24 |
+
from typing import *
|
| 25 |
+
import torch
|
| 26 |
+
import math
|
| 27 |
+
from .. import SparseTensor
|
| 28 |
+
from .. import DEBUG, ATTN
|
| 29 |
+
|
| 30 |
+
if ATTN == 'xformers':
|
| 31 |
+
import xformers.ops as xops
|
| 32 |
+
elif ATTN == 'flash_attn':
|
| 33 |
+
import flash_attn
|
| 34 |
+
else:
|
| 35 |
+
raise ValueError(f"Unknown attention module: {ATTN}")
|
| 36 |
+
|
| 37 |
+
|
| 38 |
+
__all__ = [
|
| 39 |
+
'sparse_windowed_scaled_dot_product_self_attention',
|
| 40 |
+
]
|
| 41 |
+
|
| 42 |
+
|
| 43 |
+
def calc_window_partition(
|
| 44 |
+
tensor: SparseTensor,
|
| 45 |
+
window_size: Union[int, Tuple[int, ...]],
|
| 46 |
+
shift_window: Union[int, Tuple[int, ...]] = 0
|
| 47 |
+
) -> Tuple[torch.Tensor, torch.Tensor, List[int], List[int]]:
|
| 48 |
+
"""
|
| 49 |
+
Calculate serialization and partitioning for a set of coordinates.
|
| 50 |
+
|
| 51 |
+
Args:
|
| 52 |
+
tensor (SparseTensor): The input tensor.
|
| 53 |
+
window_size (int): The window size to use.
|
| 54 |
+
shift_window (Tuple[int, ...]): The shift of serialized coordinates.
|
| 55 |
+
|
| 56 |
+
Returns:
|
| 57 |
+
(torch.Tensor): Forwards indices.
|
| 58 |
+
(torch.Tensor): Backwards indices.
|
| 59 |
+
(List[int]): Sequence lengths.
|
| 60 |
+
(List[int]): Sequence batch indices.
|
| 61 |
+
"""
|
| 62 |
+
DIM = tensor.coords.shape[1] - 1
|
| 63 |
+
shift_window = (shift_window,) * DIM if isinstance(shift_window, int) else shift_window
|
| 64 |
+
window_size = (window_size,) * DIM if isinstance(window_size, int) else window_size
|
| 65 |
+
shifted_coords = tensor.coords.clone().detach()
|
| 66 |
+
shifted_coords[:, 1:] += torch.tensor(shift_window, device=tensor.device, dtype=torch.int32).unsqueeze(0)
|
| 67 |
+
|
| 68 |
+
MAX_COORDS = shifted_coords[:, 1:].max(dim=0).values.tolist()
|
| 69 |
+
NUM_WINDOWS = [math.ceil((mc + 1) / ws) for mc, ws in zip(MAX_COORDS, window_size)]
|
| 70 |
+
OFFSET = torch.cumprod(torch.tensor([1] + NUM_WINDOWS[::-1]), dim=0).tolist()[::-1]
|
| 71 |
+
|
| 72 |
+
shifted_coords[:, 1:] //= torch.tensor(window_size, device=tensor.device, dtype=torch.int32).unsqueeze(0)
|
| 73 |
+
shifted_indices = (shifted_coords * torch.tensor(OFFSET, device=tensor.device, dtype=torch.int32).unsqueeze(0)).sum(dim=1)
|
| 74 |
+
fwd_indices = torch.argsort(shifted_indices)
|
| 75 |
+
bwd_indices = torch.empty_like(fwd_indices)
|
| 76 |
+
bwd_indices[fwd_indices] = torch.arange(fwd_indices.shape[0], device=tensor.device)
|
| 77 |
+
seq_lens = torch.bincount(shifted_indices)
|
| 78 |
+
seq_batch_indices = torch.arange(seq_lens.shape[0], device=tensor.device, dtype=torch.int32) // OFFSET[0]
|
| 79 |
+
mask = seq_lens != 0
|
| 80 |
+
seq_lens = seq_lens[mask].tolist()
|
| 81 |
+
seq_batch_indices = seq_batch_indices[mask].tolist()
|
| 82 |
+
|
| 83 |
+
return fwd_indices, bwd_indices, seq_lens, seq_batch_indices
|
| 84 |
+
|
| 85 |
+
|
| 86 |
+
def sparse_windowed_scaled_dot_product_self_attention(
|
| 87 |
+
qkv: SparseTensor,
|
| 88 |
+
window_size: int,
|
| 89 |
+
shift_window: Tuple[int, int, int] = (0, 0, 0)
|
| 90 |
+
) -> SparseTensor:
|
| 91 |
+
"""
|
| 92 |
+
Apply windowed scaled dot product self attention to a sparse tensor.
|
| 93 |
+
|
| 94 |
+
Args:
|
| 95 |
+
qkv (SparseTensor): [N, *, 3, H, C] sparse tensor containing Qs, Ks, and Vs.
|
| 96 |
+
window_size (int): The window size to use.
|
| 97 |
+
shift_window (Tuple[int, int, int]): The shift of serialized coordinates.
|
| 98 |
+
shift (int): The shift to use.
|
| 99 |
+
"""
|
| 100 |
+
assert len(qkv.shape) == 4 and qkv.shape[1] == 3, f"Invalid shape for qkv, got {qkv.shape}, expected [N, *, 3, H, C]"
|
| 101 |
+
|
| 102 |
+
serialization_spatial_cache_name = f'window_partition_{window_size}_{shift_window}_{qkv.feats.shape[0]}'
|
| 103 |
+
serialization_spatial_cache = qkv.get_spatial_cache(serialization_spatial_cache_name)
|
| 104 |
+
if serialization_spatial_cache is None:
|
| 105 |
+
fwd_indices, bwd_indices, seq_lens, seq_batch_indices = calc_window_partition(qkv, window_size, shift_window)
|
| 106 |
+
qkv.register_spatial_cache(serialization_spatial_cache_name, (fwd_indices, bwd_indices, seq_lens, seq_batch_indices))
|
| 107 |
+
else:
|
| 108 |
+
fwd_indices, bwd_indices, seq_lens, seq_batch_indices = serialization_spatial_cache
|
| 109 |
+
|
| 110 |
+
M = fwd_indices.shape[0]
|
| 111 |
+
T = qkv.feats.shape[0]
|
| 112 |
+
H = qkv.feats.shape[2]
|
| 113 |
+
C = qkv.feats.shape[3]
|
| 114 |
+
|
| 115 |
+
qkv_feats = qkv.feats[fwd_indices] # [M, 3, H, C]
|
| 116 |
+
|
| 117 |
+
if DEBUG:
|
| 118 |
+
start = 0
|
| 119 |
+
qkv_coords = qkv.coords[fwd_indices]
|
| 120 |
+
for i in range(len(seq_lens)):
|
| 121 |
+
seq_coords = qkv_coords[start:start+seq_lens[i]]
|
| 122 |
+
assert (seq_coords[:, 0] == seq_batch_indices[i]).all(), f"SparseWindowedScaledDotProductSelfAttention: batch index mismatch"
|
| 123 |
+
assert (seq_coords[:, 1:].max(dim=0).values - seq_coords[:, 1:].min(dim=0).values < window_size).all(), \
|
| 124 |
+
f"SparseWindowedScaledDotProductSelfAttention: window size exceeded"
|
| 125 |
+
start += seq_lens[i]
|
| 126 |
+
|
| 127 |
+
if all([seq_len == window_size for seq_len in seq_lens]):
|
| 128 |
+
B = len(seq_lens)
|
| 129 |
+
N = window_size
|
| 130 |
+
qkv_feats = qkv_feats.reshape(B, N, 3, H, C)
|
| 131 |
+
if ATTN == 'xformers':
|
| 132 |
+
q, k, v = qkv_feats.unbind(dim=2) # [B, N, H, C]
|
| 133 |
+
out = xops.memory_efficient_attention(q, k, v) # [B, N, H, C]
|
| 134 |
+
elif ATTN == 'flash_attn':
|
| 135 |
+
out = flash_attn.flash_attn_qkvpacked_func(qkv_feats) # [B, N, H, C]
|
| 136 |
+
else:
|
| 137 |
+
raise ValueError(f"Unknown attention module: {ATTN}")
|
| 138 |
+
out = out.reshape(B * N, H, C) # [M, H, C]
|
| 139 |
+
else:
|
| 140 |
+
if ATTN == 'xformers':
|
| 141 |
+
q, k, v = qkv_feats.unbind(dim=1) # [M, H, C]
|
| 142 |
+
q = q.unsqueeze(0) # [1, M, H, C]
|
| 143 |
+
k = k.unsqueeze(0) # [1, M, H, C]
|
| 144 |
+
v = v.unsqueeze(0) # [1, M, H, C]
|
| 145 |
+
mask = xops.fmha.BlockDiagonalMask.from_seqlens(seq_lens)
|
| 146 |
+
out = xops.memory_efficient_attention(q, k, v, mask)[0] # [M, H, C]
|
| 147 |
+
elif ATTN == 'flash_attn':
|
| 148 |
+
cu_seqlens = torch.cat([torch.tensor([0]), torch.cumsum(torch.tensor(seq_lens), dim=0)], dim=0) \
|
| 149 |
+
.to(qkv.device).int()
|
| 150 |
+
out = flash_attn.flash_attn_varlen_qkvpacked_func(qkv_feats, cu_seqlens, max(seq_lens)) # [M, H, C]
|
| 151 |
+
|
| 152 |
+
out = out[bwd_indices] # [T, H, C]
|
| 153 |
+
|
| 154 |
+
if DEBUG:
|
| 155 |
+
qkv_coords = qkv_coords[bwd_indices]
|
| 156 |
+
assert torch.equal(qkv_coords, qkv.coords), "SparseWindowedScaledDotProductSelfAttention: coordinate mismatch"
|
| 157 |
+
|
| 158 |
+
return qkv.replace(out)
|
modules/sparse/basic.py
ADDED
|
@@ -0,0 +1,482 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# MIT License
|
| 2 |
+
|
| 3 |
+
# Copyright (c) Microsoft Corporation.
|
| 4 |
+
# Copyright (c) 2025 VAST-AI-Research and contributors.
|
| 5 |
+
|
| 6 |
+
# Permission is hereby granted, free of charge, to any person obtaining a copy
|
| 7 |
+
# of this software and associated documentation files (the "Software"), to deal
|
| 8 |
+
# in the Software without restriction, including without limitation the rights
|
| 9 |
+
# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
|
| 10 |
+
# copies of the Software, and to permit persons to whom the Software is
|
| 11 |
+
# furnished to do so, subject to the following conditions:
|
| 12 |
+
|
| 13 |
+
# The above copyright notice and this permission notice shall be included in all
|
| 14 |
+
# copies or substantial portions of the Software.
|
| 15 |
+
|
| 16 |
+
# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
|
| 17 |
+
# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
|
| 18 |
+
# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
|
| 19 |
+
# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
|
| 20 |
+
# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
|
| 21 |
+
# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
|
| 22 |
+
# SOFTWARE
|
| 23 |
+
|
| 24 |
+
from typing import *
|
| 25 |
+
import torch
|
| 26 |
+
import torch.nn as nn
|
| 27 |
+
from . import BACKEND, DEBUG
|
| 28 |
+
SparseTensorData = None # Lazy import
|
| 29 |
+
|
| 30 |
+
|
| 31 |
+
__all__ = [
|
| 32 |
+
'SparseTensor',
|
| 33 |
+
'sparse_batch_broadcast',
|
| 34 |
+
'sparse_batch_op',
|
| 35 |
+
'sparse_cat',
|
| 36 |
+
'sparse_unbind',
|
| 37 |
+
]
|
| 38 |
+
|
| 39 |
+
|
| 40 |
+
class SparseTensor:
|
| 41 |
+
"""
|
| 42 |
+
Sparse tensor with support for both torchsparse and spconv backends.
|
| 43 |
+
|
| 44 |
+
Parameters:
|
| 45 |
+
- feats (torch.Tensor): Features of the sparse tensor.
|
| 46 |
+
- coords (torch.Tensor): Coordinates of the sparse tensor.
|
| 47 |
+
- shape (torch.Size): Shape of the sparse tensor.
|
| 48 |
+
- layout (List[slice]): Layout of the sparse tensor for each batch
|
| 49 |
+
- data (SparseTensorData): Sparse tensor data used for convolusion
|
| 50 |
+
|
| 51 |
+
NOTE:
|
| 52 |
+
- Data corresponding to a same batch should be contiguous.
|
| 53 |
+
- Coords should be in [0, 1023]
|
| 54 |
+
"""
|
| 55 |
+
@overload
|
| 56 |
+
def __init__(self, feats: torch.Tensor, coords: torch.Tensor, shape: Optional[torch.Size] = None, layout: Optional[List[slice]] = None, **kwargs): ...
|
| 57 |
+
|
| 58 |
+
@overload
|
| 59 |
+
def __init__(self, data, shape: Optional[torch.Size] = None, layout: Optional[List[slice]] = None, **kwargs): ...
|
| 60 |
+
|
| 61 |
+
def __init__(self, *args, **kwargs):
|
| 62 |
+
# Lazy import of sparse tensor backend
|
| 63 |
+
global SparseTensorData
|
| 64 |
+
if SparseTensorData is None:
|
| 65 |
+
import importlib
|
| 66 |
+
if BACKEND == 'torchsparse':
|
| 67 |
+
SparseTensorData = importlib.import_module('torchsparse').SparseTensor
|
| 68 |
+
elif BACKEND == 'spconv':
|
| 69 |
+
SparseTensorData = importlib.import_module('spconv.pytorch').SparseConvTensor
|
| 70 |
+
|
| 71 |
+
method_id = 0
|
| 72 |
+
if len(args) != 0:
|
| 73 |
+
method_id = 0 if isinstance(args[0], torch.Tensor) else 1
|
| 74 |
+
else:
|
| 75 |
+
method_id = 1 if 'data' in kwargs else 0
|
| 76 |
+
|
| 77 |
+
if method_id == 0:
|
| 78 |
+
feats, coords, shape, layout = args + (None,) * (4 - len(args))
|
| 79 |
+
if 'feats' in kwargs:
|
| 80 |
+
feats = kwargs['feats']
|
| 81 |
+
del kwargs['feats']
|
| 82 |
+
if 'coords' in kwargs:
|
| 83 |
+
coords = kwargs['coords']
|
| 84 |
+
del kwargs['coords']
|
| 85 |
+
if 'shape' in kwargs:
|
| 86 |
+
shape = kwargs['shape']
|
| 87 |
+
del kwargs['shape']
|
| 88 |
+
if 'layout' in kwargs:
|
| 89 |
+
layout = kwargs['layout']
|
| 90 |
+
del kwargs['layout']
|
| 91 |
+
|
| 92 |
+
if shape is None:
|
| 93 |
+
shape = self.__cal_shape(feats, coords)
|
| 94 |
+
if layout is None:
|
| 95 |
+
layout = self.__cal_layout(coords, shape[0])
|
| 96 |
+
if BACKEND == 'torchsparse':
|
| 97 |
+
self.data = SparseTensorData(feats, coords, **kwargs)
|
| 98 |
+
elif BACKEND == 'spconv':
|
| 99 |
+
spatial_shape = list(coords.max(0)[0] + 1)[1:]
|
| 100 |
+
self.data = SparseTensorData(feats.reshape(feats.shape[0], -1), coords, spatial_shape, shape[0], **kwargs)
|
| 101 |
+
self.data._features = feats
|
| 102 |
+
elif method_id == 1:
|
| 103 |
+
data, shape, layout = args + (None,) * (3 - len(args))
|
| 104 |
+
if 'data' in kwargs:
|
| 105 |
+
data = kwargs['data']
|
| 106 |
+
del kwargs['data']
|
| 107 |
+
if 'shape' in kwargs:
|
| 108 |
+
shape = kwargs['shape']
|
| 109 |
+
del kwargs['shape']
|
| 110 |
+
if 'layout' in kwargs:
|
| 111 |
+
layout = kwargs['layout']
|
| 112 |
+
del kwargs['layout']
|
| 113 |
+
|
| 114 |
+
self.data = data
|
| 115 |
+
if shape is None:
|
| 116 |
+
shape = self.__cal_shape(self.feats, self.coords)
|
| 117 |
+
if layout is None:
|
| 118 |
+
layout = self.__cal_layout(self.coords, shape[0])
|
| 119 |
+
|
| 120 |
+
self._shape = shape
|
| 121 |
+
self._layout = layout
|
| 122 |
+
self._scale = kwargs.get('scale', (1, 1, 1))
|
| 123 |
+
self._spatial_cache = kwargs.get('spatial_cache', {})
|
| 124 |
+
|
| 125 |
+
if DEBUG:
|
| 126 |
+
try:
|
| 127 |
+
assert self.feats.shape[0] == self.coords.shape[0], f"Invalid feats shape: {self.feats.shape}, coords shape: {self.coords.shape}"
|
| 128 |
+
assert self.shape == self.__cal_shape(self.feats, self.coords), f"Invalid shape: {self.shape}"
|
| 129 |
+
assert self.layout == self.__cal_layout(self.coords, self.shape[0]), f"Invalid layout: {self.layout}"
|
| 130 |
+
for i in range(self.shape[0]):
|
| 131 |
+
assert torch.all(self.coords[self.layout[i], 0] == i), f"The data of batch {i} is not contiguous"
|
| 132 |
+
except Exception as e:
|
| 133 |
+
print('Debugging information:')
|
| 134 |
+
print(f"- Shape: {self.shape}")
|
| 135 |
+
print(f"- Layout: {self.layout}")
|
| 136 |
+
print(f"- Scale: {self._scale}")
|
| 137 |
+
print(f"- Coords: {self.coords}")
|
| 138 |
+
raise e
|
| 139 |
+
|
| 140 |
+
def __cal_shape(self, feats, coords):
|
| 141 |
+
shape = []
|
| 142 |
+
shape.append(coords[:, 0].max().item() + 1)
|
| 143 |
+
shape.extend([*feats.shape[1:]])
|
| 144 |
+
return torch.Size(shape)
|
| 145 |
+
|
| 146 |
+
def __cal_layout(self, coords, batch_size):
|
| 147 |
+
seq_len = torch.bincount(coords[:, 0], minlength=batch_size)
|
| 148 |
+
offset = torch.cumsum(seq_len, dim=0)
|
| 149 |
+
layout = [slice((offset[i] - seq_len[i]).item(), offset[i].item()) for i in range(batch_size)]
|
| 150 |
+
return layout
|
| 151 |
+
|
| 152 |
+
@property
|
| 153 |
+
def shape(self) -> torch.Size:
|
| 154 |
+
return self._shape
|
| 155 |
+
|
| 156 |
+
def dim(self) -> int:
|
| 157 |
+
return len(self.shape)
|
| 158 |
+
|
| 159 |
+
@property
|
| 160 |
+
def layout(self) -> List[slice]:
|
| 161 |
+
return self._layout
|
| 162 |
+
|
| 163 |
+
@property
|
| 164 |
+
def feats(self) -> torch.Tensor:
|
| 165 |
+
if BACKEND == 'torchsparse':
|
| 166 |
+
return self.data.F
|
| 167 |
+
elif BACKEND == 'spconv':
|
| 168 |
+
return self.data.features
|
| 169 |
+
|
| 170 |
+
@feats.setter
|
| 171 |
+
def feats(self, value: torch.Tensor):
|
| 172 |
+
if BACKEND == 'torchsparse':
|
| 173 |
+
self.data.F = value
|
| 174 |
+
elif BACKEND == 'spconv':
|
| 175 |
+
self.data.features = value
|
| 176 |
+
|
| 177 |
+
@property
|
| 178 |
+
def coords(self) -> torch.Tensor:
|
| 179 |
+
if BACKEND == 'torchsparse':
|
| 180 |
+
return self.data.C
|
| 181 |
+
elif BACKEND == 'spconv':
|
| 182 |
+
return self.data.indices
|
| 183 |
+
|
| 184 |
+
@coords.setter
|
| 185 |
+
def coords(self, value: torch.Tensor):
|
| 186 |
+
if BACKEND == 'torchsparse':
|
| 187 |
+
self.data.C = value
|
| 188 |
+
elif BACKEND == 'spconv':
|
| 189 |
+
self.data.indices = value
|
| 190 |
+
|
| 191 |
+
@property
|
| 192 |
+
def dtype(self):
|
| 193 |
+
return self.feats.dtype
|
| 194 |
+
|
| 195 |
+
@property
|
| 196 |
+
def device(self):
|
| 197 |
+
return self.feats.device
|
| 198 |
+
|
| 199 |
+
@overload
|
| 200 |
+
def to(self, dtype: torch.dtype) -> 'SparseTensor': ...
|
| 201 |
+
|
| 202 |
+
@overload
|
| 203 |
+
def to(self, device: Optional[Union[str, torch.device]] = None, dtype: Optional[torch.dtype] = None) -> 'SparseTensor': ...
|
| 204 |
+
|
| 205 |
+
def to(self, *args, **kwargs) -> 'SparseTensor':
|
| 206 |
+
device = None
|
| 207 |
+
dtype = None
|
| 208 |
+
if len(args) == 2:
|
| 209 |
+
device, dtype = args
|
| 210 |
+
elif len(args) == 1:
|
| 211 |
+
if isinstance(args[0], torch.dtype):
|
| 212 |
+
dtype = args[0]
|
| 213 |
+
else:
|
| 214 |
+
device = args[0]
|
| 215 |
+
if 'dtype' in kwargs:
|
| 216 |
+
assert dtype is None, "to() received multiple values for argument 'dtype'"
|
| 217 |
+
dtype = kwargs['dtype']
|
| 218 |
+
if 'device' in kwargs:
|
| 219 |
+
assert device is None, "to() received multiple values for argument 'device'"
|
| 220 |
+
device = kwargs['device']
|
| 221 |
+
|
| 222 |
+
new_feats = self.feats.to(device=device, dtype=dtype)
|
| 223 |
+
new_coords = self.coords.to(device=device)
|
| 224 |
+
return self.replace(new_feats, new_coords)
|
| 225 |
+
|
| 226 |
+
def type(self, dtype):
|
| 227 |
+
new_feats = self.feats.type(dtype)
|
| 228 |
+
return self.replace(new_feats)
|
| 229 |
+
|
| 230 |
+
def cpu(self) -> 'SparseTensor':
|
| 231 |
+
new_feats = self.feats.cpu()
|
| 232 |
+
new_coords = self.coords.cpu()
|
| 233 |
+
return self.replace(new_feats, new_coords)
|
| 234 |
+
|
| 235 |
+
def cuda(self) -> 'SparseTensor':
|
| 236 |
+
new_feats = self.feats.cuda()
|
| 237 |
+
new_coords = self.coords.cuda()
|
| 238 |
+
return self.replace(new_feats, new_coords)
|
| 239 |
+
|
| 240 |
+
def half(self) -> 'SparseTensor':
|
| 241 |
+
new_feats = self.feats.half()
|
| 242 |
+
return self.replace(new_feats)
|
| 243 |
+
|
| 244 |
+
def float(self) -> 'SparseTensor':
|
| 245 |
+
new_feats = self.feats.float()
|
| 246 |
+
return self.replace(new_feats)
|
| 247 |
+
|
| 248 |
+
def detach(self) -> 'SparseTensor':
|
| 249 |
+
new_coords = self.coords.detach()
|
| 250 |
+
new_feats = self.feats.detach()
|
| 251 |
+
return self.replace(new_feats, new_coords)
|
| 252 |
+
|
| 253 |
+
def dense(self) -> torch.Tensor:
|
| 254 |
+
if BACKEND == 'torchsparse':
|
| 255 |
+
return self.data.dense()
|
| 256 |
+
elif BACKEND == 'spconv':
|
| 257 |
+
return self.data.dense()
|
| 258 |
+
|
| 259 |
+
def reshape(self, *shape) -> 'SparseTensor':
|
| 260 |
+
new_feats = self.feats.reshape(self.feats.shape[0], *shape)
|
| 261 |
+
return self.replace(new_feats)
|
| 262 |
+
|
| 263 |
+
def unbind(self, dim: int) -> List['SparseTensor']:
|
| 264 |
+
return sparse_unbind(self, dim)
|
| 265 |
+
|
| 266 |
+
def replace(self, feats: torch.Tensor, coords: Optional[torch.Tensor] = None) -> 'SparseTensor':
|
| 267 |
+
new_shape = [self.shape[0]]
|
| 268 |
+
new_shape.extend(feats.shape[1:])
|
| 269 |
+
if BACKEND == 'torchsparse':
|
| 270 |
+
new_data = SparseTensorData(
|
| 271 |
+
feats=feats,
|
| 272 |
+
coords=self.data.coords if coords is None else coords,
|
| 273 |
+
stride=self.data.stride,
|
| 274 |
+
spatial_range=self.data.spatial_range,
|
| 275 |
+
)
|
| 276 |
+
new_data._caches = self.data._caches
|
| 277 |
+
elif BACKEND == 'spconv':
|
| 278 |
+
new_data = SparseTensorData(
|
| 279 |
+
self.data.features.reshape(self.data.features.shape[0], -1),
|
| 280 |
+
self.data.indices,
|
| 281 |
+
self.data.spatial_shape,
|
| 282 |
+
self.data.batch_size,
|
| 283 |
+
self.data.grid,
|
| 284 |
+
self.data.voxel_num,
|
| 285 |
+
self.data.indice_dict
|
| 286 |
+
)
|
| 287 |
+
new_data._features = feats
|
| 288 |
+
new_data.benchmark = self.data.benchmark
|
| 289 |
+
new_data.benchmark_record = self.data.benchmark_record
|
| 290 |
+
new_data.thrust_allocator = self.data.thrust_allocator
|
| 291 |
+
new_data._timer = self.data._timer
|
| 292 |
+
new_data.force_algo = self.data.force_algo
|
| 293 |
+
new_data.int8_scale = self.data.int8_scale
|
| 294 |
+
if coords is not None:
|
| 295 |
+
new_data.indices = coords
|
| 296 |
+
new_tensor = SparseTensor(new_data, shape=torch.Size(new_shape), layout=self.layout, scale=self._scale, spatial_cache=self._spatial_cache)
|
| 297 |
+
return new_tensor
|
| 298 |
+
|
| 299 |
+
@staticmethod
|
| 300 |
+
def full(aabb, dim, value, dtype=torch.float32, device=None) -> 'SparseTensor':
|
| 301 |
+
N, C = dim
|
| 302 |
+
x = torch.arange(aabb[0], aabb[3] + 1)
|
| 303 |
+
y = torch.arange(aabb[1], aabb[4] + 1)
|
| 304 |
+
z = torch.arange(aabb[2], aabb[5] + 1)
|
| 305 |
+
coords = torch.stack(torch.meshgrid(x, y, z, indexing='ij'), dim=-1).reshape(-1, 3)
|
| 306 |
+
coords = torch.cat([
|
| 307 |
+
torch.arange(N).view(-1, 1).repeat(1, coords.shape[0]).view(-1, 1),
|
| 308 |
+
coords.repeat(N, 1),
|
| 309 |
+
], dim=1).to(dtype=torch.int32, device=device)
|
| 310 |
+
feats = torch.full((coords.shape[0], C), value, dtype=dtype, device=device)
|
| 311 |
+
return SparseTensor(feats=feats, coords=coords)
|
| 312 |
+
|
| 313 |
+
def __merge_sparse_cache(self, other: 'SparseTensor') -> dict:
|
| 314 |
+
new_cache = {}
|
| 315 |
+
for k in set(list(self._spatial_cache.keys()) + list(other._spatial_cache.keys())):
|
| 316 |
+
if k in self._spatial_cache:
|
| 317 |
+
new_cache[k] = self._spatial_cache[k]
|
| 318 |
+
if k in other._spatial_cache:
|
| 319 |
+
if k not in new_cache:
|
| 320 |
+
new_cache[k] = other._spatial_cache[k]
|
| 321 |
+
else:
|
| 322 |
+
new_cache[k].update(other._spatial_cache[k])
|
| 323 |
+
return new_cache
|
| 324 |
+
|
| 325 |
+
def __neg__(self) -> 'SparseTensor':
|
| 326 |
+
return self.replace(-self.feats)
|
| 327 |
+
|
| 328 |
+
def __elemwise__(self, other: Union[torch.Tensor, 'SparseTensor'], op: callable) -> 'SparseTensor':
|
| 329 |
+
if isinstance(other, torch.Tensor):
|
| 330 |
+
try:
|
| 331 |
+
other = torch.broadcast_to(other, self.shape)
|
| 332 |
+
other = sparse_batch_broadcast(self, other)
|
| 333 |
+
except:
|
| 334 |
+
pass
|
| 335 |
+
if isinstance(other, SparseTensor):
|
| 336 |
+
other = other.feats
|
| 337 |
+
new_feats = op(self.feats, other)
|
| 338 |
+
new_tensor = self.replace(new_feats)
|
| 339 |
+
if isinstance(other, SparseTensor):
|
| 340 |
+
new_tensor._spatial_cache = self.__merge_sparse_cache(other)
|
| 341 |
+
return new_tensor
|
| 342 |
+
|
| 343 |
+
def __add__(self, other: Union[torch.Tensor, 'SparseTensor', float]) -> 'SparseTensor':
|
| 344 |
+
return self.__elemwise__(other, torch.add)
|
| 345 |
+
|
| 346 |
+
def __radd__(self, other: Union[torch.Tensor, 'SparseTensor', float]) -> 'SparseTensor':
|
| 347 |
+
return self.__elemwise__(other, torch.add)
|
| 348 |
+
|
| 349 |
+
def __sub__(self, other: Union[torch.Tensor, 'SparseTensor', float]) -> 'SparseTensor':
|
| 350 |
+
return self.__elemwise__(other, torch.sub)
|
| 351 |
+
|
| 352 |
+
def __rsub__(self, other: Union[torch.Tensor, 'SparseTensor', float]) -> 'SparseTensor':
|
| 353 |
+
return self.__elemwise__(other, lambda x, y: torch.sub(y, x))
|
| 354 |
+
|
| 355 |
+
def __mul__(self, other: Union[torch.Tensor, 'SparseTensor', float]) -> 'SparseTensor':
|
| 356 |
+
return self.__elemwise__(other, torch.mul)
|
| 357 |
+
|
| 358 |
+
def __rmul__(self, other: Union[torch.Tensor, 'SparseTensor', float]) -> 'SparseTensor':
|
| 359 |
+
return self.__elemwise__(other, torch.mul)
|
| 360 |
+
|
| 361 |
+
def __truediv__(self, other: Union[torch.Tensor, 'SparseTensor', float]) -> 'SparseTensor':
|
| 362 |
+
return self.__elemwise__(other, torch.div)
|
| 363 |
+
|
| 364 |
+
def __rtruediv__(self, other: Union[torch.Tensor, 'SparseTensor', float]) -> 'SparseTensor':
|
| 365 |
+
return self.__elemwise__(other, lambda x, y: torch.div(y, x))
|
| 366 |
+
|
| 367 |
+
def __getitem__(self, idx):
|
| 368 |
+
if isinstance(idx, int):
|
| 369 |
+
idx = [idx]
|
| 370 |
+
elif isinstance(idx, slice):
|
| 371 |
+
idx = range(*idx.indices(self.shape[0]))
|
| 372 |
+
elif isinstance(idx, torch.Tensor):
|
| 373 |
+
if idx.dtype == torch.bool:
|
| 374 |
+
assert idx.shape == (self.shape[0],), f"Invalid index shape: {idx.shape}"
|
| 375 |
+
idx = idx.nonzero().squeeze(1)
|
| 376 |
+
elif idx.dtype in [torch.int32, torch.int64]:
|
| 377 |
+
assert len(idx.shape) == 1, f"Invalid index shape: {idx.shape}"
|
| 378 |
+
else:
|
| 379 |
+
raise ValueError(f"Unknown index type: {idx.dtype}")
|
| 380 |
+
else:
|
| 381 |
+
raise ValueError(f"Unknown index type: {type(idx)}")
|
| 382 |
+
|
| 383 |
+
coords = []
|
| 384 |
+
feats = []
|
| 385 |
+
for new_idx, old_idx in enumerate(idx):
|
| 386 |
+
coords.append(self.coords[self.layout[old_idx]].clone())
|
| 387 |
+
coords[-1][:, 0] = new_idx
|
| 388 |
+
feats.append(self.feats[self.layout[old_idx]])
|
| 389 |
+
coords = torch.cat(coords, dim=0).contiguous()
|
| 390 |
+
feats = torch.cat(feats, dim=0).contiguous()
|
| 391 |
+
return SparseTensor(feats=feats, coords=coords)
|
| 392 |
+
|
| 393 |
+
def register_spatial_cache(self, key, value) -> None:
|
| 394 |
+
"""
|
| 395 |
+
Register a spatial cache.
|
| 396 |
+
The spatial cache can be any thing you want to cache.
|
| 397 |
+
The registery and retrieval of the cache is based on current scale.
|
| 398 |
+
"""
|
| 399 |
+
scale_key = str(self._scale)
|
| 400 |
+
if scale_key not in self._spatial_cache:
|
| 401 |
+
self._spatial_cache[scale_key] = {}
|
| 402 |
+
self._spatial_cache[scale_key][key] = value
|
| 403 |
+
|
| 404 |
+
def get_spatial_cache(self, key=None):
|
| 405 |
+
"""
|
| 406 |
+
Get a spatial cache.
|
| 407 |
+
"""
|
| 408 |
+
scale_key = str(self._scale)
|
| 409 |
+
cur_scale_cache = self._spatial_cache.get(scale_key, {})
|
| 410 |
+
if key is None:
|
| 411 |
+
return cur_scale_cache
|
| 412 |
+
return cur_scale_cache.get(key, None)
|
| 413 |
+
|
| 414 |
+
|
| 415 |
+
def sparse_batch_broadcast(input: SparseTensor, other: torch.Tensor) -> torch.Tensor:
|
| 416 |
+
"""
|
| 417 |
+
Broadcast a 1D tensor to a sparse tensor along the batch dimension then perform an operation.
|
| 418 |
+
|
| 419 |
+
Args:
|
| 420 |
+
input (torch.Tensor): 1D tensor to broadcast.
|
| 421 |
+
target (SparseTensor): Sparse tensor to broadcast to.
|
| 422 |
+
op (callable): Operation to perform after broadcasting. Defaults to torch.add.
|
| 423 |
+
"""
|
| 424 |
+
coords, feats = input.coords, input.feats
|
| 425 |
+
broadcasted = torch.zeros_like(feats)
|
| 426 |
+
for k in range(input.shape[0]):
|
| 427 |
+
broadcasted[input.layout[k]] = other[k]
|
| 428 |
+
return broadcasted
|
| 429 |
+
|
| 430 |
+
|
| 431 |
+
def sparse_batch_op(input: SparseTensor, other: torch.Tensor, op: callable = torch.add) -> SparseTensor:
|
| 432 |
+
"""
|
| 433 |
+
Broadcast a 1D tensor to a sparse tensor along the batch dimension then perform an operation.
|
| 434 |
+
|
| 435 |
+
Args:
|
| 436 |
+
input (torch.Tensor): 1D tensor to broadcast.
|
| 437 |
+
target (SparseTensor): Sparse tensor to broadcast to.
|
| 438 |
+
op (callable): Operation to perform after broadcasting. Defaults to torch.add.
|
| 439 |
+
"""
|
| 440 |
+
return input.replace(op(input.feats, sparse_batch_broadcast(input, other)))
|
| 441 |
+
|
| 442 |
+
|
| 443 |
+
def sparse_cat(inputs: List[SparseTensor], dim: int = 0) -> SparseTensor:
|
| 444 |
+
"""
|
| 445 |
+
Concatenate a list of sparse tensors.
|
| 446 |
+
|
| 447 |
+
Args:
|
| 448 |
+
inputs (List[SparseTensor]): List of sparse tensors to concatenate.
|
| 449 |
+
"""
|
| 450 |
+
if dim == 0:
|
| 451 |
+
start = 0
|
| 452 |
+
coords = []
|
| 453 |
+
for input in inputs:
|
| 454 |
+
coords.append(input.coords.clone())
|
| 455 |
+
coords[-1][:, 0] += start
|
| 456 |
+
start += input.shape[0]
|
| 457 |
+
coords = torch.cat(coords, dim=0)
|
| 458 |
+
feats = torch.cat([input.feats for input in inputs], dim=0)
|
| 459 |
+
output = SparseTensor(
|
| 460 |
+
coords=coords,
|
| 461 |
+
feats=feats,
|
| 462 |
+
)
|
| 463 |
+
else:
|
| 464 |
+
feats = torch.cat([input.feats for input in inputs], dim=dim)
|
| 465 |
+
output = inputs[0].replace(feats)
|
| 466 |
+
|
| 467 |
+
return output
|
| 468 |
+
|
| 469 |
+
|
| 470 |
+
def sparse_unbind(input: SparseTensor, dim: int) -> List[SparseTensor]:
|
| 471 |
+
"""
|
| 472 |
+
Unbind a sparse tensor along a dimension.
|
| 473 |
+
|
| 474 |
+
Args:
|
| 475 |
+
input (SparseTensor): Sparse tensor to unbind.
|
| 476 |
+
dim (int): Dimension to unbind.
|
| 477 |
+
"""
|
| 478 |
+
if dim == 0:
|
| 479 |
+
return [input[i] for i in range(input.shape[0])]
|
| 480 |
+
else:
|
| 481 |
+
feats = input.feats.unbind(dim)
|
| 482 |
+
return [input.replace(f) for f in feats]
|
modules/sparse/blocks.py
ADDED
|
@@ -0,0 +1,71 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from typing import *
|
| 2 |
+
import torch
|
| 3 |
+
import torch.nn as nn
|
| 4 |
+
import torch.nn.functional as F
|
| 5 |
+
from ..utils import zero_module
|
| 6 |
+
from ..norm import LayerNorm32
|
| 7 |
+
from .. import sparse as sp
|
| 8 |
+
|
| 9 |
+
|
| 10 |
+
class SparseResBlock3d(nn.Module):
|
| 11 |
+
def __init__(
|
| 12 |
+
self,
|
| 13 |
+
channels: int,
|
| 14 |
+
out_channels: Optional[int] = None,
|
| 15 |
+
downsample: bool = False,
|
| 16 |
+
upsample: bool = False,
|
| 17 |
+
use_checkpoint: bool = False,
|
| 18 |
+
):
|
| 19 |
+
super().__init__()
|
| 20 |
+
self.channels = channels
|
| 21 |
+
self.out_channels = out_channels or channels
|
| 22 |
+
self.downsample = downsample
|
| 23 |
+
self.upsample = upsample
|
| 24 |
+
self.use_checkpoint = use_checkpoint
|
| 25 |
+
|
| 26 |
+
assert not (
|
| 27 |
+
downsample and upsample
|
| 28 |
+
), "Cannot downsample and upsample at the same time"
|
| 29 |
+
|
| 30 |
+
self.norm1 = LayerNorm32(channels, elementwise_affine=True, eps=1e-6)
|
| 31 |
+
self.norm2 = LayerNorm32(self.out_channels, elementwise_affine=False, eps=1e-6)
|
| 32 |
+
self.conv1 = sp.SparseConv3d(channels, self.out_channels, 3)
|
| 33 |
+
self.conv2 = zero_module(
|
| 34 |
+
sp.SparseConv3d(self.out_channels, self.out_channels, 3)
|
| 35 |
+
)
|
| 36 |
+
|
| 37 |
+
self.skip_connection = (
|
| 38 |
+
sp.SparseLinear(channels, self.out_channels)
|
| 39 |
+
if channels != self.out_channels
|
| 40 |
+
else nn.Identity()
|
| 41 |
+
)
|
| 42 |
+
self.updown = None
|
| 43 |
+
if self.downsample:
|
| 44 |
+
self.updown = sp.SparseDownsample(2)
|
| 45 |
+
elif self.upsample:
|
| 46 |
+
self.updown = sp.SparseUpsample(2)
|
| 47 |
+
|
| 48 |
+
def _updown(self, x: sp.SparseTensor) -> sp.SparseTensor:
|
| 49 |
+
if self.updown is not None:
|
| 50 |
+
x = self.updown(x)
|
| 51 |
+
return x
|
| 52 |
+
|
| 53 |
+
def _forward(self, x: sp.SparseTensor) -> sp.SparseTensor:
|
| 54 |
+
x = self._updown(x)
|
| 55 |
+
h = x.replace(self.norm1(x.feats))
|
| 56 |
+
h = h.replace(F.silu(h.feats))
|
| 57 |
+
h = self.conv1(h)
|
| 58 |
+
h = h.replace(self.norm2(h.feats))
|
| 59 |
+
h = h.replace(F.silu(h.feats))
|
| 60 |
+
h = self.conv2(h)
|
| 61 |
+
h = h + self.skip_connection(x)
|
| 62 |
+
|
| 63 |
+
return h
|
| 64 |
+
|
| 65 |
+
def forward(self, x: torch.Tensor):
|
| 66 |
+
if self.use_checkpoint:
|
| 67 |
+
return torch.utils.checkpoint.checkpoint(
|
| 68 |
+
self._forward, x, use_reentrant=False
|
| 69 |
+
)
|
| 70 |
+
else:
|
| 71 |
+
return self._forward(x)
|
modules/sparse/conv/__init__.py
ADDED
|
@@ -0,0 +1,44 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# MIT License
|
| 2 |
+
|
| 3 |
+
# Copyright (c) Microsoft Corporation.
|
| 4 |
+
# Copyright (c) 2025 VAST-AI-Research and contributors.
|
| 5 |
+
|
| 6 |
+
# Permission is hereby granted, free of charge, to any person obtaining a copy
|
| 7 |
+
# of this software and associated documentation files (the "Software"), to deal
|
| 8 |
+
# in the Software without restriction, including without limitation the rights
|
| 9 |
+
# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
|
| 10 |
+
# copies of the Software, and to permit persons to whom the Software is
|
| 11 |
+
# furnished to do so, subject to the following conditions:
|
| 12 |
+
|
| 13 |
+
# The above copyright notice and this permission notice shall be included in all
|
| 14 |
+
# copies or substantial portions of the Software.
|
| 15 |
+
|
| 16 |
+
# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
|
| 17 |
+
# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
|
| 18 |
+
# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
|
| 19 |
+
# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
|
| 20 |
+
# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
|
| 21 |
+
# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
|
| 22 |
+
# SOFTWARE
|
| 23 |
+
|
| 24 |
+
from .. import BACKEND
|
| 25 |
+
|
| 26 |
+
|
| 27 |
+
SPCONV_ALGO = 'auto' # 'auto', 'implicit_gemm', 'native'
|
| 28 |
+
|
| 29 |
+
def __from_env():
|
| 30 |
+
import os
|
| 31 |
+
|
| 32 |
+
global SPCONV_ALGO
|
| 33 |
+
env_spconv_algo = os.environ.get('SPCONV_ALGO')
|
| 34 |
+
if env_spconv_algo is not None and env_spconv_algo in ['auto', 'implicit_gemm', 'native']:
|
| 35 |
+
SPCONV_ALGO = env_spconv_algo
|
| 36 |
+
print(f"[SPARSE][CONV] spconv algo: {SPCONV_ALGO}")
|
| 37 |
+
|
| 38 |
+
|
| 39 |
+
__from_env()
|
| 40 |
+
|
| 41 |
+
if BACKEND == 'torchsparse':
|
| 42 |
+
from .conv_torchsparse import *
|
| 43 |
+
elif BACKEND == 'spconv':
|
| 44 |
+
from .conv_spconv import *
|
modules/sparse/conv/conv_spconv.py
ADDED
|
@@ -0,0 +1,107 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# MIT License
|
| 2 |
+
|
| 3 |
+
# Copyright (c) Microsoft Corporation.
|
| 4 |
+
# Copyright (c) 2025 VAST-AI-Research and contributors.
|
| 5 |
+
|
| 6 |
+
# Permission is hereby granted, free of charge, to any person obtaining a copy
|
| 7 |
+
# of this software and associated documentation files (the "Software"), to deal
|
| 8 |
+
# in the Software without restriction, including without limitation the rights
|
| 9 |
+
# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
|
| 10 |
+
# copies of the Software, and to permit persons to whom the Software is
|
| 11 |
+
# furnished to do so, subject to the following conditions:
|
| 12 |
+
|
| 13 |
+
# The above copyright notice and this permission notice shall be included in all
|
| 14 |
+
# copies or substantial portions of the Software.
|
| 15 |
+
|
| 16 |
+
# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
|
| 17 |
+
# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
|
| 18 |
+
# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
|
| 19 |
+
# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
|
| 20 |
+
# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
|
| 21 |
+
# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
|
| 22 |
+
# SOFTWARE
|
| 23 |
+
|
| 24 |
+
import torch
|
| 25 |
+
import torch.nn as nn
|
| 26 |
+
from .. import SparseTensor
|
| 27 |
+
from .. import DEBUG
|
| 28 |
+
from . import SPCONV_ALGO
|
| 29 |
+
import spconv.pytorch as spconv
|
| 30 |
+
|
| 31 |
+
class SparseConv3d(nn.Module):
|
| 32 |
+
def __init__(self, in_channels, out_channels, kernel_size, stride=1, dilation=1, padding=None, bias=True, indice_key=None):
|
| 33 |
+
super(SparseConv3d, self).__init__()
|
| 34 |
+
# if 'spconv' not in globals():
|
| 35 |
+
# import spconv.pytorch as spconv
|
| 36 |
+
algo = None
|
| 37 |
+
if SPCONV_ALGO == 'native':
|
| 38 |
+
algo = spconv.ConvAlgo.Native
|
| 39 |
+
elif SPCONV_ALGO == 'implicit_gemm':
|
| 40 |
+
algo = spconv.ConvAlgo.MaskImplicitGemm
|
| 41 |
+
if stride == 1 and (padding is None):
|
| 42 |
+
self.conv = spconv.SubMConv3d(in_channels, out_channels, kernel_size, dilation=dilation, bias=bias, indice_key=indice_key, algo=algo)
|
| 43 |
+
else:
|
| 44 |
+
self.conv = spconv.SparseConv3d(in_channels, out_channels, kernel_size, stride=stride, dilation=dilation, padding=padding, bias=bias, indice_key=indice_key, algo=algo)
|
| 45 |
+
self.stride = tuple(stride) if isinstance(stride, (list, tuple)) else (stride, stride, stride)
|
| 46 |
+
self.padding = padding
|
| 47 |
+
|
| 48 |
+
def forward(self, x: SparseTensor) -> SparseTensor:
|
| 49 |
+
spatial_changed = any(s != 1 for s in self.stride) or (self.padding is not None)
|
| 50 |
+
|
| 51 |
+
dtype_ = x.feats.dtype
|
| 52 |
+
x = x.replace(x.feats.type(torch.float32))
|
| 53 |
+
new_data = self.conv(x.data)
|
| 54 |
+
new_shape = [x.shape[0], self.conv.out_channels]
|
| 55 |
+
new_layout = None if spatial_changed else x.layout
|
| 56 |
+
|
| 57 |
+
if spatial_changed and (x.shape[0] != 1):
|
| 58 |
+
# spconv was non-1 stride will break the contiguous of the output tensor, sort by the coords
|
| 59 |
+
fwd = new_data.indices[:, 0].argsort()
|
| 60 |
+
bwd = torch.zeros_like(fwd).scatter_(0, fwd, torch.arange(fwd.shape[0], device=fwd.device))
|
| 61 |
+
sorted_feats = new_data.features[fwd]
|
| 62 |
+
sorted_coords = new_data.indices[fwd]
|
| 63 |
+
unsorted_data = new_data
|
| 64 |
+
new_data = spconv.SparseConvTensor(sorted_feats, sorted_coords, unsorted_data.spatial_shape, unsorted_data.batch_size) # type: ignore
|
| 65 |
+
|
| 66 |
+
out = SparseTensor(
|
| 67 |
+
new_data, shape=torch.Size(new_shape), layout=new_layout,
|
| 68 |
+
scale=tuple([s * stride for s, stride in zip(x._scale, self.stride)]),
|
| 69 |
+
spatial_cache=x._spatial_cache,
|
| 70 |
+
)
|
| 71 |
+
out = out.replace(out.feats.type(dtype_))
|
| 72 |
+
|
| 73 |
+
if spatial_changed and (x.shape[0] != 1):
|
| 74 |
+
out.register_spatial_cache(f'conv_{self.stride}_unsorted_data', unsorted_data)
|
| 75 |
+
out.register_spatial_cache(f'conv_{self.stride}_sort_bwd', bwd)
|
| 76 |
+
|
| 77 |
+
return out
|
| 78 |
+
|
| 79 |
+
class SparseInverseConv3d(nn.Module):
|
| 80 |
+
def __init__(self, in_channels, out_channels, kernel_size, stride=1, dilation=1, bias=True, indice_key=None):
|
| 81 |
+
super(SparseInverseConv3d, self).__init__()
|
| 82 |
+
if 'spconv' not in globals():
|
| 83 |
+
import spconv.pytorch as spconv
|
| 84 |
+
self.conv = spconv.SparseInverseConv3d(in_channels, out_channels, kernel_size, bias=bias, indice_key=indice_key)
|
| 85 |
+
self.stride = tuple(stride) if isinstance(stride, (list, tuple)) else (stride, stride, stride)
|
| 86 |
+
|
| 87 |
+
def forward(self, x: SparseTensor) -> SparseTensor:
|
| 88 |
+
spatial_changed = any(s != 1 for s in self.stride)
|
| 89 |
+
if spatial_changed:
|
| 90 |
+
# recover the original spconv order
|
| 91 |
+
data = x.get_spatial_cache(f'conv_{self.stride}_unsorted_data')
|
| 92 |
+
bwd = x.get_spatial_cache(f'conv_{self.stride}_sort_bwd')
|
| 93 |
+
data = data.replace_feature(x.feats[bwd])
|
| 94 |
+
if DEBUG:
|
| 95 |
+
assert torch.equal(data.indices, x.coords[bwd]), 'Recover the original order failed'
|
| 96 |
+
else:
|
| 97 |
+
data = x.data
|
| 98 |
+
|
| 99 |
+
new_data = self.conv(data)
|
| 100 |
+
new_shape = [x.shape[0], self.conv.out_channels]
|
| 101 |
+
new_layout = None if spatial_changed else x.layout
|
| 102 |
+
out = SparseTensor(
|
| 103 |
+
new_data, shape=torch.Size(new_shape), layout=new_layout,
|
| 104 |
+
scale=tuple([s // stride for s, stride in zip(x._scale, self.stride)]),
|
| 105 |
+
spatial_cache=x._spatial_cache,
|
| 106 |
+
)
|
| 107 |
+
return out
|
modules/sparse/conv/conv_torchsparse.py
ADDED
|
@@ -0,0 +1,60 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# MIT License
|
| 2 |
+
|
| 3 |
+
# Copyright (c) Microsoft Corporation.
|
| 4 |
+
# Copyright (c) 2025 VAST-AI-Research and contributors.
|
| 5 |
+
|
| 6 |
+
# Permission is hereby granted, free of charge, to any person obtaining a copy
|
| 7 |
+
# of this software and associated documentation files (the "Software"), to deal
|
| 8 |
+
# in the Software without restriction, including without limitation the rights
|
| 9 |
+
# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
|
| 10 |
+
# copies of the Software, and to permit persons to whom the Software is
|
| 11 |
+
# furnished to do so, subject to the following conditions:
|
| 12 |
+
|
| 13 |
+
# The above copyright notice and this permission notice shall be included in all
|
| 14 |
+
# copies or substantial portions of the Software.
|
| 15 |
+
|
| 16 |
+
# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
|
| 17 |
+
# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
|
| 18 |
+
# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
|
| 19 |
+
# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
|
| 20 |
+
# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
|
| 21 |
+
# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
|
| 22 |
+
# SOFTWARE
|
| 23 |
+
|
| 24 |
+
import torch
|
| 25 |
+
import torch.nn as nn
|
| 26 |
+
from .. import SparseTensor
|
| 27 |
+
|
| 28 |
+
class SparseConv3d(nn.Module):
|
| 29 |
+
def __init__(self, in_channels, out_channels, kernel_size, stride=1, dilation=1, bias=True, indice_key=None):
|
| 30 |
+
super(SparseConv3d, self).__init__()
|
| 31 |
+
if 'torchsparse' not in globals():
|
| 32 |
+
import torchsparse
|
| 33 |
+
self.conv = torchsparse.nn.Conv3d(in_channels, out_channels, kernel_size, stride, 0, dilation, bias)
|
| 34 |
+
|
| 35 |
+
def forward(self, x: SparseTensor) -> SparseTensor:
|
| 36 |
+
out = self.conv(x.data)
|
| 37 |
+
new_shape = [x.shape[0], self.conv.out_channels]
|
| 38 |
+
out = SparseTensor(out, shape=torch.Size(new_shape), layout=x.layout if all(s == 1 for s in self.conv.stride) else None)
|
| 39 |
+
out._spatial_cache = x._spatial_cache
|
| 40 |
+
out._scale = tuple([s * stride for s, stride in zip(x._scale, self.conv.stride)])
|
| 41 |
+
return out
|
| 42 |
+
|
| 43 |
+
|
| 44 |
+
class SparseInverseConv3d(nn.Module):
|
| 45 |
+
def __init__(self, in_channels, out_channels, kernel_size, stride=1, dilation=1, bias=True, indice_key=None):
|
| 46 |
+
super(SparseInverseConv3d, self).__init__()
|
| 47 |
+
if 'torchsparse' not in globals():
|
| 48 |
+
import torchsparse
|
| 49 |
+
self.conv = torchsparse.nn.Conv3d(in_channels, out_channels, kernel_size, stride, 0, dilation, bias, transposed=True)
|
| 50 |
+
|
| 51 |
+
def forward(self, x: SparseTensor) -> SparseTensor:
|
| 52 |
+
out = self.conv(x.data)
|
| 53 |
+
new_shape = [x.shape[0], self.conv.out_channels]
|
| 54 |
+
out = SparseTensor(out, shape=torch.Size(new_shape), layout=x.layout if all(s == 1 for s in self.conv.stride) else None)
|
| 55 |
+
out._spatial_cache = x._spatial_cache
|
| 56 |
+
out._scale = tuple([s // stride for s, stride in zip(x._scale, self.conv.stride)])
|
| 57 |
+
return out
|
| 58 |
+
|
| 59 |
+
|
| 60 |
+
|
modules/sparse/linear.py
ADDED
|
@@ -0,0 +1,38 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
|
| 2 |
+
# MIT License
|
| 3 |
+
|
| 4 |
+
# Copyright (c) Microsoft Corporation.
|
| 5 |
+
# Copyright (c) 2025 VAST-AI-Research and contributors.
|
| 6 |
+
|
| 7 |
+
# Permission is hereby granted, free of charge, to any person obtaining a copy
|
| 8 |
+
# of this software and associated documentation files (the "Software"), to deal
|
| 9 |
+
# in the Software without restriction, including without limitation the rights
|
| 10 |
+
# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
|
| 11 |
+
# copies of the Software, and to permit persons to whom the Software is
|
| 12 |
+
# furnished to do so, subject to the following conditions:
|
| 13 |
+
|
| 14 |
+
# The above copyright notice and this permission notice shall be included in all
|
| 15 |
+
# copies or substantial portions of the Software.
|
| 16 |
+
|
| 17 |
+
# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
|
| 18 |
+
# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
|
| 19 |
+
# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
|
| 20 |
+
# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
|
| 21 |
+
# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
|
| 22 |
+
# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
|
| 23 |
+
# SOFTWARE
|
| 24 |
+
|
| 25 |
+
import torch
|
| 26 |
+
import torch.nn as nn
|
| 27 |
+
from . import SparseTensor
|
| 28 |
+
|
| 29 |
+
__all__ = [
|
| 30 |
+
'SparseLinear'
|
| 31 |
+
]
|
| 32 |
+
|
| 33 |
+
class SparseLinear(nn.Linear):
|
| 34 |
+
def __init__(self, in_features, out_features, bias=True):
|
| 35 |
+
super(SparseLinear, self).__init__(in_features, out_features, bias)
|
| 36 |
+
|
| 37 |
+
def forward(self, input: SparseTensor) -> SparseTensor:
|
| 38 |
+
return input.replace(super().forward(input.feats))
|
modules/sparse/nonlinearity.py
ADDED
|
@@ -0,0 +1,58 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# MIT License
|
| 2 |
+
|
| 3 |
+
# Copyright (c) Microsoft Corporation.
|
| 4 |
+
# Copyright (c) 2025 VAST-AI-Research and contributors.
|
| 5 |
+
|
| 6 |
+
# Permission is hereby granted, free of charge, to any person obtaining a copy
|
| 7 |
+
# of this software and associated documentation files (the "Software"), to deal
|
| 8 |
+
# in the Software without restriction, including without limitation the rights
|
| 9 |
+
# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
|
| 10 |
+
# copies of the Software, and to permit persons to whom the Software is
|
| 11 |
+
# furnished to do so, subject to the following conditions:
|
| 12 |
+
|
| 13 |
+
# The above copyright notice and this permission notice shall be included in all
|
| 14 |
+
# copies or substantial portions of the Software.
|
| 15 |
+
|
| 16 |
+
# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
|
| 17 |
+
# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
|
| 18 |
+
# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
|
| 19 |
+
# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
|
| 20 |
+
# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
|
| 21 |
+
# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
|
| 22 |
+
# SOFTWARE
|
| 23 |
+
|
| 24 |
+
import torch
|
| 25 |
+
import torch.nn as nn
|
| 26 |
+
from . import SparseTensor
|
| 27 |
+
|
| 28 |
+
__all__ = [
|
| 29 |
+
'SparseReLU',
|
| 30 |
+
'SparseSiLU',
|
| 31 |
+
'SparseGELU',
|
| 32 |
+
'SparseActivation'
|
| 33 |
+
]
|
| 34 |
+
|
| 35 |
+
|
| 36 |
+
class SparseReLU(nn.ReLU):
|
| 37 |
+
def forward(self, input: SparseTensor) -> SparseTensor:
|
| 38 |
+
return input.replace(super().forward(input.feats))
|
| 39 |
+
|
| 40 |
+
|
| 41 |
+
class SparseSiLU(nn.SiLU):
|
| 42 |
+
def forward(self, input: SparseTensor) -> SparseTensor:
|
| 43 |
+
return input.replace(super().forward(input.feats))
|
| 44 |
+
|
| 45 |
+
|
| 46 |
+
class SparseGELU(nn.GELU):
|
| 47 |
+
def forward(self, input: SparseTensor) -> SparseTensor:
|
| 48 |
+
return input.replace(super().forward(input.feats))
|
| 49 |
+
|
| 50 |
+
|
| 51 |
+
class SparseActivation(nn.Module):
|
| 52 |
+
def __init__(self, activation: nn.Module):
|
| 53 |
+
super().__init__()
|
| 54 |
+
self.activation = activation
|
| 55 |
+
|
| 56 |
+
def forward(self, input: SparseTensor) -> SparseTensor:
|
| 57 |
+
return input.replace(self.activation(input.feats))
|
| 58 |
+
|
modules/sparse/norm.py
ADDED
|
@@ -0,0 +1,81 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# MIT License
|
| 2 |
+
|
| 3 |
+
# Copyright (c) Microsoft Corporation.
|
| 4 |
+
# Copyright (c) 2025 VAST-AI-Research and contributors.
|
| 5 |
+
|
| 6 |
+
# Permission is hereby granted, free of charge, to any person obtaining a copy
|
| 7 |
+
# of this software and associated documentation files (the "Software"), to deal
|
| 8 |
+
# in the Software without restriction, including without limitation the rights
|
| 9 |
+
# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
|
| 10 |
+
# copies of the Software, and to permit persons to whom the Software is
|
| 11 |
+
# furnished to do so, subject to the following conditions:
|
| 12 |
+
|
| 13 |
+
# The above copyright notice and this permission notice shall be included in all
|
| 14 |
+
# copies or substantial portions of the Software.
|
| 15 |
+
|
| 16 |
+
# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
|
| 17 |
+
# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
|
| 18 |
+
# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
|
| 19 |
+
# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
|
| 20 |
+
# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
|
| 21 |
+
# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
|
| 22 |
+
# SOFTWARE
|
| 23 |
+
|
| 24 |
+
import torch
|
| 25 |
+
import torch.nn as nn
|
| 26 |
+
from . import SparseTensor
|
| 27 |
+
from . import DEBUG
|
| 28 |
+
|
| 29 |
+
__all__ = [
|
| 30 |
+
'SparseGroupNorm',
|
| 31 |
+
'SparseLayerNorm',
|
| 32 |
+
'SparseGroupNorm32',
|
| 33 |
+
'SparseLayerNorm32',
|
| 34 |
+
]
|
| 35 |
+
|
| 36 |
+
|
| 37 |
+
class SparseGroupNorm(nn.GroupNorm):
|
| 38 |
+
def __init__(self, num_groups, num_channels, eps=1e-5, affine=True):
|
| 39 |
+
super(SparseGroupNorm, self).__init__(num_groups, num_channels, eps, affine)
|
| 40 |
+
|
| 41 |
+
def forward(self, input: SparseTensor) -> SparseTensor:
|
| 42 |
+
nfeats = torch.zeros_like(input.feats)
|
| 43 |
+
for k in range(input.shape[0]):
|
| 44 |
+
if DEBUG:
|
| 45 |
+
assert (input.coords[input.layout[k], 0] == k).all(), f"SparseGroupNorm: batch index mismatch"
|
| 46 |
+
bfeats = input.feats[input.layout[k]]
|
| 47 |
+
bfeats = bfeats.permute(1, 0).reshape(1, input.shape[1], -1)
|
| 48 |
+
bfeats = super().forward(bfeats)
|
| 49 |
+
bfeats = bfeats.reshape(input.shape[1], -1).permute(1, 0)
|
| 50 |
+
nfeats[input.layout[k]] = bfeats
|
| 51 |
+
return input.replace(nfeats)
|
| 52 |
+
|
| 53 |
+
|
| 54 |
+
class SparseLayerNorm(nn.LayerNorm):
|
| 55 |
+
def __init__(self, normalized_shape, eps=1e-5, elementwise_affine=True):
|
| 56 |
+
super(SparseLayerNorm, self).__init__(normalized_shape, eps, elementwise_affine)
|
| 57 |
+
|
| 58 |
+
def forward(self, input: SparseTensor) -> SparseTensor:
|
| 59 |
+
nfeats = torch.zeros_like(input.feats)
|
| 60 |
+
for k in range(input.shape[0]):
|
| 61 |
+
bfeats = input.feats[input.layout[k]]
|
| 62 |
+
bfeats = bfeats.permute(1, 0).reshape(1, input.shape[1], -1)
|
| 63 |
+
bfeats = super().forward(bfeats)
|
| 64 |
+
bfeats = bfeats.reshape(input.shape[1], -1).permute(1, 0)
|
| 65 |
+
nfeats[input.layout[k]] = bfeats
|
| 66 |
+
return input.replace(nfeats)
|
| 67 |
+
|
| 68 |
+
|
| 69 |
+
class SparseGroupNorm32(SparseGroupNorm):
|
| 70 |
+
"""
|
| 71 |
+
A GroupNorm layer that converts to float32 before the forward pass.
|
| 72 |
+
"""
|
| 73 |
+
def forward(self, x: SparseTensor) -> SparseTensor:
|
| 74 |
+
return super().forward(x.float()).type(x.dtype)
|
| 75 |
+
|
| 76 |
+
class SparseLayerNorm32(SparseLayerNorm):
|
| 77 |
+
"""
|
| 78 |
+
A LayerNorm layer that converts to float32 before the forward pass.
|
| 79 |
+
"""
|
| 80 |
+
def forward(self, x: SparseTensor) -> SparseTensor:
|
| 81 |
+
return super().forward(x.float()).type(x.dtype)
|
modules/sparse/spatial.py
ADDED
|
@@ -0,0 +1,158 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# MIT License
|
| 2 |
+
|
| 3 |
+
# Copyright (c) Microsoft Corporation.
|
| 4 |
+
# Copyright (c) 2025 VAST-AI-Research and contributors.
|
| 5 |
+
|
| 6 |
+
# Permission is hereby granted, free of charge, to any person obtaining a copy
|
| 7 |
+
# of this software and associated documentation files (the "Software"), to deal
|
| 8 |
+
# in the Software without restriction, including without limitation the rights
|
| 9 |
+
# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
|
| 10 |
+
# copies of the Software, and to permit persons to whom the Software is
|
| 11 |
+
# furnished to do so, subject to the following conditions:
|
| 12 |
+
|
| 13 |
+
# The above copyright notice and this permission notice shall be included in all
|
| 14 |
+
# copies or substantial portions of the Software.
|
| 15 |
+
|
| 16 |
+
# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
|
| 17 |
+
# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
|
| 18 |
+
# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
|
| 19 |
+
# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
|
| 20 |
+
# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
|
| 21 |
+
# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
|
| 22 |
+
# SOFTWARE
|
| 23 |
+
|
| 24 |
+
from typing import *
|
| 25 |
+
import torch
|
| 26 |
+
import torch.nn as nn
|
| 27 |
+
from . import SparseTensor
|
| 28 |
+
|
| 29 |
+
__all__ = [
|
| 30 |
+
"SparseDownsample",
|
| 31 |
+
"SparseUpsample",
|
| 32 |
+
"SparseSubdivide",
|
| 33 |
+
]
|
| 34 |
+
|
| 35 |
+
|
| 36 |
+
class SparseDownsample(nn.Module):
|
| 37 |
+
"""
|
| 38 |
+
Downsample a sparse tensor by a factor of `factor`.
|
| 39 |
+
Implemented as average pooling.
|
| 40 |
+
"""
|
| 41 |
+
|
| 42 |
+
def __init__(self, factor: Union[int, Tuple[int, ...], List[int]]):
|
| 43 |
+
super(SparseDownsample, self).__init__()
|
| 44 |
+
self.factor = tuple(factor) if isinstance(factor, (list, tuple)) else factor
|
| 45 |
+
|
| 46 |
+
def forward(self, input: SparseTensor) -> SparseTensor:
|
| 47 |
+
DIM = input.coords.shape[-1] - 1
|
| 48 |
+
factor = self.factor if isinstance(self.factor, tuple) else (self.factor,) * DIM
|
| 49 |
+
assert DIM == len(
|
| 50 |
+
factor
|
| 51 |
+
), "Input coordinates must have the same dimension as the downsample factor."
|
| 52 |
+
|
| 53 |
+
coord = list(input.coords.unbind(dim=-1))
|
| 54 |
+
for i, f in enumerate(factor):
|
| 55 |
+
coord[i + 1] = coord[i + 1] // f
|
| 56 |
+
|
| 57 |
+
MAX = [coord[i + 1].max().item() + 1 for i in range(DIM)]
|
| 58 |
+
OFFSET = torch.cumprod(torch.tensor(MAX[::-1]), 0).tolist()[::-1] + [1]
|
| 59 |
+
code = sum([c * o for c, o in zip(coord, OFFSET)])
|
| 60 |
+
code, idx = code.unique(return_inverse=True)
|
| 61 |
+
|
| 62 |
+
new_feats = torch.scatter_reduce(
|
| 63 |
+
torch.zeros(
|
| 64 |
+
code.shape[0],
|
| 65 |
+
input.feats.shape[1],
|
| 66 |
+
device=input.feats.device,
|
| 67 |
+
dtype=input.feats.dtype,
|
| 68 |
+
),
|
| 69 |
+
dim=0,
|
| 70 |
+
index=idx.unsqueeze(1).expand(-1, input.feats.shape[1]),
|
| 71 |
+
src=input.feats,
|
| 72 |
+
# reduce='mean',
|
| 73 |
+
reduce="amax",
|
| 74 |
+
)
|
| 75 |
+
new_coords = torch.stack(
|
| 76 |
+
[code // OFFSET[0]]
|
| 77 |
+
+ [(code // OFFSET[i + 1]) % MAX[i] for i in range(DIM)],
|
| 78 |
+
dim=-1,
|
| 79 |
+
)
|
| 80 |
+
out = SparseTensor(
|
| 81 |
+
new_feats,
|
| 82 |
+
new_coords,
|
| 83 |
+
input.shape,
|
| 84 |
+
)
|
| 85 |
+
out._scale = tuple([s // f for s, f in zip(input._scale, factor)])
|
| 86 |
+
out._spatial_cache = input._spatial_cache
|
| 87 |
+
|
| 88 |
+
out.register_spatial_cache(f"upsample_{factor}_coords", input.coords)
|
| 89 |
+
out.register_spatial_cache(f"upsample_{factor}_layout", input.layout)
|
| 90 |
+
out.register_spatial_cache(f"upsample_{factor}_idx", idx)
|
| 91 |
+
|
| 92 |
+
return out
|
| 93 |
+
|
| 94 |
+
|
| 95 |
+
class SparseUpsample(nn.Module):
|
| 96 |
+
"""
|
| 97 |
+
Upsample a sparse tensor by a factor of `factor`.
|
| 98 |
+
Implemented as nearest neighbor interpolation.
|
| 99 |
+
"""
|
| 100 |
+
|
| 101 |
+
def __init__(self, factor: Union[int, Tuple[int, int, int], List[int]]):
|
| 102 |
+
super(SparseUpsample, self).__init__()
|
| 103 |
+
self.factor = tuple(factor) if isinstance(factor, (list, tuple)) else factor
|
| 104 |
+
|
| 105 |
+
def forward(self, input: SparseTensor) -> SparseTensor:
|
| 106 |
+
DIM = input.coords.shape[-1] - 1
|
| 107 |
+
factor = self.factor if isinstance(self.factor, tuple) else (self.factor,) * DIM
|
| 108 |
+
assert DIM == len(
|
| 109 |
+
factor
|
| 110 |
+
), "Input coordinates must have the same dimension as the upsample factor."
|
| 111 |
+
|
| 112 |
+
new_coords = input.get_spatial_cache(f"upsample_{factor}_coords")
|
| 113 |
+
new_layout = input.get_spatial_cache(f"upsample_{factor}_layout")
|
| 114 |
+
idx = input.get_spatial_cache(f"upsample_{factor}_idx")
|
| 115 |
+
if any([x is None for x in [new_coords, new_layout, idx]]):
|
| 116 |
+
raise ValueError(
|
| 117 |
+
"Upsample cache not found. SparseUpsample must be paired with SparseDownsample."
|
| 118 |
+
)
|
| 119 |
+
new_feats = input.feats[idx]
|
| 120 |
+
out = SparseTensor(new_feats, new_coords, input.shape, new_layout)
|
| 121 |
+
out._scale = tuple([s * f for s, f in zip(input._scale, factor)])
|
| 122 |
+
out._spatial_cache = input._spatial_cache
|
| 123 |
+
return out
|
| 124 |
+
|
| 125 |
+
|
| 126 |
+
class SparseSubdivide(nn.Module):
|
| 127 |
+
"""
|
| 128 |
+
Upsample a sparse tensor by a factor of `factor`.
|
| 129 |
+
Implemented as nearest neighbor interpolation.
|
| 130 |
+
"""
|
| 131 |
+
|
| 132 |
+
def __init__(self):
|
| 133 |
+
super(SparseSubdivide, self).__init__()
|
| 134 |
+
|
| 135 |
+
def forward(self, input: SparseTensor) -> SparseTensor:
|
| 136 |
+
DIM = input.coords.shape[-1] - 1
|
| 137 |
+
# upsample scale=2^DIM
|
| 138 |
+
n_cube = torch.ones([2] * DIM, device=input.device, dtype=torch.int)
|
| 139 |
+
n_coords = torch.nonzero(n_cube)
|
| 140 |
+
n_coords = torch.cat([torch.zeros_like(n_coords[:, :1]), n_coords], dim=-1)
|
| 141 |
+
factor = n_coords.shape[0]
|
| 142 |
+
assert factor == 2**DIM
|
| 143 |
+
# print(n_coords.shape)
|
| 144 |
+
new_coords = input.coords.clone()
|
| 145 |
+
new_coords[:, 1:] *= 2
|
| 146 |
+
new_coords = new_coords.unsqueeze(1) + n_coords.unsqueeze(0).to(
|
| 147 |
+
new_coords.dtype
|
| 148 |
+
)
|
| 149 |
+
|
| 150 |
+
new_feats = input.feats.unsqueeze(1).expand(
|
| 151 |
+
input.feats.shape[0], factor, *input.feats.shape[1:]
|
| 152 |
+
)
|
| 153 |
+
out = SparseTensor(
|
| 154 |
+
new_feats.flatten(0, 1), new_coords.flatten(0, 1), input.shape
|
| 155 |
+
)
|
| 156 |
+
out._scale = input._scale * 2
|
| 157 |
+
out._spatial_cache = input._spatial_cache
|
| 158 |
+
return out
|
modules/sparse/transformer/__init__.py
ADDED
|
@@ -0,0 +1,26 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# MIT License
|
| 2 |
+
|
| 3 |
+
# Copyright (c) Microsoft Corporation.
|
| 4 |
+
# Copyright (c) 2025 VAST-AI-Research and contributors.
|
| 5 |
+
|
| 6 |
+
# Permission is hereby granted, free of charge, to any person obtaining a copy
|
| 7 |
+
# of this software and associated documentation files (the "Software"), to deal
|
| 8 |
+
# in the Software without restriction, including without limitation the rights
|
| 9 |
+
# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
|
| 10 |
+
# copies of the Software, and to permit persons to whom the Software is
|
| 11 |
+
# furnished to do so, subject to the following conditions:
|
| 12 |
+
|
| 13 |
+
# The above copyright notice and this permission notice shall be included in all
|
| 14 |
+
# copies or substantial portions of the Software.
|
| 15 |
+
|
| 16 |
+
# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
|
| 17 |
+
# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
|
| 18 |
+
# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
|
| 19 |
+
# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
|
| 20 |
+
# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
|
| 21 |
+
# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
|
| 22 |
+
# SOFTWARE
|
| 23 |
+
|
| 24 |
+
from .blocks import *
|
| 25 |
+
from .modulated import *
|
| 26 |
+
from .bases import *
|
modules/sparse/transformer/bases.py
ADDED
|
@@ -0,0 +1,234 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# MIT License
|
| 2 |
+
|
| 3 |
+
# Copyright (c) Microsoft Corporation.
|
| 4 |
+
# Copyright (c) 2025 VAST-AI-Research and contributors.
|
| 5 |
+
|
| 6 |
+
# Permission is hereby granted, free of charge, to any person obtaining a copy
|
| 7 |
+
# of this software and associated documentation files (the "Software"), to deal
|
| 8 |
+
# in the Software without restriction, including without limitation the rights
|
| 9 |
+
# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
|
| 10 |
+
# copies of the Software, and to permit persons to whom the Software is
|
| 11 |
+
# furnished to do so, subject to the following conditions:
|
| 12 |
+
|
| 13 |
+
# The above copyright notice and this permission notice shall be included in all
|
| 14 |
+
# copies or substantial portions of the Software.
|
| 15 |
+
|
| 16 |
+
# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
|
| 17 |
+
# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
|
| 18 |
+
# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
|
| 19 |
+
# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
|
| 20 |
+
# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
|
| 21 |
+
# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
|
| 22 |
+
# SOFTWARE
|
| 23 |
+
|
| 24 |
+
from typing import *
|
| 25 |
+
import torch
|
| 26 |
+
import torch.nn as nn
|
| 27 |
+
from ...utils import convert_module_to_f16, convert_module_to_f32
|
| 28 |
+
from ...transformer import AbsolutePositionEmbedder
|
| 29 |
+
from modules import sparse as sp
|
| 30 |
+
from .blocks import SparseTransformerBlock, SparseTransformerCrossBlock
|
| 31 |
+
|
| 32 |
+
|
| 33 |
+
def block_attn_config(self):
|
| 34 |
+
"""
|
| 35 |
+
Return the attention configuration of the model.
|
| 36 |
+
"""
|
| 37 |
+
for i in range(self.num_blocks):
|
| 38 |
+
if self.attn_mode == "shift_window":
|
| 39 |
+
yield "serialized", self.window_size, 0, (16 * (i % 2),) * 3, sp.SerializeMode.Z_ORDER
|
| 40 |
+
elif self.attn_mode == "shift_sequence":
|
| 41 |
+
yield "serialized", self.window_size, self.window_size // 2 * (i % 2), (0, 0, 0), sp.SerializeMode.Z_ORDER
|
| 42 |
+
elif self.attn_mode == "shift_order":
|
| 43 |
+
yield "serialized", self.window_size, 0, (0, 0, 0), sp.SerializeModes[i % 4]
|
| 44 |
+
elif self.attn_mode == "full":
|
| 45 |
+
yield "full", None, None, None, None
|
| 46 |
+
elif self.attn_mode == "swin":
|
| 47 |
+
yield "windowed", self.window_size, None, self.window_size // 2 * (i % 2), None
|
| 48 |
+
|
| 49 |
+
|
| 50 |
+
class SparseTransformerBase(nn.Module):
|
| 51 |
+
"""
|
| 52 |
+
Sparse Transformer without output layers.
|
| 53 |
+
Serve as the base class for encoder and decoder.
|
| 54 |
+
"""
|
| 55 |
+
def __init__(
|
| 56 |
+
self,
|
| 57 |
+
in_channels: int,
|
| 58 |
+
model_channels: int,
|
| 59 |
+
num_blocks: int,
|
| 60 |
+
num_heads: Optional[int] = None,
|
| 61 |
+
num_head_channels: Optional[int] = 64,
|
| 62 |
+
mlp_ratio: float = 4.0,
|
| 63 |
+
attn_mode: Literal["full", "shift_window", "shift_sequence", "shift_order", "swin"] = "full",
|
| 64 |
+
window_size: Optional[int] = None,
|
| 65 |
+
pe_mode: Literal["ape", "rope"] = "ape",
|
| 66 |
+
use_fp16: bool = False,
|
| 67 |
+
use_checkpoint: bool = False,
|
| 68 |
+
qk_rms_norm: bool = False,
|
| 69 |
+
):
|
| 70 |
+
super().__init__()
|
| 71 |
+
self.in_channels = in_channels
|
| 72 |
+
self.model_channels = model_channels
|
| 73 |
+
self.num_blocks = num_blocks
|
| 74 |
+
self.window_size = window_size
|
| 75 |
+
self.num_heads = num_heads or model_channels // num_head_channels
|
| 76 |
+
self.mlp_ratio = mlp_ratio
|
| 77 |
+
self.attn_mode = attn_mode
|
| 78 |
+
self.pe_mode = pe_mode
|
| 79 |
+
self.use_fp16 = use_fp16
|
| 80 |
+
self.use_checkpoint = use_checkpoint
|
| 81 |
+
self.qk_rms_norm = qk_rms_norm
|
| 82 |
+
self.dtype = torch.float16 if use_fp16 else torch.float32
|
| 83 |
+
|
| 84 |
+
if pe_mode == "ape":
|
| 85 |
+
self.pos_embedder = AbsolutePositionEmbedder(model_channels)
|
| 86 |
+
|
| 87 |
+
self.input_layer = sp.SparseLinear(in_channels, model_channels)
|
| 88 |
+
self.blocks = nn.ModuleList([
|
| 89 |
+
SparseTransformerBlock(
|
| 90 |
+
model_channels,
|
| 91 |
+
num_heads=self.num_heads,
|
| 92 |
+
mlp_ratio=self.mlp_ratio,
|
| 93 |
+
attn_mode=attn_mode,
|
| 94 |
+
window_size=window_size,
|
| 95 |
+
shift_sequence=shift_sequence,
|
| 96 |
+
shift_window=shift_window,
|
| 97 |
+
serialize_mode=serialize_mode,
|
| 98 |
+
use_checkpoint=self.use_checkpoint,
|
| 99 |
+
use_rope=(pe_mode == "rope"),
|
| 100 |
+
qk_rms_norm=self.qk_rms_norm,
|
| 101 |
+
)
|
| 102 |
+
for attn_mode, window_size, shift_sequence, shift_window, serialize_mode in block_attn_config(self)
|
| 103 |
+
])
|
| 104 |
+
|
| 105 |
+
@property
|
| 106 |
+
def device(self) -> torch.device:
|
| 107 |
+
"""
|
| 108 |
+
Return the device of the model.
|
| 109 |
+
"""
|
| 110 |
+
return next(self.parameters()).device
|
| 111 |
+
|
| 112 |
+
def convert_to_fp16(self) -> None:
|
| 113 |
+
"""
|
| 114 |
+
Convert the torso of the model to float16.
|
| 115 |
+
"""
|
| 116 |
+
self.blocks.apply(convert_module_to_f16)
|
| 117 |
+
|
| 118 |
+
def convert_to_fp32(self) -> None:
|
| 119 |
+
"""
|
| 120 |
+
Convert the torso of the model to float32.
|
| 121 |
+
"""
|
| 122 |
+
self.blocks.apply(convert_module_to_f32)
|
| 123 |
+
|
| 124 |
+
def initialize_weights(self) -> None:
|
| 125 |
+
# Initialize transformer layers:
|
| 126 |
+
def _basic_init(module):
|
| 127 |
+
if isinstance(module, nn.Linear):
|
| 128 |
+
torch.nn.init.xavier_uniform_(module.weight)
|
| 129 |
+
if module.bias is not None:
|
| 130 |
+
nn.init.constant_(module.bias, 0)
|
| 131 |
+
self.apply(_basic_init)
|
| 132 |
+
|
| 133 |
+
def forward(self, x: sp.SparseTensor) -> sp.SparseTensor:
|
| 134 |
+
h = self.input_layer(x)
|
| 135 |
+
if self.pe_mode == "ape" and len(self.blocks) != 0:
|
| 136 |
+
h = h + self.pos_embedder(x.coords[:, 1:])
|
| 137 |
+
for block in self.blocks:
|
| 138 |
+
h = block(h)
|
| 139 |
+
return h
|
| 140 |
+
|
| 141 |
+
class SparseTransformerCrossBase(nn.Module):
|
| 142 |
+
"""
|
| 143 |
+
Sparse Transformer without output layers.
|
| 144 |
+
Serve as the base class for encoder and decoder.
|
| 145 |
+
"""
|
| 146 |
+
def __init__(
|
| 147 |
+
self,
|
| 148 |
+
in_channels: int,
|
| 149 |
+
model_channels: int,
|
| 150 |
+
context_channels: int,
|
| 151 |
+
num_blocks: int,
|
| 152 |
+
num_heads: Optional[int] = None,
|
| 153 |
+
num_head_channels: Optional[int] = 64,
|
| 154 |
+
mlp_ratio: float = 4.0,
|
| 155 |
+
attn_mode: Literal["full", "shift_window", "shift_sequence", "shift_order", "swin"] = "full",
|
| 156 |
+
window_size: Optional[int] = None,
|
| 157 |
+
pe_mode: Literal["ape", "rope"] = "ape",
|
| 158 |
+
use_fp16: bool = False,
|
| 159 |
+
use_checkpoint: bool = False,
|
| 160 |
+
qk_rms_norm: bool = False,
|
| 161 |
+
):
|
| 162 |
+
super().__init__()
|
| 163 |
+
self.in_channels = in_channels
|
| 164 |
+
self.model_channels = model_channels
|
| 165 |
+
self.num_blocks = num_blocks
|
| 166 |
+
self.window_size = window_size
|
| 167 |
+
self.num_heads = num_heads or model_channels // num_head_channels
|
| 168 |
+
self.mlp_ratio = mlp_ratio
|
| 169 |
+
self.attn_mode = attn_mode
|
| 170 |
+
self.pe_mode = pe_mode
|
| 171 |
+
self.use_fp16 = use_fp16
|
| 172 |
+
self.use_checkpoint = use_checkpoint
|
| 173 |
+
self.qk_rms_norm = qk_rms_norm
|
| 174 |
+
self.dtype = torch.float16 if use_fp16 else torch.float32
|
| 175 |
+
|
| 176 |
+
if pe_mode == "ape":
|
| 177 |
+
self.pos_embedder_x = AbsolutePositionEmbedder(model_channels)
|
| 178 |
+
self.pos_embedder_ctx = AbsolutePositionEmbedder(context_channels)
|
| 179 |
+
|
| 180 |
+
self.input_layer = sp.SparseLinear(in_channels, model_channels)
|
| 181 |
+
self.blocks = nn.ModuleList([
|
| 182 |
+
SparseTransformerCrossBlock(
|
| 183 |
+
model_channels,
|
| 184 |
+
num_heads=self.num_heads,
|
| 185 |
+
ctx_channels=context_channels,
|
| 186 |
+
mlp_ratio=self.mlp_ratio,
|
| 187 |
+
attn_mode=attn_mode,
|
| 188 |
+
window_size=window_size,
|
| 189 |
+
shift_sequence=shift_sequence,
|
| 190 |
+
shift_window=shift_window,
|
| 191 |
+
serialize_mode=serialize_mode,
|
| 192 |
+
use_checkpoint=self.use_checkpoint,
|
| 193 |
+
use_rope=(pe_mode == "rope"),
|
| 194 |
+
qk_rms_norm=self.qk_rms_norm,
|
| 195 |
+
)
|
| 196 |
+
for attn_mode, window_size, shift_sequence, shift_window, serialize_mode in block_attn_config(self)
|
| 197 |
+
])
|
| 198 |
+
|
| 199 |
+
@property
|
| 200 |
+
def device(self) -> torch.device:
|
| 201 |
+
"""
|
| 202 |
+
Return the device of the model.
|
| 203 |
+
"""
|
| 204 |
+
return next(self.parameters()).device
|
| 205 |
+
|
| 206 |
+
def convert_to_fp16(self) -> None:
|
| 207 |
+
"""
|
| 208 |
+
Convert the torso of the model to float16.
|
| 209 |
+
"""
|
| 210 |
+
self.blocks.apply(convert_module_to_f16)
|
| 211 |
+
|
| 212 |
+
def convert_to_fp32(self) -> None:
|
| 213 |
+
"""
|
| 214 |
+
Convert the torso of the model to float32.
|
| 215 |
+
"""
|
| 216 |
+
self.blocks.apply(convert_module_to_f32)
|
| 217 |
+
|
| 218 |
+
def initialize_weights(self) -> None:
|
| 219 |
+
# Initialize transformer layers:
|
| 220 |
+
def _basic_init(module):
|
| 221 |
+
if isinstance(module, nn.Linear):
|
| 222 |
+
torch.nn.init.xavier_uniform_(module.weight)
|
| 223 |
+
if module.bias is not None:
|
| 224 |
+
nn.init.constant_(module.bias, 0)
|
| 225 |
+
self.apply(_basic_init)
|
| 226 |
+
|
| 227 |
+
def forward(self, x: sp.SparseTensor, context: sp.SparseTensor) -> sp.SparseTensor:
|
| 228 |
+
h = self.input_layer(x)
|
| 229 |
+
if self.pe_mode == "ape" and len(self.blocks) != 0:
|
| 230 |
+
h = h + self.pos_embedder_x(x.coords[:, 1:])
|
| 231 |
+
context = context + self.pos_embedder_ctx(context.coords[:, 1:])
|
| 232 |
+
for block in self.blocks:
|
| 233 |
+
h = block(h, context)
|
| 234 |
+
return h
|
modules/sparse/transformer/blocks.py
ADDED
|
@@ -0,0 +1,165 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from typing import *
|
| 2 |
+
import torch
|
| 3 |
+
import torch.nn as nn
|
| 4 |
+
from ..basic import SparseTensor
|
| 5 |
+
from ..linear import SparseLinear
|
| 6 |
+
from ..nonlinearity import SparseGELU
|
| 7 |
+
from ..attention import SparseMultiHeadAttention, SerializeMode
|
| 8 |
+
from ...norm import LayerNorm32
|
| 9 |
+
|
| 10 |
+
|
| 11 |
+
class SparseFeedForwardNet(nn.Module):
|
| 12 |
+
def __init__(self, channels: int, mlp_ratio: float = 4.0):
|
| 13 |
+
super().__init__()
|
| 14 |
+
self.mlp = nn.Sequential(
|
| 15 |
+
SparseLinear(channels, int(channels * mlp_ratio)),
|
| 16 |
+
SparseGELU(approximate="tanh"),
|
| 17 |
+
SparseLinear(int(channels * mlp_ratio), channels),
|
| 18 |
+
)
|
| 19 |
+
|
| 20 |
+
def forward(self, x: SparseTensor) -> SparseTensor:
|
| 21 |
+
return self.mlp(x)
|
| 22 |
+
|
| 23 |
+
|
| 24 |
+
class SparseTransformerBlock(nn.Module):
|
| 25 |
+
"""
|
| 26 |
+
Sparse Transformer block (MSA + FFN).
|
| 27 |
+
"""
|
| 28 |
+
|
| 29 |
+
def __init__(
|
| 30 |
+
self,
|
| 31 |
+
channels: int,
|
| 32 |
+
num_heads: int,
|
| 33 |
+
mlp_ratio: float = 4.0,
|
| 34 |
+
attn_mode: Literal[
|
| 35 |
+
"full", "shift_window", "shift_sequence", "shift_order", "swin"
|
| 36 |
+
] = "full",
|
| 37 |
+
window_size: Optional[int] = None,
|
| 38 |
+
shift_sequence: Optional[int] = None,
|
| 39 |
+
shift_window: Optional[Tuple[int, int, int]] = None,
|
| 40 |
+
serialize_mode: Optional[SerializeMode] = None,
|
| 41 |
+
use_checkpoint: bool = False,
|
| 42 |
+
use_rope: bool = False,
|
| 43 |
+
qk_rms_norm: bool = False,
|
| 44 |
+
qkv_bias: bool = True,
|
| 45 |
+
ln_affine: bool = False,
|
| 46 |
+
):
|
| 47 |
+
super().__init__()
|
| 48 |
+
self.use_checkpoint = use_checkpoint
|
| 49 |
+
self.norm1 = LayerNorm32(channels, elementwise_affine=ln_affine, eps=1e-6)
|
| 50 |
+
self.norm2 = LayerNorm32(channels, elementwise_affine=ln_affine, eps=1e-6)
|
| 51 |
+
self.attn = SparseMultiHeadAttention(
|
| 52 |
+
channels,
|
| 53 |
+
num_heads=num_heads,
|
| 54 |
+
attn_mode=attn_mode,
|
| 55 |
+
window_size=window_size,
|
| 56 |
+
shift_sequence=shift_sequence,
|
| 57 |
+
shift_window=shift_window,
|
| 58 |
+
serialize_mode=serialize_mode,
|
| 59 |
+
qkv_bias=qkv_bias,
|
| 60 |
+
use_rope=use_rope,
|
| 61 |
+
qk_rms_norm=qk_rms_norm,
|
| 62 |
+
)
|
| 63 |
+
self.mlp = SparseFeedForwardNet(
|
| 64 |
+
channels,
|
| 65 |
+
mlp_ratio=mlp_ratio,
|
| 66 |
+
)
|
| 67 |
+
|
| 68 |
+
def _forward(self, x: SparseTensor) -> SparseTensor:
|
| 69 |
+
h = x.replace(self.norm1(x.feats))
|
| 70 |
+
h = self.attn(h)
|
| 71 |
+
x = x + h
|
| 72 |
+
h = x.replace(self.norm2(x.feats))
|
| 73 |
+
h = self.mlp(h)
|
| 74 |
+
x = x + h
|
| 75 |
+
return x
|
| 76 |
+
|
| 77 |
+
def forward(self, x: SparseTensor) -> SparseTensor:
|
| 78 |
+
if self.use_checkpoint:
|
| 79 |
+
return torch.utils.checkpoint.checkpoint(
|
| 80 |
+
self._forward, x, use_reentrant=False
|
| 81 |
+
)
|
| 82 |
+
else:
|
| 83 |
+
return self._forward(x)
|
| 84 |
+
|
| 85 |
+
|
| 86 |
+
class SparseTransformerCrossBlock(nn.Module):
|
| 87 |
+
"""
|
| 88 |
+
Sparse Transformer cross-attention block (MSA + MCA + FFN).
|
| 89 |
+
"""
|
| 90 |
+
|
| 91 |
+
def __init__(
|
| 92 |
+
self,
|
| 93 |
+
channels: int,
|
| 94 |
+
ctx_channels: int,
|
| 95 |
+
num_heads: int,
|
| 96 |
+
mlp_ratio: float = 4.0,
|
| 97 |
+
attn_mode: Literal[
|
| 98 |
+
"full", "shift_window", "shift_sequence", "shift_order", "swin"
|
| 99 |
+
] = "full",
|
| 100 |
+
window_size: Optional[int] = None,
|
| 101 |
+
shift_sequence: Optional[int] = None,
|
| 102 |
+
shift_window: Optional[Tuple[int, int, int]] = None,
|
| 103 |
+
serialize_mode: Optional[SerializeMode] = None,
|
| 104 |
+
use_checkpoint: bool = False,
|
| 105 |
+
use_rope: bool = False,
|
| 106 |
+
qk_rms_norm: bool = False,
|
| 107 |
+
qk_rms_norm_cross: bool = False,
|
| 108 |
+
qkv_bias: bool = True,
|
| 109 |
+
ln_affine: bool = False,
|
| 110 |
+
):
|
| 111 |
+
super().__init__()
|
| 112 |
+
self.use_checkpoint = use_checkpoint
|
| 113 |
+
self.norm1 = LayerNorm32(channels, elementwise_affine=ln_affine, eps=1e-6)
|
| 114 |
+
self.norm2 = LayerNorm32(channels, elementwise_affine=ln_affine, eps=1e-6)
|
| 115 |
+
self.norm3 = LayerNorm32(channels, elementwise_affine=ln_affine, eps=1e-6)
|
| 116 |
+
self.context_norm = LayerNorm32(
|
| 117 |
+
ctx_channels, elementwise_affine=ln_affine, eps=1e-6
|
| 118 |
+
)
|
| 119 |
+
self.self_attn = SparseMultiHeadAttention(
|
| 120 |
+
channels,
|
| 121 |
+
num_heads=num_heads,
|
| 122 |
+
type="self",
|
| 123 |
+
attn_mode=attn_mode,
|
| 124 |
+
window_size=window_size,
|
| 125 |
+
shift_sequence=shift_sequence,
|
| 126 |
+
shift_window=shift_window,
|
| 127 |
+
serialize_mode=serialize_mode,
|
| 128 |
+
qkv_bias=qkv_bias,
|
| 129 |
+
use_rope=use_rope,
|
| 130 |
+
qk_rms_norm=qk_rms_norm,
|
| 131 |
+
)
|
| 132 |
+
self.cross_attn = SparseMultiHeadAttention(
|
| 133 |
+
channels,
|
| 134 |
+
ctx_channels=ctx_channels,
|
| 135 |
+
num_heads=num_heads,
|
| 136 |
+
type="cross",
|
| 137 |
+
attn_mode="full",
|
| 138 |
+
qkv_bias=qkv_bias,
|
| 139 |
+
qk_rms_norm=qk_rms_norm_cross,
|
| 140 |
+
)
|
| 141 |
+
self.mlp = SparseFeedForwardNet(
|
| 142 |
+
channels,
|
| 143 |
+
mlp_ratio=mlp_ratio,
|
| 144 |
+
)
|
| 145 |
+
|
| 146 |
+
def _forward(self, x: SparseTensor, context: torch.Tensor):
|
| 147 |
+
h = x.replace(self.norm1(x.feats))
|
| 148 |
+
h = self.self_attn(h)
|
| 149 |
+
x = x + h
|
| 150 |
+
h = x.replace(self.norm2(x.feats))
|
| 151 |
+
|
| 152 |
+
h = self.cross_attn(h, context)
|
| 153 |
+
x = x + h
|
| 154 |
+
h = x.replace(self.norm3(x.feats))
|
| 155 |
+
h = self.mlp(h)
|
| 156 |
+
x = x + h
|
| 157 |
+
return x
|
| 158 |
+
|
| 159 |
+
def forward(self, x: SparseTensor, context: torch.Tensor):
|
| 160 |
+
if self.use_checkpoint:
|
| 161 |
+
return torch.utils.checkpoint.checkpoint(
|
| 162 |
+
self._forward, x, context, use_reentrant=False
|
| 163 |
+
)
|
| 164 |
+
else:
|
| 165 |
+
return self._forward(x, context)
|
modules/sparse/transformer/modulated.py
ADDED
|
@@ -0,0 +1,119 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# MIT License
|
| 2 |
+
|
| 3 |
+
# Copyright (c) Microsoft Corporation.
|
| 4 |
+
# Copyright (c) 2025 VAST-AI-Research and contributors.
|
| 5 |
+
|
| 6 |
+
# Permission is hereby granted, free of charge, to any person obtaining a copy
|
| 7 |
+
# of this software and associated documentation files (the "Software"), to deal
|
| 8 |
+
# in the Software without restriction, including without limitation the rights
|
| 9 |
+
# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
|
| 10 |
+
# copies of the Software, and to permit persons to whom the Software is
|
| 11 |
+
# furnished to do so, subject to the following conditions:
|
| 12 |
+
|
| 13 |
+
# The above copyright notice and this permission notice shall be included in all
|
| 14 |
+
# copies or substantial portions of the Software.
|
| 15 |
+
|
| 16 |
+
# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
|
| 17 |
+
# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
|
| 18 |
+
# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
|
| 19 |
+
# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
|
| 20 |
+
# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
|
| 21 |
+
# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
|
| 22 |
+
# SOFTWARE
|
| 23 |
+
|
| 24 |
+
from typing import *
|
| 25 |
+
import torch
|
| 26 |
+
import torch.nn as nn
|
| 27 |
+
import torch.utils.checkpoint
|
| 28 |
+
from ..basic import SparseTensor
|
| 29 |
+
from ..attention import SparseMultiHeadAttention, SerializeMode
|
| 30 |
+
from ...norm import LayerNorm32
|
| 31 |
+
from .blocks import SparseFeedForwardNet
|
| 32 |
+
|
| 33 |
+
|
| 34 |
+
class ModulatedSparseTransformerCrossBlock(nn.Module):
|
| 35 |
+
"""
|
| 36 |
+
Sparse Transformer cross-attention block (MSA + MCA + FFN) with adaptive layer norm conditioning.
|
| 37 |
+
"""
|
| 38 |
+
def __init__(
|
| 39 |
+
self,
|
| 40 |
+
channels: int,
|
| 41 |
+
ctx_channels: int,
|
| 42 |
+
num_heads: int,
|
| 43 |
+
mlp_ratio: float = 4.0,
|
| 44 |
+
attn_mode: Literal["full", "shift_window", "shift_sequence", "shift_order", "swin"] = "full",
|
| 45 |
+
window_size: Optional[int] = None,
|
| 46 |
+
shift_sequence: Optional[int] = None,
|
| 47 |
+
shift_window: Optional[Tuple[int, int, int]] = None,
|
| 48 |
+
serialize_mode: Optional[SerializeMode] = None,
|
| 49 |
+
use_checkpoint: bool = False,
|
| 50 |
+
use_rope: bool = False,
|
| 51 |
+
qk_rms_norm: bool = False,
|
| 52 |
+
qk_rms_norm_cross: bool = False,
|
| 53 |
+
qkv_bias: bool = True,
|
| 54 |
+
share_mod: bool = False,
|
| 55 |
+
|
| 56 |
+
):
|
| 57 |
+
super().__init__()
|
| 58 |
+
self.use_checkpoint = use_checkpoint
|
| 59 |
+
self.share_mod = share_mod
|
| 60 |
+
self.norm1 = LayerNorm32(channels, elementwise_affine=False, eps=1e-6)
|
| 61 |
+
self.norm2 = LayerNorm32(channels, elementwise_affine=True, eps=1e-6)
|
| 62 |
+
self.norm3 = LayerNorm32(channels, elementwise_affine=False, eps=1e-6)
|
| 63 |
+
self.self_attn = SparseMultiHeadAttention(
|
| 64 |
+
channels,
|
| 65 |
+
num_heads=num_heads,
|
| 66 |
+
type="self",
|
| 67 |
+
attn_mode=attn_mode,
|
| 68 |
+
window_size=window_size,
|
| 69 |
+
shift_sequence=shift_sequence,
|
| 70 |
+
shift_window=shift_window,
|
| 71 |
+
serialize_mode=serialize_mode,
|
| 72 |
+
qkv_bias=qkv_bias,
|
| 73 |
+
use_rope=use_rope,
|
| 74 |
+
qk_rms_norm=qk_rms_norm,
|
| 75 |
+
)
|
| 76 |
+
self.cross_attn = SparseMultiHeadAttention(
|
| 77 |
+
channels,
|
| 78 |
+
ctx_channels=ctx_channels,
|
| 79 |
+
num_heads=num_heads,
|
| 80 |
+
type="cross",
|
| 81 |
+
attn_mode="full",
|
| 82 |
+
qkv_bias=qkv_bias,
|
| 83 |
+
qk_rms_norm=qk_rms_norm_cross,
|
| 84 |
+
)
|
| 85 |
+
self.mlp = SparseFeedForwardNet(
|
| 86 |
+
channels,
|
| 87 |
+
mlp_ratio=mlp_ratio,
|
| 88 |
+
)
|
| 89 |
+
if not share_mod:
|
| 90 |
+
self.adaLN_modulation = nn.Sequential(
|
| 91 |
+
nn.SiLU(),
|
| 92 |
+
nn.Linear(channels, 6 * channels, bias=True)
|
| 93 |
+
)
|
| 94 |
+
|
| 95 |
+
def _forward(self, x: SparseTensor, mod: torch.Tensor, context: torch.Tensor) -> SparseTensor:
|
| 96 |
+
if self.share_mod:
|
| 97 |
+
shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = mod.chunk(6, dim=1)
|
| 98 |
+
else:
|
| 99 |
+
shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = self.adaLN_modulation(mod).chunk(6, dim=1)
|
| 100 |
+
h = x.replace(self.norm1(x.feats))
|
| 101 |
+
h = h * (1 + scale_msa) + shift_msa
|
| 102 |
+
h = self.self_attn(h)
|
| 103 |
+
h = h * gate_msa
|
| 104 |
+
x = x + h
|
| 105 |
+
h = x.replace(self.norm2(x.feats))
|
| 106 |
+
h = self.cross_attn(h, context)
|
| 107 |
+
x = x + h
|
| 108 |
+
h = x.replace(self.norm3(x.feats))
|
| 109 |
+
h = h * (1 + scale_mlp) + shift_mlp
|
| 110 |
+
h = self.mlp(h)
|
| 111 |
+
h = h * gate_mlp
|
| 112 |
+
x = x + h
|
| 113 |
+
return x
|
| 114 |
+
|
| 115 |
+
def forward(self, x: SparseTensor, mod: torch.Tensor, context: torch.Tensor) -> SparseTensor:
|
| 116 |
+
if self.use_checkpoint:
|
| 117 |
+
return torch.utils.checkpoint.checkpoint(self._forward, x, mod, context, use_reentrant=False)
|
| 118 |
+
else:
|
| 119 |
+
return self._forward(x, mod, context)
|
modules/transformer/__init__.py
ADDED
|
@@ -0,0 +1,24 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# MIT License
|
| 2 |
+
|
| 3 |
+
# Copyright (c) Microsoft Corporation.
|
| 4 |
+
# Copyright (c) 2025 VAST-AI-Research and contributors.
|
| 5 |
+
|
| 6 |
+
# Permission is hereby granted, free of charge, to any person obtaining a copy
|
| 7 |
+
# of this software and associated documentation files (the "Software"), to deal
|
| 8 |
+
# in the Software without restriction, including without limitation the rights
|
| 9 |
+
# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
|
| 10 |
+
# copies of the Software, and to permit persons to whom the Software is
|
| 11 |
+
# furnished to do so, subject to the following conditions:
|
| 12 |
+
|
| 13 |
+
# The above copyright notice and this permission notice shall be included in all
|
| 14 |
+
# copies or substantial portions of the Software.
|
| 15 |
+
|
| 16 |
+
# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
|
| 17 |
+
# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
|
| 18 |
+
# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
|
| 19 |
+
# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
|
| 20 |
+
# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
|
| 21 |
+
# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
|
| 22 |
+
# SOFTWARE
|
| 23 |
+
|
| 24 |
+
from .blocks import *
|
modules/transformer/blocks.py
ADDED
|
@@ -0,0 +1,276 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# MIT License
|
| 2 |
+
|
| 3 |
+
# Copyright (c) Microsoft Corporation.
|
| 4 |
+
# Copyright (c) 2025 VAST-AI-Research and contributors.
|
| 5 |
+
|
| 6 |
+
# Permission is hereby granted, free of charge, to any person obtaining a copy
|
| 7 |
+
# of this software and associated documentation files (the "Software"), to deal
|
| 8 |
+
# in the Software without restriction, including without limitation the rights
|
| 9 |
+
# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
|
| 10 |
+
# copies of the Software, and to permit persons to whom the Software is
|
| 11 |
+
# furnished to do so, subject to the following conditions:
|
| 12 |
+
|
| 13 |
+
# The above copyright notice and this permission notice shall be included in all
|
| 14 |
+
# copies or substantial portions of the Software.
|
| 15 |
+
|
| 16 |
+
# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
|
| 17 |
+
# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
|
| 18 |
+
# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
|
| 19 |
+
# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
|
| 20 |
+
# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
|
| 21 |
+
# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
|
| 22 |
+
# SOFTWARE
|
| 23 |
+
|
| 24 |
+
from typing import *
|
| 25 |
+
import numpy as np
|
| 26 |
+
import torch
|
| 27 |
+
import torch.nn as nn
|
| 28 |
+
import torch.nn.functional as F
|
| 29 |
+
|
| 30 |
+
|
| 31 |
+
class AbsolutePositionEmbedder(nn.Module):
|
| 32 |
+
"""
|
| 33 |
+
Embeds spatial positions into vector representations.
|
| 34 |
+
"""
|
| 35 |
+
|
| 36 |
+
def __init__(self, channels: int, in_channels: int = 3):
|
| 37 |
+
super().__init__()
|
| 38 |
+
self.channels = channels
|
| 39 |
+
self.in_channels = in_channels
|
| 40 |
+
self.freq_dim = channels // in_channels // 2
|
| 41 |
+
self.freqs = torch.arange(self.freq_dim, dtype=torch.float32) / self.freq_dim
|
| 42 |
+
self.freqs = 1.0 / (10000**self.freqs)
|
| 43 |
+
|
| 44 |
+
def _sin_cos_embedding(self, x: torch.Tensor) -> torch.Tensor:
|
| 45 |
+
"""
|
| 46 |
+
Create sinusoidal position embeddings.
|
| 47 |
+
|
| 48 |
+
Args:
|
| 49 |
+
x: a 1-D Tensor of N indices
|
| 50 |
+
|
| 51 |
+
Returns:
|
| 52 |
+
an (N, D) Tensor of positional embeddings.
|
| 53 |
+
"""
|
| 54 |
+
self.freqs = self.freqs.to(x.device)
|
| 55 |
+
out = torch.outer(x, self.freqs)
|
| 56 |
+
out = torch.cat([torch.sin(out), torch.cos(out)], dim=-1)
|
| 57 |
+
return out
|
| 58 |
+
|
| 59 |
+
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
| 60 |
+
"""
|
| 61 |
+
Args:
|
| 62 |
+
x (torch.Tensor): (N, D) tensor of spatial positions
|
| 63 |
+
"""
|
| 64 |
+
N, D = x.shape
|
| 65 |
+
assert (
|
| 66 |
+
D == self.in_channels
|
| 67 |
+
), "Input dimension must match number of input channels"
|
| 68 |
+
embed = self._sin_cos_embedding(x.reshape(-1))
|
| 69 |
+
embed = embed.reshape(N, -1)
|
| 70 |
+
if embed.shape[1] < self.channels:
|
| 71 |
+
embed = torch.cat(
|
| 72 |
+
[
|
| 73 |
+
embed,
|
| 74 |
+
torch.zeros(N, self.channels - embed.shape[1], device=embed.device),
|
| 75 |
+
],
|
| 76 |
+
dim=-1,
|
| 77 |
+
)
|
| 78 |
+
return embed
|
| 79 |
+
|
| 80 |
+
|
| 81 |
+
class RotaryPositionPhasesEmbedder(nn.Module):
|
| 82 |
+
def __init__(
|
| 83 |
+
self,
|
| 84 |
+
head_dim: int,
|
| 85 |
+
dim: int = 3,
|
| 86 |
+
rope_freq: Tuple[float, float] = (1.0, 10000.0),
|
| 87 |
+
):
|
| 88 |
+
super().__init__()
|
| 89 |
+
assert head_dim % 2 == 0, "Head dim must be divisible by 2"
|
| 90 |
+
self.head_dim = head_dim
|
| 91 |
+
self.dim = dim
|
| 92 |
+
self.rope_freq = rope_freq
|
| 93 |
+
self.freq_dim = head_dim // 2 // dim
|
| 94 |
+
self.freqs = torch.arange(self.freq_dim, dtype=torch.float32) / self.freq_dim
|
| 95 |
+
self.freqs = rope_freq[0] / (rope_freq[1] ** (self.freqs))
|
| 96 |
+
|
| 97 |
+
def _get_phases(self, indices: torch.Tensor) -> torch.Tensor:
|
| 98 |
+
self.freqs = self.freqs.to(indices.device)
|
| 99 |
+
phases = torch.outer(indices, self.freqs)
|
| 100 |
+
phases = torch.polar(torch.ones_like(phases), phases)
|
| 101 |
+
return phases
|
| 102 |
+
|
| 103 |
+
@staticmethod
|
| 104 |
+
def apply_rotary_embedding(x: torch.Tensor, phases: torch.Tensor) -> torch.Tensor:
|
| 105 |
+
x_complex = torch.view_as_complex(x.float().reshape(*x.shape[:-1], -1, 2))
|
| 106 |
+
if phases.ndim == 3:
|
| 107 |
+
phases = phases.unsqueeze(1)
|
| 108 |
+
x_rotated = x_complex * phases
|
| 109 |
+
x_embed = (
|
| 110 |
+
torch.view_as_real(x_rotated).reshape(*x_rotated.shape[:-1], -1).to(x.dtype)
|
| 111 |
+
)
|
| 112 |
+
return x_embed
|
| 113 |
+
|
| 114 |
+
def forward(self, indices: torch.Tensor) -> torch.Tensor:
|
| 115 |
+
assert indices.shape[-1] == self.dim, f"Last dim of indices must be {self.dim}"
|
| 116 |
+
phases = self._get_phases(indices.reshape(-1)).reshape(*indices.shape[:-1], -1)
|
| 117 |
+
if phases.shape[-1] < self.head_dim // 2:
|
| 118 |
+
padn = self.head_dim // 2 - phases.shape[-1]
|
| 119 |
+
phases = torch.cat(
|
| 120 |
+
[
|
| 121 |
+
phases,
|
| 122 |
+
torch.polar(
|
| 123 |
+
torch.ones(*phases.shape[:-1], padn, device=phases.device),
|
| 124 |
+
torch.zeros(*phases.shape[:-1], padn, device=phases.device),
|
| 125 |
+
),
|
| 126 |
+
],
|
| 127 |
+
dim=-1,
|
| 128 |
+
)
|
| 129 |
+
return phases
|
| 130 |
+
|
| 131 |
+
|
| 132 |
+
class TimestepEmbedder(nn.Module):
|
| 133 |
+
"""
|
| 134 |
+
Embeds scalar timesteps into vector representations.
|
| 135 |
+
"""
|
| 136 |
+
|
| 137 |
+
def __init__(self, hidden_size, frequency_embedding_size=256):
|
| 138 |
+
super().__init__()
|
| 139 |
+
self.mlp = nn.Sequential(
|
| 140 |
+
nn.Linear(frequency_embedding_size, hidden_size, bias=True),
|
| 141 |
+
nn.SiLU(),
|
| 142 |
+
nn.Linear(hidden_size, hidden_size, bias=True),
|
| 143 |
+
)
|
| 144 |
+
self.frequency_embedding_size = frequency_embedding_size
|
| 145 |
+
|
| 146 |
+
@staticmethod
|
| 147 |
+
def timestep_embedding(t, dim, max_period=10000):
|
| 148 |
+
"""
|
| 149 |
+
Create sinusoidal timestep embeddings.
|
| 150 |
+
|
| 151 |
+
Args:
|
| 152 |
+
t: a 1-D Tensor of N indices, one per batch element.
|
| 153 |
+
These may be fractional.
|
| 154 |
+
dim: the dimension of the output.
|
| 155 |
+
max_period: controls the minimum frequency of the embeddings.
|
| 156 |
+
|
| 157 |
+
Returns:
|
| 158 |
+
an (N, D) Tensor of positional embeddings.
|
| 159 |
+
"""
|
| 160 |
+
# https://github.com/openai/glide-text2im/blob/main/glide_text2im/nn.py
|
| 161 |
+
half = dim // 2
|
| 162 |
+
freqs = torch.exp(
|
| 163 |
+
-np.log(max_period)
|
| 164 |
+
* torch.arange(start=0, end=half, dtype=torch.float32)
|
| 165 |
+
/ half
|
| 166 |
+
).to(device=t.device)
|
| 167 |
+
args = t[:, None].float() * freqs[None]
|
| 168 |
+
embedding = torch.cat([torch.cos(args), torch.sin(args)], dim=-1)
|
| 169 |
+
if dim % 2:
|
| 170 |
+
embedding = torch.cat(
|
| 171 |
+
[embedding, torch.zeros_like(embedding[:, :1])], dim=-1
|
| 172 |
+
)
|
| 173 |
+
return embedding
|
| 174 |
+
|
| 175 |
+
def forward(self, t):
|
| 176 |
+
t_freq = self.timestep_embedding(t, self.frequency_embedding_size)
|
| 177 |
+
t_emb = self.mlp(t_freq)
|
| 178 |
+
return t_emb
|
| 179 |
+
|
| 180 |
+
|
| 181 |
+
class PointEmbed(nn.Module):
|
| 182 |
+
def __init__(self, hidden_dim=48, dim=128):
|
| 183 |
+
super().__init__()
|
| 184 |
+
|
| 185 |
+
assert hidden_dim % 6 == 0
|
| 186 |
+
|
| 187 |
+
self.embedding_dim = hidden_dim
|
| 188 |
+
e = torch.pow(2, torch.arange(self.embedding_dim // 6)).float() * np.pi
|
| 189 |
+
e = torch.stack(
|
| 190 |
+
[
|
| 191 |
+
torch.cat(
|
| 192 |
+
[
|
| 193 |
+
e,
|
| 194 |
+
torch.zeros(self.embedding_dim // 6),
|
| 195 |
+
torch.zeros(self.embedding_dim // 6),
|
| 196 |
+
]
|
| 197 |
+
),
|
| 198 |
+
torch.cat(
|
| 199 |
+
[
|
| 200 |
+
torch.zeros(self.embedding_dim // 6),
|
| 201 |
+
e,
|
| 202 |
+
torch.zeros(self.embedding_dim // 6),
|
| 203 |
+
]
|
| 204 |
+
),
|
| 205 |
+
torch.cat(
|
| 206 |
+
[
|
| 207 |
+
torch.zeros(self.embedding_dim // 6),
|
| 208 |
+
torch.zeros(self.embedding_dim // 6),
|
| 209 |
+
e,
|
| 210 |
+
]
|
| 211 |
+
),
|
| 212 |
+
]
|
| 213 |
+
)
|
| 214 |
+
self.register_buffer("basis", e) # 3 x 16
|
| 215 |
+
|
| 216 |
+
self.mlp = nn.Linear(self.embedding_dim + 3, dim)
|
| 217 |
+
|
| 218 |
+
@staticmethod
|
| 219 |
+
def embed(input, basis):
|
| 220 |
+
projections = torch.einsum("bnd,de->bne", input, basis)
|
| 221 |
+
embeddings = torch.cat([projections.sin(), projections.cos()], dim=2)
|
| 222 |
+
return embeddings
|
| 223 |
+
|
| 224 |
+
def forward(self, input):
|
| 225 |
+
dt = self.mlp.weight.dtype
|
| 226 |
+
if input.dtype != dt:
|
| 227 |
+
input = input.to(dtype=dt)
|
| 228 |
+
basis = self.basis.to(dtype=dt)
|
| 229 |
+
embed = self.mlp(torch.cat([self.embed(input, basis), input], dim=2))
|
| 230 |
+
return embed
|
| 231 |
+
|
| 232 |
+
|
| 233 |
+
class MaskedTransformerCrossAttnBlock(nn.Module):
|
| 234 |
+
def __init__(self, hidden_size: int, num_heads: int, cond_dim: int):
|
| 235 |
+
super().__init__()
|
| 236 |
+
self.hidden_size = hidden_size
|
| 237 |
+
self.num_heads = num_heads
|
| 238 |
+
self.head_dim = hidden_size // num_heads
|
| 239 |
+
|
| 240 |
+
self.norm = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6)
|
| 241 |
+
self.q_cross = nn.Linear(hidden_size, hidden_size, bias=True)
|
| 242 |
+
self.kv_cross = nn.Linear(cond_dim, hidden_size * 2, bias=True)
|
| 243 |
+
self.proj_out_cross = nn.Linear(hidden_size, hidden_size, bias=True)
|
| 244 |
+
self.scale_cross = nn.Parameter(torch.zeros(hidden_size))
|
| 245 |
+
|
| 246 |
+
def forward(
|
| 247 |
+
self,
|
| 248 |
+
x: torch.Tensor,
|
| 249 |
+
c_tokens: torch.Tensor,
|
| 250 |
+
x_mask: Optional[torch.Tensor] = None,
|
| 251 |
+
c_mask: Optional[torch.Tensor] = None,
|
| 252 |
+
) -> torch.Tensor:
|
| 253 |
+
b, n, d = x.shape
|
| 254 |
+
q_c = (
|
| 255 |
+
self.q_cross(self.norm(x))
|
| 256 |
+
.view(b, n, self.num_heads, self.head_dim)
|
| 257 |
+
.transpose(1, 2)
|
| 258 |
+
)
|
| 259 |
+
kv_c = (
|
| 260 |
+
self.kv_cross(c_tokens)
|
| 261 |
+
.view(b, c_tokens.shape[1], 2, self.num_heads, self.head_dim)
|
| 262 |
+
.permute(2, 0, 3, 1, 4)
|
| 263 |
+
)
|
| 264 |
+
k_c, v_c = kv_c[0], kv_c[1]
|
| 265 |
+
cross_attn_mask = c_mask.view(b, 1, 1, -1) if c_mask is not None else None
|
| 266 |
+
cross_out = F.scaled_dot_product_attention(
|
| 267 |
+
q_c,
|
| 268 |
+
k_c,
|
| 269 |
+
v_c,
|
| 270 |
+
attn_mask=cross_attn_mask,
|
| 271 |
+
)
|
| 272 |
+
cross_out = cross_out.transpose(1, 2).reshape(b, n, d)
|
| 273 |
+
x = x + self.scale_cross * self.proj_out_cross(cross_out)
|
| 274 |
+
if x_mask is not None:
|
| 275 |
+
x = torch.where(x_mask.unsqueeze(-1), x, torch.zeros_like(x))
|
| 276 |
+
return x
|
modules/transformer/hybrid.py
ADDED
|
@@ -0,0 +1,236 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from __future__ import annotations
|
| 2 |
+
|
| 3 |
+
from typing import Optional
|
| 4 |
+
import torch
|
| 5 |
+
import torch.nn as nn
|
| 6 |
+
import torch.nn.functional as F
|
| 7 |
+
from torch.utils.checkpoint import checkpoint
|
| 8 |
+
|
| 9 |
+
from ..attention import (
|
| 10 |
+
can_flash_varlen,
|
| 11 |
+
flash_varlen_self_attention,
|
| 12 |
+
graph_adj_varlen_attention,
|
| 13 |
+
sdpa_padding_mask,
|
| 14 |
+
)
|
| 15 |
+
from .blocks import RotaryPositionPhasesEmbedder
|
| 16 |
+
|
| 17 |
+
|
| 18 |
+
class GraphAttnVarlenBlock(nn.Module):
|
| 19 |
+
def __init__(
|
| 20 |
+
self, hidden_size: int, num_heads: int, gradient_checkpointing: bool = False
|
| 21 |
+
):
|
| 22 |
+
super().__init__()
|
| 23 |
+
self.hidden_size = hidden_size
|
| 24 |
+
self.num_heads = num_heads
|
| 25 |
+
self.head_dim = hidden_size // num_heads
|
| 26 |
+
|
| 27 |
+
self.norm1 = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6)
|
| 28 |
+
self.norm2 = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6)
|
| 29 |
+
self.qkv = nn.Linear(hidden_size, hidden_size * 3, bias=True)
|
| 30 |
+
self.proj_out = nn.Linear(hidden_size, hidden_size, bias=True)
|
| 31 |
+
self.ffn = nn.Sequential(
|
| 32 |
+
nn.Linear(hidden_size, hidden_size * 4, bias=True),
|
| 33 |
+
nn.GELU(approximate="tanh"),
|
| 34 |
+
nn.Linear(hidden_size * 4, hidden_size, bias=True),
|
| 35 |
+
)
|
| 36 |
+
self.scale_msa = nn.Parameter(torch.zeros(hidden_size))
|
| 37 |
+
self.scale_mlp = nn.Parameter(torch.zeros(hidden_size))
|
| 38 |
+
self.gradient_checkpointing = bool(gradient_checkpointing)
|
| 39 |
+
|
| 40 |
+
def _forward_once(
|
| 41 |
+
self,
|
| 42 |
+
x: torch.Tensor,
|
| 43 |
+
x_mask: Optional[torch.Tensor],
|
| 44 |
+
adj_matrix: Optional[torch.Tensor],
|
| 45 |
+
rope_phases: Optional[torch.Tensor],
|
| 46 |
+
) -> torch.Tensor:
|
| 47 |
+
B, N, D = x.shape
|
| 48 |
+
qkv = (
|
| 49 |
+
self.qkv(self.norm1(x))
|
| 50 |
+
.view(B, N, 3, self.num_heads, self.head_dim)
|
| 51 |
+
.permute(2, 0, 3, 1, 4)
|
| 52 |
+
)
|
| 53 |
+
q, k, v = qkv[0], qkv[1], qkv[2]
|
| 54 |
+
|
| 55 |
+
if rope_phases is not None:
|
| 56 |
+
q = RotaryPositionPhasesEmbedder.apply_rotary_embedding(q, rope_phases)
|
| 57 |
+
k = RotaryPositionPhasesEmbedder.apply_rotary_embedding(k, rope_phases)
|
| 58 |
+
|
| 59 |
+
if x_mask is None:
|
| 60 |
+
attn_mask = None
|
| 61 |
+
if adj_matrix is not None:
|
| 62 |
+
adj_mask = adj_matrix.bool()
|
| 63 |
+
eye = torch.eye(N, dtype=torch.bool, device=x.device).unsqueeze(0)
|
| 64 |
+
attn_mask = (adj_mask | eye).unsqueeze(1)
|
| 65 |
+
attn_out = F.scaled_dot_product_attention(q, k, v, attn_mask=attn_mask)
|
| 66 |
+
else:
|
| 67 |
+
attn_out = graph_adj_varlen_attention(q, k, v, x_mask, adj_matrix)
|
| 68 |
+
|
| 69 |
+
attn_out = attn_out.transpose(1, 2).reshape(B, N, D)
|
| 70 |
+
x = x + self.scale_msa * self.proj_out(attn_out)
|
| 71 |
+
x = x + self.scale_mlp * self.ffn(self.norm2(x))
|
| 72 |
+
if x_mask is not None:
|
| 73 |
+
x = torch.where(x_mask.unsqueeze(-1), x, torch.zeros_like(x))
|
| 74 |
+
return x
|
| 75 |
+
|
| 76 |
+
def forward(
|
| 77 |
+
self,
|
| 78 |
+
x: torch.Tensor,
|
| 79 |
+
x_mask: Optional[torch.Tensor],
|
| 80 |
+
adj_matrix: Optional[torch.Tensor],
|
| 81 |
+
rope_phases: Optional[torch.Tensor] = None,
|
| 82 |
+
) -> torch.Tensor:
|
| 83 |
+
if self.training and self.gradient_checkpointing:
|
| 84 |
+
return checkpoint(
|
| 85 |
+
self._forward_once,
|
| 86 |
+
x,
|
| 87 |
+
x_mask,
|
| 88 |
+
adj_matrix,
|
| 89 |
+
rope_phases,
|
| 90 |
+
use_reentrant=False,
|
| 91 |
+
)
|
| 92 |
+
return self._forward_once(x, x_mask, adj_matrix, rope_phases)
|
| 93 |
+
|
| 94 |
+
|
| 95 |
+
class FlashVarlenTransformerBlock(nn.Module):
|
| 96 |
+
def __init__(
|
| 97 |
+
self, hidden_size: int, num_heads: int, gradient_checkpointing: bool = False
|
| 98 |
+
):
|
| 99 |
+
super().__init__()
|
| 100 |
+
self.hidden_size = hidden_size
|
| 101 |
+
self.num_heads = num_heads
|
| 102 |
+
self.head_dim = hidden_size // num_heads
|
| 103 |
+
|
| 104 |
+
self.norm1 = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6)
|
| 105 |
+
self.norm2 = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6)
|
| 106 |
+
self.qkv = nn.Linear(hidden_size, hidden_size * 3, bias=True)
|
| 107 |
+
self.proj_out = nn.Linear(hidden_size, hidden_size, bias=True)
|
| 108 |
+
self.ffn = nn.Sequential(
|
| 109 |
+
nn.Linear(hidden_size, hidden_size * 4, bias=True),
|
| 110 |
+
nn.GELU(approximate="tanh"),
|
| 111 |
+
nn.Linear(hidden_size * 4, hidden_size, bias=True),
|
| 112 |
+
)
|
| 113 |
+
self.scale_msa = nn.Parameter(torch.zeros(hidden_size))
|
| 114 |
+
self.scale_mlp = nn.Parameter(torch.zeros(hidden_size))
|
| 115 |
+
self.gradient_checkpointing = bool(gradient_checkpointing)
|
| 116 |
+
|
| 117 |
+
def _forward_once(
|
| 118 |
+
self,
|
| 119 |
+
x: torch.Tensor,
|
| 120 |
+
x_mask: Optional[torch.Tensor],
|
| 121 |
+
rope_phases: Optional[torch.Tensor],
|
| 122 |
+
) -> torch.Tensor:
|
| 123 |
+
B, N, D = x.shape
|
| 124 |
+
qkv = (
|
| 125 |
+
self.qkv(self.norm1(x))
|
| 126 |
+
.view(B, N, 3, self.num_heads, self.head_dim)
|
| 127 |
+
.permute(2, 0, 3, 1, 4)
|
| 128 |
+
)
|
| 129 |
+
q, k, v = qkv[0], qkv[1], qkv[2]
|
| 130 |
+
|
| 131 |
+
if rope_phases is not None:
|
| 132 |
+
q = RotaryPositionPhasesEmbedder.apply_rotary_embedding(q, rope_phases)
|
| 133 |
+
k = RotaryPositionPhasesEmbedder.apply_rotary_embedding(k, rope_phases)
|
| 134 |
+
|
| 135 |
+
if can_flash_varlen(q, x_mask):
|
| 136 |
+
attn_out = flash_varlen_self_attention(q, k, v, x_mask)
|
| 137 |
+
elif x_mask is not None:
|
| 138 |
+
pad_mask = sdpa_padding_mask(x_mask)
|
| 139 |
+
attn_out = F.scaled_dot_product_attention(q, k, v, attn_mask=pad_mask)
|
| 140 |
+
else:
|
| 141 |
+
attn_out = F.scaled_dot_product_attention(q, k, v, attn_mask=None)
|
| 142 |
+
|
| 143 |
+
attn_out = attn_out.transpose(1, 2).reshape(B, N, D)
|
| 144 |
+
x = x + self.scale_msa * self.proj_out(attn_out)
|
| 145 |
+
x = x + self.scale_mlp * self.ffn(self.norm2(x))
|
| 146 |
+
if x_mask is not None:
|
| 147 |
+
x = torch.where(x_mask.unsqueeze(-1), x, torch.zeros_like(x))
|
| 148 |
+
return x
|
| 149 |
+
|
| 150 |
+
def forward(
|
| 151 |
+
self,
|
| 152 |
+
x: torch.Tensor,
|
| 153 |
+
x_mask: Optional[torch.Tensor],
|
| 154 |
+
rope_phases: Optional[torch.Tensor] = None,
|
| 155 |
+
) -> torch.Tensor:
|
| 156 |
+
if self.training and self.gradient_checkpointing:
|
| 157 |
+
return checkpoint(
|
| 158 |
+
self._forward_once,
|
| 159 |
+
x,
|
| 160 |
+
x_mask,
|
| 161 |
+
rope_phases,
|
| 162 |
+
use_reentrant=False,
|
| 163 |
+
)
|
| 164 |
+
return self._forward_once(x, x_mask, rope_phases)
|
| 165 |
+
|
| 166 |
+
|
| 167 |
+
class HybridGraphFlashStage(nn.Module):
|
| 168 |
+
def __init__(
|
| 169 |
+
self,
|
| 170 |
+
hidden_size: int,
|
| 171 |
+
num_heads: int,
|
| 172 |
+
num_flash: int,
|
| 173 |
+
gradient_checkpointing: bool = False,
|
| 174 |
+
):
|
| 175 |
+
super().__init__()
|
| 176 |
+
self.graph_block = GraphAttnVarlenBlock(
|
| 177 |
+
hidden_size, num_heads, gradient_checkpointing=gradient_checkpointing
|
| 178 |
+
)
|
| 179 |
+
self.flash_blocks = nn.ModuleList(
|
| 180 |
+
[
|
| 181 |
+
FlashVarlenTransformerBlock(
|
| 182 |
+
hidden_size,
|
| 183 |
+
num_heads,
|
| 184 |
+
gradient_checkpointing=gradient_checkpointing,
|
| 185 |
+
)
|
| 186 |
+
for _ in range(num_flash)
|
| 187 |
+
]
|
| 188 |
+
)
|
| 189 |
+
|
| 190 |
+
def forward(
|
| 191 |
+
self,
|
| 192 |
+
x: torch.Tensor,
|
| 193 |
+
x_mask: Optional[torch.Tensor],
|
| 194 |
+
adj_matrix: Optional[torch.Tensor],
|
| 195 |
+
rope_phases: Optional[torch.Tensor],
|
| 196 |
+
) -> torch.Tensor:
|
| 197 |
+
x = self.graph_block(
|
| 198 |
+
x, x_mask=x_mask, adj_matrix=adj_matrix, rope_phases=rope_phases
|
| 199 |
+
)
|
| 200 |
+
for fb in self.flash_blocks:
|
| 201 |
+
x = fb(x, x_mask=x_mask, rope_phases=rope_phases)
|
| 202 |
+
return x
|
| 203 |
+
|
| 204 |
+
|
| 205 |
+
class HybridGraphFlashStack(nn.Module):
|
| 206 |
+
def __init__(
|
| 207 |
+
self,
|
| 208 |
+
hidden_size: int,
|
| 209 |
+
num_heads: int,
|
| 210 |
+
num_stages: int,
|
| 211 |
+
num_flash_per_stage: int,
|
| 212 |
+
gradient_checkpointing: bool = False,
|
| 213 |
+
):
|
| 214 |
+
super().__init__()
|
| 215 |
+
self.stages = nn.ModuleList(
|
| 216 |
+
[
|
| 217 |
+
HybridGraphFlashStage(
|
| 218 |
+
hidden_size,
|
| 219 |
+
num_heads,
|
| 220 |
+
num_flash_per_stage,
|
| 221 |
+
gradient_checkpointing=gradient_checkpointing,
|
| 222 |
+
)
|
| 223 |
+
for _ in range(num_stages)
|
| 224 |
+
]
|
| 225 |
+
)
|
| 226 |
+
|
| 227 |
+
def forward(
|
| 228 |
+
self,
|
| 229 |
+
x: torch.Tensor,
|
| 230 |
+
x_mask: Optional[torch.Tensor],
|
| 231 |
+
adj_matrix: Optional[torch.Tensor],
|
| 232 |
+
rope_phases: Optional[torch.Tensor],
|
| 233 |
+
) -> torch.Tensor:
|
| 234 |
+
for stage in self.stages:
|
| 235 |
+
x = stage(x, x_mask=x_mask, adj_matrix=adj_matrix, rope_phases=rope_phases)
|
| 236 |
+
return x
|
modules/utils.py
ADDED
|
@@ -0,0 +1,145 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import torch
|
| 2 |
+
import torch.nn as nn
|
| 3 |
+
from typing import *
|
| 4 |
+
import numpy as np
|
| 5 |
+
from modules import sparse as sp
|
| 6 |
+
|
| 7 |
+
FP16_MODULES = (
|
| 8 |
+
nn.Conv1d,
|
| 9 |
+
nn.Conv2d,
|
| 10 |
+
nn.Conv3d,
|
| 11 |
+
nn.ConvTranspose1d,
|
| 12 |
+
nn.ConvTranspose2d,
|
| 13 |
+
nn.ConvTranspose3d,
|
| 14 |
+
nn.Linear,
|
| 15 |
+
sp.SparseConv3d,
|
| 16 |
+
sp.SparseInverseConv3d,
|
| 17 |
+
sp.SparseLinear,
|
| 18 |
+
)
|
| 19 |
+
|
| 20 |
+
|
| 21 |
+
def convert_module_to_f16(l):
|
| 22 |
+
"""
|
| 23 |
+
Convert primitive modules to float16.
|
| 24 |
+
"""
|
| 25 |
+
if isinstance(l, FP16_MODULES):
|
| 26 |
+
for p in l.parameters():
|
| 27 |
+
p.data = p.data.half()
|
| 28 |
+
|
| 29 |
+
|
| 30 |
+
def convert_module_to_f32(l):
|
| 31 |
+
"""
|
| 32 |
+
Convert primitive modules to float32, undoing convert_module_to_f16().
|
| 33 |
+
"""
|
| 34 |
+
if isinstance(l, FP16_MODULES):
|
| 35 |
+
for p in l.parameters():
|
| 36 |
+
p.data = p.data.float()
|
| 37 |
+
|
| 38 |
+
|
| 39 |
+
def zero_module(module):
|
| 40 |
+
"""
|
| 41 |
+
Zero out the parameters of a module and return it.
|
| 42 |
+
"""
|
| 43 |
+
for p in module.parameters():
|
| 44 |
+
p.detach().zero_()
|
| 45 |
+
return module
|
| 46 |
+
|
| 47 |
+
|
| 48 |
+
def modulate(x, shift, scale):
|
| 49 |
+
return x * (1 + scale.unsqueeze(1)) + shift.unsqueeze(1)
|
| 50 |
+
|
| 51 |
+
|
| 52 |
+
class DiagonalGaussianDistribution(object):
|
| 53 |
+
def __init__(
|
| 54 |
+
self,
|
| 55 |
+
parameters: Union[torch.Tensor, List[torch.Tensor]],
|
| 56 |
+
deterministic=False,
|
| 57 |
+
feat_dim=1,
|
| 58 |
+
):
|
| 59 |
+
self.feat_dim = feat_dim
|
| 60 |
+
self.parameters = parameters
|
| 61 |
+
|
| 62 |
+
if isinstance(parameters, list):
|
| 63 |
+
self.mean = parameters[0]
|
| 64 |
+
self.logvar = parameters[1]
|
| 65 |
+
else:
|
| 66 |
+
self.mean, self.logvar = torch.chunk(parameters, 2, dim=feat_dim)
|
| 67 |
+
self.logvar = torch.clamp(self.logvar, -30.0, 20.0)
|
| 68 |
+
self.deterministic = deterministic
|
| 69 |
+
self.std = torch.exp(0.5 * self.logvar)
|
| 70 |
+
self.var = torch.exp(self.logvar)
|
| 71 |
+
if self.deterministic:
|
| 72 |
+
self.var = self.std = torch.zeros_like(self.mean)
|
| 73 |
+
|
| 74 |
+
def sample(self):
|
| 75 |
+
x = self.mean + self.std * torch.randn_like(self.mean)
|
| 76 |
+
return x
|
| 77 |
+
|
| 78 |
+
def kl(self, other=None, dims=(1, 2, 3)):
|
| 79 |
+
if self.deterministic:
|
| 80 |
+
return torch.Tensor([0.0])
|
| 81 |
+
else:
|
| 82 |
+
if other is None:
|
| 83 |
+
return 0.5 * torch.mean(
|
| 84 |
+
torch.pow(self.mean, 2) + self.var - 1.0 - self.logvar, dim=dims
|
| 85 |
+
)
|
| 86 |
+
else:
|
| 87 |
+
return 0.5 * torch.mean(
|
| 88 |
+
torch.pow(self.mean - other.mean, 2) / other.var
|
| 89 |
+
+ self.var / other.var
|
| 90 |
+
- 1.0
|
| 91 |
+
- self.logvar
|
| 92 |
+
+ other.logvar,
|
| 93 |
+
dim=dims,
|
| 94 |
+
)
|
| 95 |
+
|
| 96 |
+
def nll(self, sample, dims=(1, 2, 3)):
|
| 97 |
+
if self.deterministic:
|
| 98 |
+
return torch.Tensor([0.0])
|
| 99 |
+
logtwopi = np.log(2.0 * np.pi)
|
| 100 |
+
return 0.5 * torch.sum(
|
| 101 |
+
logtwopi + self.logvar + torch.pow(sample - self.mean, 2) / self.var,
|
| 102 |
+
dim=dims,
|
| 103 |
+
)
|
| 104 |
+
|
| 105 |
+
def mode(self):
|
| 106 |
+
return self.mean
|
| 107 |
+
|
| 108 |
+
|
| 109 |
+
def per_batch_counts(batch_indices: torch.Tensor, num_batches: int) -> List[int]:
|
| 110 |
+
"""Count elements per batch, returned as a list of length num_batches."""
|
| 111 |
+
return torch.bincount(batch_indices.long(), minlength=num_batches).tolist()
|
| 112 |
+
|
| 113 |
+
|
| 114 |
+
def flatten_coords(coords_4d: torch.Tensor):
|
| 115 |
+
coords_4d_long = coords_4d.long()
|
| 116 |
+
|
| 117 |
+
base_x = 1024
|
| 118 |
+
base_y = 1024 * 1024
|
| 119 |
+
base_z = 1024 * 1024 * 1024
|
| 120 |
+
|
| 121 |
+
flat_coords = (
|
| 122 |
+
coords_4d_long[:, 0] * base_z
|
| 123 |
+
+ coords_4d_long[:, 1] * base_y
|
| 124 |
+
+ coords_4d_long[:, 2] * base_x
|
| 125 |
+
+ coords_4d_long[:, 3]
|
| 126 |
+
)
|
| 127 |
+
return flat_coords
|
| 128 |
+
|
| 129 |
+
def manual_cast(tensor, dtype):
|
| 130 |
+
if not torch.is_autocast_enabled():
|
| 131 |
+
return tensor.type(dtype)
|
| 132 |
+
return tensor
|
| 133 |
+
|
| 134 |
+
|
| 135 |
+
def str_to_dtype(dtype_str: str):
|
| 136 |
+
return {
|
| 137 |
+
"f16": torch.float16,
|
| 138 |
+
"fp16": torch.float16,
|
| 139 |
+
"float16": torch.float16,
|
| 140 |
+
"bf16": torch.bfloat16,
|
| 141 |
+
"bfloat16": torch.bfloat16,
|
| 142 |
+
"f32": torch.float32,
|
| 143 |
+
"fp32": torch.float32,
|
| 144 |
+
"float32": torch.float32,
|
| 145 |
+
}[dtype_str]
|
requirements.txt
ADDED
|
@@ -0,0 +1,23 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# ──────────────────────────────────────────────────────────────────────
|
| 2 |
+
# LATO.2 Gradio App — requirements.txt
|
| 3 |
+
# ──────────────────────────────────────────────────────────────────────
|
| 4 |
+
# Install AFTER running `setup.sh --all` (which sets up the base
|
| 5 |
+
# conda env with PyTorch, spconv, flash-attn, o_voxel, etc.)
|
| 6 |
+
#
|
| 7 |
+
# pip install -r requirements.txt
|
| 8 |
+
# ──────────────────────────────────────────────────────────────────────
|
| 9 |
+
|
| 10 |
+
# Gradio app framework
|
| 11 |
+
gradio>=4.44.0
|
| 12 |
+
gradio_rerun>=0.0.4
|
| 13 |
+
|
| 14 |
+
# Rerun 3D viewer SDK
|
| 15 |
+
rerun-sdk>=0.22.0
|
| 16 |
+
|
| 17 |
+
# ── Already in setup.sh but listed for completeness ──────────────────
|
| 18 |
+
numpy
|
| 19 |
+
trimesh
|
| 20 |
+
tqdm
|
| 21 |
+
pillow
|
| 22 |
+
huggingface_hub
|
| 23 |
+
open3d==0.19.0
|
scripts/ckpt_download.py
ADDED
|
@@ -0,0 +1,90 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""
|
| 2 |
+
Usage:
|
| 3 |
+
python scripts/ckpt_download.py \
|
| 4 |
+
[--out_dir <path>]
|
| 5 |
+
"""
|
| 6 |
+
|
| 7 |
+
import argparse
|
| 8 |
+
import os
|
| 9 |
+
import sys
|
| 10 |
+
|
| 11 |
+
ROOT = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
|
| 12 |
+
sys.path.insert(0, ROOT)
|
| 13 |
+
|
| 14 |
+
from utils import logging
|
| 15 |
+
|
| 16 |
+
DEFAULT_REPO_ID = "0x4c48/LATO.2"
|
| 17 |
+
DEFAULT_OUT_DIR = os.path.join(ROOT, "ckpt")
|
| 18 |
+
|
| 19 |
+
|
| 20 |
+
def parse_args():
|
| 21 |
+
p = argparse.ArgumentParser(
|
| 22 |
+
description="Download LATO.2 checkpoints from the Hugging Face Hub.",
|
| 23 |
+
formatter_class=argparse.ArgumentDefaultsHelpFormatter,
|
| 24 |
+
)
|
| 25 |
+
p.add_argument("--repo_id", default=DEFAULT_REPO_ID, help="HF model repo id")
|
| 26 |
+
p.add_argument(
|
| 27 |
+
"--out_dir",
|
| 28 |
+
default=DEFAULT_OUT_DIR,
|
| 29 |
+
help="Directory to download checkpoints into",
|
| 30 |
+
)
|
| 31 |
+
p.add_argument("--revision", default=None, help="Git revision / branch / tag")
|
| 32 |
+
p.add_argument(
|
| 33 |
+
"--token",
|
| 34 |
+
default=os.environ.get("HF_TOKEN"),
|
| 35 |
+
help="HF access token for gated/private repos (or set HF_TOKEN)",
|
| 36 |
+
)
|
| 37 |
+
p.add_argument(
|
| 38 |
+
"--all",
|
| 39 |
+
action="store_true",
|
| 40 |
+
help="Download every file in the repo (not just *.pt)",
|
| 41 |
+
)
|
| 42 |
+
p.add_argument(
|
| 43 |
+
"--include-readme",
|
| 44 |
+
action="store_true",
|
| 45 |
+
help="Also download README.md alongside the *.pt weights",
|
| 46 |
+
)
|
| 47 |
+
return p.parse_args()
|
| 48 |
+
|
| 49 |
+
|
| 50 |
+
def main():
|
| 51 |
+
args = parse_args()
|
| 52 |
+
|
| 53 |
+
try:
|
| 54 |
+
from huggingface_hub import snapshot_download
|
| 55 |
+
except ImportError:
|
| 56 |
+
sys.exit(
|
| 57 |
+
"huggingface_hub is not installed. Activate the `trellis2` conda env "
|
| 58 |
+
"or run: pip install -U huggingface_hub"
|
| 59 |
+
)
|
| 60 |
+
|
| 61 |
+
if args.all:
|
| 62 |
+
allow_patterns = None
|
| 63 |
+
else:
|
| 64 |
+
allow_patterns = ["*.pt"]
|
| 65 |
+
if args.include_readme:
|
| 66 |
+
allow_patterns.append("README.md")
|
| 67 |
+
|
| 68 |
+
os.makedirs(args.out_dir, exist_ok=True)
|
| 69 |
+
logging.info(f"Downloading {args.repo_id} -> {args.out_dir}")
|
| 70 |
+
if allow_patterns:
|
| 71 |
+
logging.info(f" patterns: {allow_patterns}")
|
| 72 |
+
|
| 73 |
+
path = snapshot_download(
|
| 74 |
+
repo_id=args.repo_id,
|
| 75 |
+
repo_type="model",
|
| 76 |
+
revision=args.revision,
|
| 77 |
+
local_dir=args.out_dir,
|
| 78 |
+
allow_patterns=allow_patterns,
|
| 79 |
+
token=args.token,
|
| 80 |
+
)
|
| 81 |
+
|
| 82 |
+
logging.info(f"\nDone. Checkpoints available in: {path}")
|
| 83 |
+
files = sorted(f for f in os.listdir(path) if os.path.isfile(os.path.join(path, f)))
|
| 84 |
+
for f in files:
|
| 85 |
+
size = os.path.getsize(os.path.join(path, f)) / (1024 * 1024)
|
| 86 |
+
logging.info(f" {f:24s} {size:8.1f} MB")
|
| 87 |
+
|
| 88 |
+
|
| 89 |
+
if __name__ == "__main__":
|
| 90 |
+
main()
|
scripts/e2e_inference.py
ADDED
|
@@ -0,0 +1,380 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""
|
| 2 |
+
Outputs into --out_dir:
|
| 3 |
+
<mesh_id>_pred.ply generated vertices (+offset head) in [-0.5, 0.5]
|
| 4 |
+
<mesh_id>_pred_coords.ply generated vertex voxel coords in [0, 1024)
|
| 5 |
+
<mesh_id>_pred.obj generated mesh (offset vertices, faces) in [-0.5, 0.5]
|
| 6 |
+
<mesh_id>_pred_coords.obj generated mesh on integer voxel coords in [0, 1024)
|
| 7 |
+
<mesh_id>_render.png the conditioning view fed to DINO-v2
|
| 8 |
+
|
| 9 |
+
Usage:
|
| 10 |
+
python scripts/e2e_inference.py --mesh_dir <dir> --out_dir outputs/e2e_run/<dir> \
|
| 11 |
+
[--vert_num 2000] [--cfg_strength 3.0] [--vflow_steps 24] [--tflow_steps 50] \
|
| 12 |
+
[--render_azimuth 45 --render_elevation 30] [--no-fill_quad_rings]
|
| 13 |
+
"""
|
| 14 |
+
|
| 15 |
+
import argparse
|
| 16 |
+
import os
|
| 17 |
+
import sys
|
| 18 |
+
import time
|
| 19 |
+
from collections import Counter
|
| 20 |
+
from functools import partial
|
| 21 |
+
|
| 22 |
+
import tqdm
|
| 23 |
+
|
| 24 |
+
ROOT = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
|
| 25 |
+
sys.path.insert(0, ROOT)
|
| 26 |
+
os.environ.setdefault("PYTORCH_CUDA_ALLOC_CONF", "expandable_segments:True")
|
| 27 |
+
os.environ.setdefault("XDG_RUNTIME_DIR", "/tmp/runtime-root")
|
| 28 |
+
os.makedirs(os.environ["XDG_RUNTIME_DIR"], exist_ok=True)
|
| 29 |
+
# Open3D headless rendering: without this the default EGL platform can hang
|
| 30 |
+
# (e.g. when every GPU is busy or no display device is exposed).
|
| 31 |
+
os.environ.setdefault("EGL_PLATFORM", "surfaceless")
|
| 32 |
+
|
| 33 |
+
import numpy as np
|
| 34 |
+
import torch
|
| 35 |
+
import trimesh
|
| 36 |
+
from PIL import Image
|
| 37 |
+
from torch.utils.data import DataLoader
|
| 38 |
+
|
| 39 |
+
from dataset.voxel_dataset import VoxelVertexDataset, collate_fn
|
| 40 |
+
from models import (
|
| 41 |
+
DinoV2Encoder,
|
| 42 |
+
OffsetHead,
|
| 43 |
+
TopoFlowEulerSampler,
|
| 44 |
+
TopologySiTFlow,
|
| 45 |
+
TopologyVAE,
|
| 46 |
+
VertexSLatFlowModel,
|
| 47 |
+
VertFlowEulerCfgSampler,
|
| 48 |
+
VertexVAE,
|
| 49 |
+
VoxelFieldConditioner,
|
| 50 |
+
)
|
| 51 |
+
from modules.sparse import SparseTensor
|
| 52 |
+
import utils.logging as logging
|
| 53 |
+
from utils.export import export_vertex
|
| 54 |
+
from utils.inference import (
|
| 55 |
+
build_voxel_fields,
|
| 56 |
+
compute_density,
|
| 57 |
+
decode_vertices,
|
| 58 |
+
edges_to_faces,
|
| 59 |
+
pad_verts,
|
| 60 |
+
worker_init,
|
| 61 |
+
)
|
| 62 |
+
from utils.load import load_latov2_model
|
| 63 |
+
|
| 64 |
+
|
| 65 |
+
def parse_args():
|
| 66 |
+
p = argparse.ArgumentParser(
|
| 67 |
+
description="end-to-end vertex + topology generation inference"
|
| 68 |
+
)
|
| 69 |
+
p.add_argument("--mesh_dir", required=True, help="directory of input meshes")
|
| 70 |
+
p.add_argument(
|
| 71 |
+
"--out_dir", required=True, help="output directory for the PLYs / OBJs"
|
| 72 |
+
)
|
| 73 |
+
p.add_argument("--vflow_ckpt", default=os.path.join(ROOT, "ckpt", "vflow.pt"))
|
| 74 |
+
p.add_argument("--vvae_ckpt", default=os.path.join(ROOT, "ckpt", "vvae.pt"))
|
| 75 |
+
p.add_argument(
|
| 76 |
+
"--offset_head_ckpt", default=os.path.join(ROOT, "ckpt", "offset_head.pt")
|
| 77 |
+
)
|
| 78 |
+
p.add_argument("--tflow_ckpt", default=os.path.join(ROOT, "ckpt", "tflow.pt"))
|
| 79 |
+
p.add_argument("--tvae_ckpt", default=os.path.join(ROOT, "ckpt", "tvae.pt"))
|
| 80 |
+
p.add_argument(
|
| 81 |
+
"--voxel_encoder_ckpt",
|
| 82 |
+
default=os.path.join(ROOT, "ckpt", "voxel_encoder.pt"),
|
| 83 |
+
)
|
| 84 |
+
p.add_argument("--batch_size", type=int, default=1)
|
| 85 |
+
p.add_argument(
|
| 86 |
+
"--num_samples", type=int, default=None, help="only run the first N meshes"
|
| 87 |
+
)
|
| 88 |
+
p.add_argument("--num_workers", type=int, default=4)
|
| 89 |
+
p.add_argument("--inference_threshold", type=float, default=0.5)
|
| 90 |
+
p.add_argument("--seed", type=int, default=42)
|
| 91 |
+
# vertex flow sampling
|
| 92 |
+
p.add_argument("--vflow_steps", type=int, default=24, help="V-Flow Euler steps")
|
| 93 |
+
p.add_argument("--cfg_strength", type=float, default=3.0)
|
| 94 |
+
p.add_argument("--rescale_t", type=float, default=1.0)
|
| 95 |
+
# vertex-count density conditioning
|
| 96 |
+
p.add_argument("--vert_num", type=int, default=2000, help="target vertex count")
|
| 97 |
+
p.add_argument(
|
| 98 |
+
"--use_gt_vert_count",
|
| 99 |
+
action=argparse.BooleanOptionalAction,
|
| 100 |
+
default=False,
|
| 101 |
+
help="condition on the GT quantized vertex count instead of --vert_num",
|
| 102 |
+
)
|
| 103 |
+
p.add_argument(
|
| 104 |
+
"--scaler",
|
| 105 |
+
type=float,
|
| 106 |
+
default=1.0,
|
| 107 |
+
help="multiplier on the GT count when --use_gt_vert_count",
|
| 108 |
+
)
|
| 109 |
+
p.add_argument("--min_verts", type=float, default=200.0)
|
| 110 |
+
p.add_argument("--max_verts", type=float, default=5000.0)
|
| 111 |
+
# topology flow sampling / decoding
|
| 112 |
+
p.add_argument("--tflow_steps", type=int, default=50, help="T-Flow Euler steps")
|
| 113 |
+
p.add_argument("--edge_threshold", type=float, default=0.0)
|
| 114 |
+
p.add_argument("--chunk_size", type=int, default=20000)
|
| 115 |
+
p.add_argument(
|
| 116 |
+
"--fill_quad_rings",
|
| 117 |
+
action=argparse.BooleanOptionalAction,
|
| 118 |
+
default=True,
|
| 119 |
+
help=(
|
| 120 |
+
"post-process: split chordless 4-vertex rings into two triangles "
|
| 121 |
+
"(pure topology, not the voxel support filter)"
|
| 122 |
+
),
|
| 123 |
+
)
|
| 124 |
+
# conditioning render
|
| 125 |
+
p.add_argument("--render_azimuth", type=float, default=45.0)
|
| 126 |
+
p.add_argument("--render_elevation", type=float, default=30.0)
|
| 127 |
+
p.add_argument("--img_res", type=int, default=518)
|
| 128 |
+
p.add_argument(
|
| 129 |
+
"--dino_hub_dir",
|
| 130 |
+
default=os.path.join(ROOT, "ckpt", "dinov2"),
|
| 131 |
+
help="torch.hub cache for DINO-v2; reused when present, downloaded otherwise",
|
| 132 |
+
)
|
| 133 |
+
args = p.parse_args()
|
| 134 |
+
if args.num_samples is not None and args.num_samples <= 0:
|
| 135 |
+
args.num_samples = None # <= 0 means "all", not python slice semantics
|
| 136 |
+
return args
|
| 137 |
+
|
| 138 |
+
|
| 139 |
+
def export_mesh(out_dir, base_name, vert_int, vert_offsets, faces, resolution):
|
| 140 |
+
"""OBJ pair matching export_vertex's PLY conventions (offset verts / int coords)."""
|
| 141 |
+
res = float(resolution)
|
| 142 |
+
vert_with_offset = (
|
| 143 |
+
vert_int.astype(np.float64) / res
|
| 144 |
+
- 0.5
|
| 145 |
+
+ vert_offsets.astype(np.float64) / (res * 2.0)
|
| 146 |
+
)
|
| 147 |
+
trimesh.Trimesh(vertices=vert_with_offset, faces=faces).export(
|
| 148 |
+
os.path.join(out_dir, f"{base_name}_pred.obj")
|
| 149 |
+
)
|
| 150 |
+
trimesh.Trimesh(vertices=vert_int.astype(np.float64), faces=faces).export(
|
| 151 |
+
os.path.join(out_dir, f"{base_name}_pred_coords.obj")
|
| 152 |
+
)
|
| 153 |
+
|
| 154 |
+
|
| 155 |
+
def main():
|
| 156 |
+
logging.info("End-to-end inference starting...")
|
| 157 |
+
|
| 158 |
+
args = parse_args()
|
| 159 |
+
device = torch.device("cuda")
|
| 160 |
+
torch.manual_seed(args.seed)
|
| 161 |
+
np.random.seed(args.seed)
|
| 162 |
+
os.makedirs(args.out_dir, exist_ok=True)
|
| 163 |
+
|
| 164 |
+
# stage 1: vertex generation
|
| 165 |
+
vflow, vflow_cfg = load_latov2_model(VertexSLatFlowModel, args.vflow_ckpt, device)
|
| 166 |
+
vvae, vvae_cfg = load_latov2_model(VertexVAE, args.vvae_ckpt, device)
|
| 167 |
+
offset_head, _ = load_latov2_model(OffsetHead, args.offset_head_ckpt, device)
|
| 168 |
+
# stage 2: topology generation
|
| 169 |
+
tflow, tflow_cfg = load_latov2_model(TopologySiTFlow, args.tflow_ckpt, device)
|
| 170 |
+
tvae, _ = load_latov2_model(TopologyVAE, args.tvae_ckpt, device)
|
| 171 |
+
voxel_encoder, venc_cfg = load_latov2_model(
|
| 172 |
+
VoxelFieldConditioner, args.voxel_encoder_ckpt, device
|
| 173 |
+
)
|
| 174 |
+
|
| 175 |
+
res = vvae_cfg["resolution"]
|
| 176 |
+
min_res = vvae_cfg["min_resolution"]
|
| 177 |
+
latent_dim = vflow_cfg["latent_dim"]
|
| 178 |
+
density_max = vflow_cfg["max_vertex_num"]
|
| 179 |
+
z_dim = int(tflow_cfg["args"]["z_dim"])
|
| 180 |
+
num_discrete = int(tflow_cfg["args"]["num_discrete"])
|
| 181 |
+
max_vertices = int(tflow_cfg["args"]["max_vertices"])
|
| 182 |
+
latent_scale = float(tflow_cfg["latent_scale"])
|
| 183 |
+
voxel_res = int(venc_cfg["resolution"])
|
| 184 |
+
if num_discrete != res:
|
| 185 |
+
raise ValueError(
|
| 186 |
+
f"T-Flow num_discrete={num_discrete} != V-VAE resolution={res}; "
|
| 187 |
+
"the generated vertex voxels would be in the wrong coordinate space."
|
| 188 |
+
)
|
| 189 |
+
if voxel_res != min_res:
|
| 190 |
+
raise ValueError(
|
| 191 |
+
f"voxel encoder resolution={voxel_res} != V-VAE min_resolution={min_res}; "
|
| 192 |
+
"both stages must share the same active-voxel conditioning grid."
|
| 193 |
+
)
|
| 194 |
+
|
| 195 |
+
dino = (
|
| 196 |
+
DinoV2Encoder(
|
| 197 |
+
model_name=vflow_cfg["dino_version"],
|
| 198 |
+
hub_dir=args.dino_hub_dir,
|
| 199 |
+
img_res=vflow_cfg["image_resolution"],
|
| 200 |
+
)
|
| 201 |
+
.to(device)
|
| 202 |
+
.eval()
|
| 203 |
+
)
|
| 204 |
+
logging.info(f"loaded {vflow_cfg['dino_version']} from {args.dino_hub_dir}")
|
| 205 |
+
vertex_sampler = VertFlowEulerCfgSampler()
|
| 206 |
+
topo_sampler = TopoFlowEulerSampler()
|
| 207 |
+
|
| 208 |
+
dataset = VoxelVertexDataset(
|
| 209 |
+
root_dir=args.mesh_dir,
|
| 210 |
+
resolution=res,
|
| 211 |
+
min_resolution=min_res,
|
| 212 |
+
need_encoder_inputs=False,
|
| 213 |
+
num_samples=args.num_samples,
|
| 214 |
+
render=True,
|
| 215 |
+
img_res=args.img_res,
|
| 216 |
+
render_azimuth=args.render_azimuth,
|
| 217 |
+
render_elevation=args.render_elevation,
|
| 218 |
+
)
|
| 219 |
+
loader = DataLoader(
|
| 220 |
+
dataset,
|
| 221 |
+
batch_size=args.batch_size,
|
| 222 |
+
shuffle=False,
|
| 223 |
+
collate_fn=partial(collate_fn, resolution=res, min_resolution=min_res),
|
| 224 |
+
num_workers=args.num_workers,
|
| 225 |
+
pin_memory=True,
|
| 226 |
+
# EGL rendering hangs inside fork-ed children of a CUDA-initialized
|
| 227 |
+
# parent; spawn gives each worker a clean process for its EGL context.
|
| 228 |
+
multiprocessing_context="spawn" if args.num_workers > 0 else None,
|
| 229 |
+
worker_init_fn=worker_init if args.num_workers > 0 else None,
|
| 230 |
+
)
|
| 231 |
+
dupes = sorted(
|
| 232 |
+
s
|
| 233 |
+
for s, c in Counter(os.path.splitext(f)[0] for f in dataset.files).items()
|
| 234 |
+
if c > 1
|
| 235 |
+
)
|
| 236 |
+
if dupes:
|
| 237 |
+
logging.warning(
|
| 238 |
+
f"WARNING: {len(dupes)} duplicate mesh basename(s) — later samples will overwrite earlier outputs."
|
| 239 |
+
)
|
| 240 |
+
logging.info(
|
| 241 |
+
f"{len(dataset)} meshes from {args.mesh_dir} "
|
| 242 |
+
f"(vflow_steps={args.vflow_steps}, cfg={args.cfg_strength}, "
|
| 243 |
+
f"vert_num={args.vert_num}, use_gt_vert_count={args.use_gt_vert_count}, "
|
| 244 |
+
f"scaler={args.scaler}, density_max={density_max}, tflow_steps={args.tflow_steps}, "
|
| 245 |
+
f"view=az{args.render_azimuth}/el{args.render_elevation}, seed={args.seed}) "
|
| 246 |
+
f"-> {args.out_dir}"
|
| 247 |
+
)
|
| 248 |
+
|
| 249 |
+
n_ok = n_no_topo = n_fail = 0
|
| 250 |
+
t_start = time.time()
|
| 251 |
+
qbar = tqdm.tqdm(loader, desc="inference", unit="batch", dynamic_ncols=True)
|
| 252 |
+
for batch in qbar:
|
| 253 |
+
for err in batch["errors"]:
|
| 254 |
+
n_fail += 1
|
| 255 |
+
logging.error(
|
| 256 |
+
f"{err['name']}: FAILED during preprocessing: {err['error'].splitlines()[0]}"
|
| 257 |
+
)
|
| 258 |
+
if "name" not in batch:
|
| 259 |
+
continue
|
| 260 |
+
|
| 261 |
+
density = compute_density(batch, args, density_max, device)
|
| 262 |
+
with torch.no_grad():
|
| 263 |
+
cond = dino(np.stack(batch["image"])).float()
|
| 264 |
+
neg_cond = torch.zeros_like(cond)
|
| 265 |
+
|
| 266 |
+
# ---- stage 1: V-Flow on the 64^3 active voxels -> V-VAE vertex decode ----
|
| 267 |
+
min_active = batch[f"active_voxels_{min_res}"]
|
| 268 |
+
with torch.no_grad(), torch.autocast("cuda", dtype=torch.bfloat16):
|
| 269 |
+
min_active_coords = min_active.to(device)
|
| 270 |
+
noise = SparseTensor(
|
| 271 |
+
coords=min_active_coords.int(),
|
| 272 |
+
feats=torch.randn(
|
| 273 |
+
min_active_coords.shape[0], latent_dim, device=device
|
| 274 |
+
),
|
| 275 |
+
)
|
| 276 |
+
z_pred = vertex_sampler.sample(
|
| 277 |
+
model=vflow,
|
| 278 |
+
noise=noise,
|
| 279 |
+
cond=cond,
|
| 280 |
+
neg_cond=neg_cond,
|
| 281 |
+
steps=args.vflow_steps,
|
| 282 |
+
cfg_strength=args.cfg_strength,
|
| 283 |
+
rescale_t=args.rescale_t,
|
| 284 |
+
density=density,
|
| 285 |
+
)
|
| 286 |
+
pred_coords, pred_offsets = decode_vertices(
|
| 287 |
+
vvae, offset_head, z_pred, args.inference_threshold
|
| 288 |
+
)
|
| 289 |
+
|
| 290 |
+
keep_idx, verts_list, offsets_list = [], [], []
|
| 291 |
+
for b, name in enumerate(batch["name"]):
|
| 292 |
+
pred_sel = pred_coords[:, 0] == b
|
| 293 |
+
vert_int = pred_coords[pred_sel, 1:].long()
|
| 294 |
+
vert_off = pred_offsets[pred_sel]
|
| 295 |
+
export_vertex(
|
| 296 |
+
args.out_dir,
|
| 297 |
+
name,
|
| 298 |
+
type_name="pred",
|
| 299 |
+
vert_int=vert_int.numpy(),
|
| 300 |
+
vert_offsets=vert_off.numpy(),
|
| 301 |
+
resolution=res,
|
| 302 |
+
)
|
| 303 |
+
Image.fromarray(batch["image"][b]).save(
|
| 304 |
+
os.path.join(args.out_dir, f"{name}_render.png")
|
| 305 |
+
)
|
| 306 |
+
num_pred = int(vert_int.shape[0])
|
| 307 |
+
if num_pred < 3:
|
| 308 |
+
n_no_topo += 1
|
| 309 |
+
logging.warning(
|
| 310 |
+
f"{name}: only {num_pred} generated vertices; skipping topology."
|
| 311 |
+
)
|
| 312 |
+
elif num_pred > max_vertices:
|
| 313 |
+
n_no_topo += 1
|
| 314 |
+
logging.warning(
|
| 315 |
+
f"{name}: {num_pred} generated vertices exceed T-Flow "
|
| 316 |
+
f"max_vertices={max_vertices}; skipping topology."
|
| 317 |
+
)
|
| 318 |
+
else:
|
| 319 |
+
keep_idx.append(b)
|
| 320 |
+
verts_list.append(vert_int)
|
| 321 |
+
offsets_list.append(vert_off)
|
| 322 |
+
if not keep_idx:
|
| 323 |
+
continue
|
| 324 |
+
|
| 325 |
+
# ---- stage 2: T-Flow on the generated vertices -> T-VAE edge decode ----
|
| 326 |
+
with torch.no_grad():
|
| 327 |
+
verts, mask, lengths = pad_verts(verts_list, device)
|
| 328 |
+
voxel_list = [
|
| 329 |
+
min_active[min_active[:, 0] == b, 1:].long() for b in keep_idx
|
| 330 |
+
]
|
| 331 |
+
field = build_voxel_fields(voxel_list, voxel_res, device) # (B', R, R, R)
|
| 332 |
+
cond_vox = voxel_encoder(field) # (B', R'^3, cond_in_dim)
|
| 333 |
+
|
| 334 |
+
z0 = torch.randn(verts.shape[0], verts.shape[1], z_dim, device=device)
|
| 335 |
+
z_flow = topo_sampler.sample(
|
| 336 |
+
model=tflow,
|
| 337 |
+
noise=z0,
|
| 338 |
+
verts=verts,
|
| 339 |
+
mask=mask,
|
| 340 |
+
cond=cond_vox,
|
| 341 |
+
steps=args.tflow_steps,
|
| 342 |
+
)
|
| 343 |
+
z = z_flow.float() / latent_scale
|
| 344 |
+
|
| 345 |
+
with torch.autocast("cuda", dtype=torch.bfloat16):
|
| 346 |
+
edges_list = tvae.decode(
|
| 347 |
+
z,
|
| 348 |
+
verts=verts,
|
| 349 |
+
verts_mask=mask,
|
| 350 |
+
chunk_size=args.chunk_size,
|
| 351 |
+
threshold=args.edge_threshold,
|
| 352 |
+
)
|
| 353 |
+
|
| 354 |
+
for k, b in enumerate(keep_idx):
|
| 355 |
+
name = batch["name"][b]
|
| 356 |
+
faces = edges_to_faces(edges_list[k], lengths[k], args.fill_quad_rings)
|
| 357 |
+
if faces.shape[0] == 0:
|
| 358 |
+
n_no_topo += 1
|
| 359 |
+
logging.warning(
|
| 360 |
+
f"{name}: no faces decoded; the _pred PLYs are the only outputs."
|
| 361 |
+
)
|
| 362 |
+
continue
|
| 363 |
+
export_mesh(
|
| 364 |
+
args.out_dir,
|
| 365 |
+
name,
|
| 366 |
+
vert_int=verts_list[k].numpy(),
|
| 367 |
+
vert_offsets=offsets_list[k].numpy(),
|
| 368 |
+
faces=faces,
|
| 369 |
+
resolution=res,
|
| 370 |
+
)
|
| 371 |
+
n_ok += 1
|
| 372 |
+
|
| 373 |
+
logging.info(
|
| 374 |
+
f"done: {n_ok} ok, {n_no_topo} without topology, {n_fail} failed "
|
| 375 |
+
f"in {time.time() - t_start:.0f}s -> {args.out_dir}"
|
| 376 |
+
)
|
| 377 |
+
|
| 378 |
+
|
| 379 |
+
if __name__ == "__main__":
|
| 380 |
+
main()
|
scripts/tflow_inference.py
ADDED
|
@@ -0,0 +1,221 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""
|
| 2 |
+
Outputs into --out_dir:
|
| 3 |
+
<mesh_id>_pred.obj generated mesh (known verts, faces) in [-0.5, 0.5]
|
| 4 |
+
<mesh_id>_pred.ply fallback point cloud when no faces were generated
|
| 5 |
+
<mesh_id>_known.ply the known (dequantized) vertices fed to the flow
|
| 6 |
+
<mesh_id>_voxel_field.ply the active-voxel conditioning field (debug, --save_voxel_field)
|
| 7 |
+
|
| 8 |
+
Usage:
|
| 9 |
+
python scripts/tflow_inference.py --mesh_dir <dir> --out_dir outputs/tflow_run/<dir> \
|
| 10 |
+
[--steps 50] [--no-use_cond] [--no-fill_quad_rings]
|
| 11 |
+
"""
|
| 12 |
+
|
| 13 |
+
import argparse
|
| 14 |
+
import os
|
| 15 |
+
import sys
|
| 16 |
+
import time
|
| 17 |
+
|
| 18 |
+
import numpy as np
|
| 19 |
+
import torch
|
| 20 |
+
import trimesh
|
| 21 |
+
import tqdm
|
| 22 |
+
|
| 23 |
+
ROOT = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
|
| 24 |
+
sys.path.insert(0, ROOT)
|
| 25 |
+
os.environ.setdefault("PYTORCH_CUDA_ALLOC_CONF", "expandable_segments:True")
|
| 26 |
+
|
| 27 |
+
from torch.utils.data import DataLoader
|
| 28 |
+
|
| 29 |
+
from dataset.topo_dataset import TopoVoxelDataset, collate_fn
|
| 30 |
+
from models import (
|
| 31 |
+
TopologyVAE,
|
| 32 |
+
TopologySiTFlow,
|
| 33 |
+
TopoFlowEulerSampler,
|
| 34 |
+
VoxelFieldConditioner,
|
| 35 |
+
)
|
| 36 |
+
import utils.logging as logging
|
| 37 |
+
from utils.inference import build_voxel_fields, edges_to_faces, pad_verts
|
| 38 |
+
from utils.load import load_latov2_model
|
| 39 |
+
|
| 40 |
+
|
| 41 |
+
def parse_args():
|
| 42 |
+
p = argparse.ArgumentParser(description="T-Flow topology generation inference")
|
| 43 |
+
p.add_argument("--mesh_dir", required=True, help="directory of input meshes")
|
| 44 |
+
p.add_argument("--out_dir", required=True, help="output directory for the meshes")
|
| 45 |
+
p.add_argument("--tflow_ckpt", default=os.path.join(ROOT, "ckpt", "tflow.pt"))
|
| 46 |
+
p.add_argument("--tvae_ckpt", default=os.path.join(ROOT, "ckpt", "tvae.pt"))
|
| 47 |
+
p.add_argument(
|
| 48 |
+
"--voxel_encoder_ckpt",
|
| 49 |
+
default=os.path.join(ROOT, "ckpt", "voxel_encoder.pt"),
|
| 50 |
+
)
|
| 51 |
+
p.add_argument("--batch_size", type=int, default=1)
|
| 52 |
+
p.add_argument(
|
| 53 |
+
"--num_samples", type=int, default=None, help="only run the first N meshes"
|
| 54 |
+
)
|
| 55 |
+
p.add_argument("--num_workers", type=int, default=4)
|
| 56 |
+
p.add_argument("--seed", type=int, default=42)
|
| 57 |
+
# flow sampling
|
| 58 |
+
p.add_argument("--steps", type=int, default=50, help="Euler steps")
|
| 59 |
+
p.add_argument(
|
| 60 |
+
"--use_cond",
|
| 61 |
+
action=argparse.BooleanOptionalAction,
|
| 62 |
+
default=True,
|
| 63 |
+
help=(
|
| 64 |
+
"condition on the active-voxel field. --no-use_cond runs the flow "
|
| 65 |
+
"unconditionally (the model's learned null token)."
|
| 66 |
+
),
|
| 67 |
+
)
|
| 68 |
+
# topology decoding
|
| 69 |
+
p.add_argument("--edge_threshold", type=float, default=0.0)
|
| 70 |
+
p.add_argument("--chunk_size", type=int, default=20000)
|
| 71 |
+
p.add_argument(
|
| 72 |
+
"--fill_quad_rings",
|
| 73 |
+
action=argparse.BooleanOptionalAction,
|
| 74 |
+
default=True,
|
| 75 |
+
help=(
|
| 76 |
+
"post-process: split chordless 4-vertex rings into two triangles "
|
| 77 |
+
"(pure topology, not the voxel support filter)"
|
| 78 |
+
),
|
| 79 |
+
)
|
| 80 |
+
p.add_argument(
|
| 81 |
+
"--save_voxel_field",
|
| 82 |
+
action=argparse.BooleanOptionalAction,
|
| 83 |
+
default=True,
|
| 84 |
+
help="also dump the active-voxel conditioning field as a point cloud",
|
| 85 |
+
)
|
| 86 |
+
args = p.parse_args()
|
| 87 |
+
if args.num_samples is not None and args.num_samples <= 0:
|
| 88 |
+
args.num_samples = None # <= 0 means "all", not python slice semantics
|
| 89 |
+
return args
|
| 90 |
+
|
| 91 |
+
|
| 92 |
+
def main():
|
| 93 |
+
logging.info("T-Flow inference starting...")
|
| 94 |
+
|
| 95 |
+
args = parse_args()
|
| 96 |
+
device = torch.device("cuda")
|
| 97 |
+
torch.manual_seed(args.seed)
|
| 98 |
+
np.random.seed(args.seed)
|
| 99 |
+
os.makedirs(args.out_dir, exist_ok=True)
|
| 100 |
+
|
| 101 |
+
tflow, tflow_cfg = load_latov2_model(TopologySiTFlow, args.tflow_ckpt, device)
|
| 102 |
+
tvae, _ = load_latov2_model(TopologyVAE, args.tvae_ckpt, device)
|
| 103 |
+
voxel_encoder, venc_cfg = load_latov2_model(
|
| 104 |
+
VoxelFieldConditioner, args.voxel_encoder_ckpt, device
|
| 105 |
+
)
|
| 106 |
+
z_dim = int(tflow_cfg["args"]["z_dim"])
|
| 107 |
+
num_discrete = int(tflow_cfg["args"]["num_discrete"])
|
| 108 |
+
max_vertices = int(tflow_cfg["args"]["max_vertices"])
|
| 109 |
+
latent_scale = float(tflow_cfg["latent_scale"])
|
| 110 |
+
voxel_res = int(venc_cfg["resolution"])
|
| 111 |
+
sampler = TopoFlowEulerSampler()
|
| 112 |
+
|
| 113 |
+
dataset = TopoVoxelDataset(
|
| 114 |
+
root_dir=args.mesh_dir,
|
| 115 |
+
num_discrete=num_discrete,
|
| 116 |
+
voxel_res=voxel_res,
|
| 117 |
+
max_vertices=max_vertices,
|
| 118 |
+
num_samples=args.num_samples,
|
| 119 |
+
)
|
| 120 |
+
loader = DataLoader(
|
| 121 |
+
dataset,
|
| 122 |
+
batch_size=args.batch_size,
|
| 123 |
+
shuffle=False,
|
| 124 |
+
collate_fn=collate_fn,
|
| 125 |
+
num_workers=args.num_workers,
|
| 126 |
+
pin_memory=True,
|
| 127 |
+
)
|
| 128 |
+
logging.info(
|
| 129 |
+
f"{len(dataset)} meshes from {args.mesh_dir} "
|
| 130 |
+
f"(steps={args.steps}, use_cond={args.use_cond}, num_discrete={num_discrete}, "
|
| 131 |
+
f"voxel_res={voxel_res}, latent_scale={latent_scale}, seed={args.seed}) "
|
| 132 |
+
f"-> {args.out_dir}"
|
| 133 |
+
)
|
| 134 |
+
|
| 135 |
+
n_ok = n_fail = 0
|
| 136 |
+
t_start = time.time()
|
| 137 |
+
qbar = tqdm.tqdm(loader, desc="inference", unit="batch", dynamic_ncols=True)
|
| 138 |
+
for batch in qbar:
|
| 139 |
+
for err in batch["errors"]:
|
| 140 |
+
n_fail += 1
|
| 141 |
+
logging.error(
|
| 142 |
+
f"{err['name']}: FAILED during preprocessing: {err['error'].splitlines()[0]}"
|
| 143 |
+
)
|
| 144 |
+
if "name" not in batch:
|
| 145 |
+
continue
|
| 146 |
+
|
| 147 |
+
names = batch["name"]
|
| 148 |
+
verts_list = batch["vertices"] # list of (N_i, 3) long in [0, num_discrete)
|
| 149 |
+
voxel_list = batch["voxel_coords"] # list of (M_i, 3) long in [0, voxel_res)
|
| 150 |
+
|
| 151 |
+
with torch.no_grad():
|
| 152 |
+
verts, mask, lengths = pad_verts(
|
| 153 |
+
verts_list, device
|
| 154 |
+
) # (B, N_max, 3), (B, N_max)
|
| 155 |
+
if args.use_cond:
|
| 156 |
+
field = build_voxel_fields(
|
| 157 |
+
voxel_list, voxel_res, device
|
| 158 |
+
) # (B, R, R, R)
|
| 159 |
+
cond = voxel_encoder(field) # (B, R'^3, cond_in_dim)
|
| 160 |
+
else:
|
| 161 |
+
cond = None # unconditional
|
| 162 |
+
|
| 163 |
+
z0 = torch.randn(verts.shape[0], verts.shape[1], z_dim, device=device)
|
| 164 |
+
z_flow = sampler.sample(
|
| 165 |
+
model=tflow,
|
| 166 |
+
noise=z0,
|
| 167 |
+
verts=verts,
|
| 168 |
+
mask=mask,
|
| 169 |
+
cond=cond,
|
| 170 |
+
steps=args.steps,
|
| 171 |
+
)
|
| 172 |
+
z = z_flow.float() / latent_scale
|
| 173 |
+
|
| 174 |
+
with torch.autocast("cuda", dtype=torch.bfloat16):
|
| 175 |
+
edges_list = tvae.decode(
|
| 176 |
+
z,
|
| 177 |
+
verts=verts,
|
| 178 |
+
verts_mask=mask,
|
| 179 |
+
chunk_size=args.chunk_size,
|
| 180 |
+
threshold=args.edge_threshold,
|
| 181 |
+
)
|
| 182 |
+
|
| 183 |
+
for b, name in enumerate(names):
|
| 184 |
+
num_vertices = lengths[b]
|
| 185 |
+
if num_vertices == 0:
|
| 186 |
+
logging.warning(f"{name}: no known vertices; skipping.")
|
| 187 |
+
continue
|
| 188 |
+
edges = edges_list[b]
|
| 189 |
+
faces = edges_to_faces(edges, num_vertices, args.fill_quad_rings)
|
| 190 |
+
|
| 191 |
+
verts_int = verts_list[b]
|
| 192 |
+
disp = (verts_int.numpy().astype(np.float64) + 0.5) / num_discrete - 0.5
|
| 193 |
+
|
| 194 |
+
if faces.shape[0] > 0:
|
| 195 |
+
trimesh.Trimesh(vertices=disp, faces=faces).export(
|
| 196 |
+
os.path.join(args.out_dir, f"{name}_pred.obj")
|
| 197 |
+
)
|
| 198 |
+
else:
|
| 199 |
+
trimesh.PointCloud(disp).export(
|
| 200 |
+
os.path.join(args.out_dir, f"{name}_pred.ply")
|
| 201 |
+
)
|
| 202 |
+
trimesh.PointCloud(disp).export(
|
| 203 |
+
os.path.join(args.out_dir, f"{name}_known.ply")
|
| 204 |
+
)
|
| 205 |
+
if args.use_cond and args.save_voxel_field and voxel_list[b].shape[0] > 0:
|
| 206 |
+
vox_pts = (
|
| 207 |
+
voxel_list[b].numpy().astype(np.float64) + 0.5
|
| 208 |
+
) / voxel_res - 0.5
|
| 209 |
+
trimesh.PointCloud(vox_pts).export(
|
| 210 |
+
os.path.join(args.out_dir, f"{name}_voxel_field.ply")
|
| 211 |
+
)
|
| 212 |
+
|
| 213 |
+
n_ok += 1
|
| 214 |
+
|
| 215 |
+
logging.info(
|
| 216 |
+
f"done: {n_ok} ok, {n_fail} failed in {time.time() - t_start:.0f}s -> {args.out_dir}"
|
| 217 |
+
)
|
| 218 |
+
|
| 219 |
+
|
| 220 |
+
if __name__ == "__main__":
|
| 221 |
+
main()
|