prithivMLmods commited on
Commit
1cc6214
·
verified ·
1 Parent(s): e262431

Upload 56 files

Browse files
This view is limited to 50 files because it contains too many changes.   See raw diff
Files changed (50) hide show
  1. .gitattributes +4 -0
  2. app.py +706 -0
  3. assets/example_mesh/crocodile.glb +3 -0
  4. assets/example_mesh/dragon.glb +3 -0
  5. assets/example_mesh/spaceman.glb +3 -0
  6. assets/teaser.png +3 -0
  7. dataset/mesh_render.py +265 -0
  8. dataset/topo_dataset.py +88 -0
  9. dataset/utils.py +283 -0
  10. dataset/voxel_dataset.py +211 -0
  11. models/__init__.py +9 -0
  12. models/dino_encoder.py +73 -0
  13. models/flow_sampler.py +58 -0
  14. models/offset_head.py +22 -0
  15. models/topo_autoencoder.py +297 -0
  16. models/topo_flow.py +400 -0
  17. models/vdf_encoder.py +55 -0
  18. models/vertex_autoencoder.py +595 -0
  19. models/vertex_structured_flow.py +147 -0
  20. models/voxel_encoder.py +183 -0
  21. modules/attention.py +160 -0
  22. modules/norm.py +41 -0
  23. modules/pointnet.py +330 -0
  24. modules/sparse/__init__.py +130 -0
  25. modules/sparse/attention/__init__.py +27 -0
  26. modules/sparse/attention/full_attn.py +238 -0
  27. modules/sparse/attention/modules.py +214 -0
  28. modules/sparse/attention/serialized_attn.py +217 -0
  29. modules/sparse/attention/windowed_attn.py +158 -0
  30. modules/sparse/basic.py +482 -0
  31. modules/sparse/blocks.py +71 -0
  32. modules/sparse/conv/__init__.py +44 -0
  33. modules/sparse/conv/conv_spconv.py +107 -0
  34. modules/sparse/conv/conv_torchsparse.py +60 -0
  35. modules/sparse/linear.py +38 -0
  36. modules/sparse/nonlinearity.py +58 -0
  37. modules/sparse/norm.py +81 -0
  38. modules/sparse/spatial.py +158 -0
  39. modules/sparse/transformer/__init__.py +26 -0
  40. modules/sparse/transformer/bases.py +234 -0
  41. modules/sparse/transformer/blocks.py +165 -0
  42. modules/sparse/transformer/modulated.py +119 -0
  43. modules/transformer/__init__.py +24 -0
  44. modules/transformer/blocks.py +276 -0
  45. modules/transformer/hybrid.py +236 -0
  46. modules/utils.py +145 -0
  47. requirements.txt +23 -0
  48. scripts/ckpt_download.py +90 -0
  49. scripts/e2e_inference.py +380 -0
  50. 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

  • SHA256: 3a8551b9f4b9b0c5becfa6225ca554237d10f3eea283e1b51790288e27776bc4
  • Pointer size: 133 Bytes
  • Size of remote file: 12.9 MB
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()