diff --git a/.gitattributes b/.gitattributes
index d08e4ab61215c868fc40d52ff43066554c99803c..100d0bb3ebc6786ce06bb5ee510dfea89d1a780f 100644
--- a/.gitattributes
+++ b/.gitattributes
@@ -41,3 +41,7 @@ examples/example-04.jpg filter=lfs diff=lfs merge=lfs -text
examples/example-05.pdf filter=lfs diff=lfs merge=lfs -text
examples/2.jpg filter=lfs diff=lfs merge=lfs -text
examples/4.jpg filter=lfs diff=lfs merge=lfs -text
+assets/example_mesh/crocodile.glb filter=lfs diff=lfs merge=lfs -text
+assets/example_mesh/dragon.glb filter=lfs diff=lfs merge=lfs -text
+assets/example_mesh/spaceman.glb filter=lfs diff=lfs merge=lfs -text
+assets/teaser.png filter=lfs diff=lfs merge=lfs -text
diff --git a/app.py b/app.py
new file mode 100644
index 0000000000000000000000000000000000000000..843ccf1dbae8c9eb8a359144e24a10b90d400780
--- /dev/null
+++ b/app.py
@@ -0,0 +1,706 @@
+"""
+LATO.2 Gradio App — Image-to-3D Mesh Generation
+=================================================
+Factorized 3D Mesh Generation with Vertex and Topology Flow.
+
+Launches a Gradio interface that accepts an input image (or mesh),
+runs the full V-Flow → T-Flow pipeline, and displays the result
+with Rerun 3D viewer + GLB download.
+
+Usage:
+ python app.py [--share] [--port 7860]
+"""
+
+import argparse
+import os
+import sys
+import tempfile
+import time
+import uuid
+from pathlib import Path
+
+# ── project root on sys.path ──────────────────────────────────────────────────
+ROOT = os.path.dirname(os.path.abspath(__file__))
+sys.path.insert(0, ROOT)
+os.environ.setdefault("PYTORCH_CUDA_ALLOC_CONF", "expandable_segments:True")
+os.environ.setdefault("XDG_RUNTIME_DIR", os.path.join(tempfile.gettempdir(), "runtime-root"))
+os.makedirs(os.environ["XDG_RUNTIME_DIR"], exist_ok=True)
+os.environ.setdefault("EGL_PLATFORM", "surfaceless")
+
+import gradio as gr
+import numpy as np
+import rerun as rr
+import torch
+import trimesh
+from gradio_rerun import Rerun
+from PIL import Image
+
+from dataset.utils import (
+ MESH_EXTENSIONS,
+ dedup_quantized_mesh,
+ extract_active_voxels,
+ quantize_mesh_clustering,
+)
+from models import (
+ DinoV2Encoder,
+ OffsetHead,
+ TopoFlowEulerSampler,
+ TopologySiTFlow,
+ TopologyVAE,
+ VertexSLatFlowModel,
+ VertFlowEulerCfgSampler,
+ VertexVAE,
+ VoxelFieldConditioner,
+)
+from modules.sparse import SparseTensor
+import utils.logging as logging
+from utils.inference import (
+ build_voxel_fields,
+ decode_vertices,
+ edges_to_faces,
+ pad_verts,
+)
+from utils.load import load_latov2_model
+
+
+# ═══════════════════════════════════════════════════════════════════════════════
+# Global model state (lazy-loaded once on first inference)
+# ═══════════════════════════════════════════════════════════════════════════════
+_models = {}
+_configs = {}
+_device = None
+
+OUTPUT_DIR = os.path.join(ROOT, "gradio_outputs")
+os.makedirs(OUTPUT_DIR, exist_ok=True)
+
+
+def _load_models():
+ """Load all LATO.2 sub-models once (idempotent)."""
+ global _models, _configs, _device
+
+ if _models:
+ return # already loaded
+
+ _device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
+ device = _device
+ ckpt = os.path.join(ROOT, "ckpt")
+
+ logging.info("Loading LATO.2 models …")
+
+ # Stage 1 — vertex generation
+ vflow, vflow_cfg = load_latov2_model(VertexSLatFlowModel, os.path.join(ckpt, "vflow.pt"), device)
+ vvae, vvae_cfg = load_latov2_model(VertexVAE, os.path.join(ckpt, "vvae.pt"), device)
+ offset_head, _ = load_latov2_model(OffsetHead, os.path.join(ckpt, "offset_head.pt"), device)
+
+ # Stage 2 — topology generation
+ tflow, tflow_cfg = load_latov2_model(TopologySiTFlow, os.path.join(ckpt, "tflow.pt"), device)
+ tvae, _ = load_latov2_model(TopologyVAE, os.path.join(ckpt, "tvae.pt"), device)
+ voxel_encoder, venc_cfg = load_latov2_model(VoxelFieldConditioner, os.path.join(ckpt, "voxel_encoder.pt"), device)
+
+ # DINO-v2 image encoder
+ dino = (
+ DinoV2Encoder(
+ model_name=vflow_cfg["dino_version"],
+ hub_dir=os.path.join(ckpt, "dinov2"),
+ img_res=vflow_cfg["image_resolution"],
+ )
+ .to(device)
+ .eval()
+ )
+
+ _models = dict(
+ vflow=vflow, vvae=vvae, offset_head=offset_head,
+ tflow=tflow, tvae=tvae, voxel_encoder=voxel_encoder,
+ dino=dino,
+ vertex_sampler=VertFlowEulerCfgSampler(),
+ topo_sampler=TopoFlowEulerSampler(),
+ )
+ _configs = dict(vflow=vflow_cfg, vvae=vvae_cfg, tflow=tflow_cfg, venc=venc_cfg)
+ logging.info("All models loaded ✓")
+
+
+# ═══════════════════════════════════════════════════════════════════════════════
+# Helper: prepare input mesh → active voxels + quantized vertices
+# ═══════════════════════════════════════════════════════════════════════════════
+
+def _prepare_mesh_data(mesh_path: str, resolution: int, min_resolution: int):
+ """Quantize and extract voxels from a mesh file for the pipeline."""
+ quantized = quantize_mesh_clustering(mesh_path, resolution=resolution)
+ if quantized is None:
+ raise ValueError("Input mesh is empty or could not be loaded.")
+ v_int, offsets, faces = quantized
+ if len(faces) < 1 or len(v_int) < 3:
+ raise ValueError("Mesh is degenerate after quantization.")
+
+ gt_int, gt_offsets, gt_faces = dedup_quantized_mesh(v_int, offsets, faces, resolution)
+ if len(gt_int) < 3 or len(gt_faces) < 1:
+ raise ValueError("Too few vertices/faces after deduplication.")
+
+ quant_v = gt_int.astype(np.float64) / (resolution - 1.0) - 0.5
+ quant_v = np.clip(quant_v, -0.5 + 1e-6, 0.5 - 1e-6).astype(np.float32)
+
+ min_active = extract_active_voxels(quant_v, gt_faces, min_resolution)
+ return gt_int, gt_offsets, gt_faces, quant_v, min_active
+
+
+def _render_mesh_to_image(mesh_path: str, resolution: int, img_res: int = 518,
+ azimuth: float = 45.0, elevation: float = 30.0):
+ """Render a conditioning view from a mesh (used when no user image is provided)."""
+ quantized = quantize_mesh_clustering(mesh_path, resolution=resolution)
+ if quantized is None:
+ return None
+ v_int, _, faces = quantized
+ render_v = v_int.astype(np.float64) / resolution - 0.5
+
+ from dataset.mesh_render import WhiteModelRenderer
+ renderer = WhiteModelRenderer(
+ img_res=img_res,
+ mesh_color=(0.78, 0.78, 0.82),
+ bg_color=(0.0, 0.0, 0.0),
+ up_axis="y",
+ add_ground=False,
+ shadow=True,
+ crop_to_object=True,
+ crop_padding=1.2,
+ )
+ imgs, _ = renderer.render(
+ np.asarray(render_v, dtype=np.float64),
+ np.asarray(faces, dtype=np.int64),
+ num_views=1,
+ azimuths=[azimuth],
+ elevations=[elevation],
+ )
+ return imgs[0] # (H, W, 3) uint8
+
+
+# ═══════════════════════════════════════════════════════════════════════════════
+# Core generation pipeline
+# ═══════════════════════════════════════════════════════════════════════════════
+
+def generate_mesh(
+ input_image: np.ndarray | None,
+ input_mesh_path: str | None,
+ vert_num: int,
+ cfg_strength: float,
+ vflow_steps: int,
+ tflow_steps: int,
+ seed: int,
+ progress=gr.Progress(track_tqdm=True),
+):
+ """
+ Main generation function.
+ - input_image: user-uploaded image (H, W, 3) uint8 — used as DINOv2 conditioning
+ - input_mesh_path: reference mesh file — provides the voxel scaffold
+ If only an image is supplied, the user must also supply a reference mesh for
+ the voxel scaffold (or we use one of the bundled examples).
+ """
+ _load_models() # ensure models are loaded
+
+ device = _device
+ m = _models
+ c = _configs
+
+ torch.manual_seed(seed)
+ np.random.seed(seed)
+
+ res = c["vvae"]["resolution"]
+ min_res = c["vvae"]["min_resolution"]
+ latent_dim = c["vflow"]["latent_dim"]
+ density_max = c["vflow"]["max_vertex_num"]
+ z_dim = int(c["tflow"]["args"]["z_dim"])
+ max_vertices = int(c["tflow"]["args"]["max_vertices"])
+ latent_scale = float(c["tflow"]["latent_scale"])
+ voxel_res = int(c["venc"]["resolution"])
+ inference_threshold = 0.5
+
+ run_id = str(uuid.uuid4())[:8]
+
+ # ── Resolve mesh scaffold ─────────────────────────────────────────────
+ if input_mesh_path is None or not os.path.isfile(input_mesh_path):
+ raise gr.Error(
+ "A reference mesh file is required to provide the voxel scaffold. "
+ "Please upload a .glb / .obj / .ply / .stl mesh."
+ )
+
+ progress(0.05, desc="Quantizing mesh & extracting voxels …")
+ gt_int, gt_offsets, gt_faces, quant_v, min_active = _prepare_mesh_data(
+ input_mesh_path, res, min_res
+ )
+
+ # ── Resolve conditioning image ────────────────────────────────────────
+ if input_image is not None:
+ cond_img = np.asarray(input_image, dtype=np.uint8)
+ if cond_img.ndim == 2:
+ cond_img = np.stack([cond_img] * 3, axis=-1)
+ elif cond_img.shape[-1] == 4:
+ cond_img = cond_img[:, :, :3]
+ else:
+ progress(0.08, desc="Rendering conditioning view from mesh …")
+ cond_img = _render_mesh_to_image(input_mesh_path, res)
+ if cond_img is None:
+ raise gr.Error("Could not render a conditioning view from the mesh.")
+
+ # ── Compute density conditioning ──────────────────────────────────────
+ clamped = float(min(max(vert_num, 200), 5000))
+ density = torch.tensor([clamped], dtype=torch.float32, device=device)
+ density = density / density_max * 1000.0
+
+ # ── Stage 1: V-Flow → V-VAE ──────────────────────────────────────────
+ progress(0.12, desc="Running V-Flow (vertex generation) …")
+ min_active_batched = torch.cat(
+ [torch.zeros(min_active.shape[0], 1, dtype=torch.int32), min_active], dim=1
+ )
+
+ with torch.no_grad():
+ cond = m["dino"](cond_img).float()
+ neg_cond = torch.zeros_like(cond)
+
+ with torch.no_grad(), torch.autocast("cuda", dtype=torch.bfloat16):
+ min_active_coords = min_active_batched.to(device)
+ noise = SparseTensor(
+ coords=min_active_coords.int(),
+ feats=torch.randn(
+ min_active_coords.shape[0], latent_dim, device=device
+ ),
+ )
+ z_pred = m["vertex_sampler"].sample(
+ model=m["vflow"],
+ noise=noise,
+ cond=cond,
+ neg_cond=neg_cond,
+ steps=vflow_steps,
+ cfg_strength=cfg_strength,
+ rescale_t=1.0,
+ density=density,
+ )
+ pred_coords, pred_offsets = decode_vertices(
+ m["vvae"], m["offset_head"], z_pred, inference_threshold
+ )
+
+ progress(0.55, desc="Decoding vertices …")
+
+ # Filter to batch 0
+ pred_sel = pred_coords[:, 0] == 0
+ vert_int = pred_coords[pred_sel, 1:].long()
+ vert_off = pred_offsets[pred_sel]
+ num_pred = int(vert_int.shape[0])
+
+ if num_pred < 3:
+ raise gr.Error(f"Only {num_pred} vertices generated. Try different parameters.")
+ if num_pred > max_vertices:
+ raise gr.Error(
+ f"Generated {num_pred} vertices exceeds T-Flow max ({max_vertices}). "
+ "Try reducing the vertex count."
+ )
+
+ # ── Stage 2: T-Flow → T-VAE ──────────────────────────────────────────
+ progress(0.60, desc="Running T-Flow (topology generation) …")
+
+ with torch.no_grad():
+ verts, mask, lengths = pad_verts([vert_int], device)
+ voxel_list = [min_active.long()]
+ field = build_voxel_fields(voxel_list, voxel_res, device)
+ cond_vox = m["voxel_encoder"](field)
+
+ z0 = torch.randn(verts.shape[0], verts.shape[1], z_dim, device=device)
+ z_flow = m["topo_sampler"].sample(
+ model=m["tflow"],
+ noise=z0,
+ verts=verts,
+ mask=mask,
+ cond=cond_vox,
+ steps=tflow_steps,
+ )
+ z = z_flow.float() / latent_scale
+
+ with torch.autocast("cuda", dtype=torch.bfloat16):
+ edges_list = m["tvae"].decode(
+ z,
+ verts=verts,
+ verts_mask=mask,
+ chunk_size=20000,
+ threshold=0.0,
+ )
+
+ progress(0.85, desc="Assembling faces & exporting …")
+ faces = edges_to_faces(edges_list[0], lengths[0], fill_quad_rings=True)
+
+ if faces.shape[0] == 0:
+ raise gr.Error("No faces were generated. Try different parameters or a different input.")
+
+ # ── Build final mesh ──────────────────────────────────────────────────
+ vert_np = vert_int.numpy()
+ off_np = vert_off.numpy()
+ vert_with_offset = (
+ vert_np.astype(np.float64) / res
+ - 0.5
+ + off_np.astype(np.float64) / (res * 2.0)
+ )
+ mesh = trimesh.Trimesh(vertices=vert_with_offset, faces=faces, process=False)
+
+ # ── Export to GLB ─────────────────────────────────────────────────────
+ glb_filename = f"lato2_{run_id}.glb"
+ glb_path = os.path.join(OUTPUT_DIR, glb_filename)
+ mesh.export(glb_path, file_type="glb")
+
+ # ── Build Rerun visualization ─────────────────────────────────────────
+ progress(0.92, desc="Building 3D visualization …")
+ rr_data = _build_rerun_stream(mesh, cond_img, run_id)
+
+ progress(1.0, desc="Done ✓")
+ return rr_data, glb_path
+
+
+# ═══════════════════════════════════════════════════════════════════════════════
+# Rerun 3D Visualization
+# ═══════════════════════════════════════════════════════════════════════════════
+
+def _build_rerun_stream(mesh: trimesh.Trimesh, cond_img: np.ndarray, run_id: str):
+ """Create an .rrd byte stream with the generated mesh + conditioning image."""
+ rrd_path = os.path.join(OUTPUT_DIR, f"lato2_{run_id}.rrd")
+
+ rr.init("LATO.2 — 3D Mesh Generation", spawn=False)
+ rec = rr.new_recording(application_id="LATO.2", recording_id=run_id)
+
+ vertices = np.asarray(mesh.vertices, dtype=np.float32)
+ faces = np.asarray(mesh.faces, dtype=np.uint32)
+
+ # Compute vertex normals for nicer shading
+ if mesh.vertex_normals is not None and len(mesh.vertex_normals) > 0:
+ normals = np.asarray(mesh.vertex_normals, dtype=np.float32)
+ else:
+ normals = None
+
+ # Log the generated mesh
+ rec.log(
+ "world/generated_mesh",
+ rr.Mesh3D(
+ vertex_positions=vertices,
+ triangle_indices=faces,
+ vertex_normals=normals,
+ ),
+ )
+
+ # Log the conditioning image
+ if cond_img is not None:
+ rec.log("conditioning_image", rr.Image(cond_img))
+
+ # Log mesh stats as text
+ rec.log(
+ "world/stats",
+ rr.TextDocument(
+ f"Vertices: {len(vertices)}\n"
+ f"Faces: {len(faces)}\n"
+ f"Bounding box: {vertices.min(axis=0).tolist()} → {vertices.max(axis=0).tolist()}"
+ ),
+ )
+
+ rrd_bytes = rec.memory_recording()
+ return rrd_bytes
+
+
+# ═══════════════════════════════════════════════════════════════════════════════
+# Gradio UI
+# ═══════════════════════════════════════════════════════════════════════════════
+
+TITLE = "LATO.2: Factorized 3D Mesh Generation"
+DESCRIPTION = """
+**LATO.2** factorizes mesh generation into a **Vertex Flow (V-Flow)** for vertex positions
+and a **Topology Flow (T-Flow)** for connectivity prediction.
+
+### How to use
+1. **Upload a conditioning image** — this drives the DINOv2 shape conditioning.
+2. **Upload a reference mesh** (.glb / .obj / .ply / .stl) — this provides the coarse voxel scaffold.
+ *If you skip the image, a rendered view of the mesh is used as conditioning instead.*
+3. Adjust parameters and click **🚀 Generate 3D Mesh**.
+4. Explore the result in the **Rerun 3D Viewer** and download the **.glb** file.
+"""
+
+EXAMPLES_DIR = os.path.join(ROOT, "assets", "example_mesh")
+
+CSS = """
+/* ── Dark premium theme overrides ──────────────────────────────────── */
+.gradio-container {
+ max-width: 1400px !important;
+ margin: auto;
+}
+
+#app-title {
+ text-align: center;
+ background: linear-gradient(135deg, #667eea 0%, #764ba2 100%);
+ -webkit-background-clip: text;
+ -webkit-text-fill-color: transparent;
+ background-clip: text;
+ font-size: 2.4rem;
+ font-weight: 800;
+ letter-spacing: -0.02em;
+ margin-bottom: 0.2em;
+ font-family: 'Inter', 'Segoe UI', sans-serif;
+}
+
+#app-subtitle {
+ text-align: center;
+ color: #9ca3af;
+ font-size: 1.05rem;
+ margin-top: -0.5em;
+ margin-bottom: 1.5em;
+}
+
+.generate-btn {
+ background: linear-gradient(135deg, #667eea 0%, #764ba2 100%) !important;
+ border: none !important;
+ color: white !important;
+ font-weight: 700 !important;
+ font-size: 1.15rem !important;
+ padding: 14px 32px !important;
+ border-radius: 12px !important;
+ box-shadow: 0 4px 15px rgba(102, 126, 234, 0.4) !important;
+ transition: all 0.3s ease !important;
+ letter-spacing: 0.02em !important;
+}
+
+.generate-btn:hover {
+ transform: translateY(-2px) !important;
+ box-shadow: 0 8px 25px rgba(102, 126, 234, 0.55) !important;
+}
+
+.param-accordion {
+ border: 1px solid rgba(102, 126, 234, 0.2) !important;
+ border-radius: 12px !important;
+ margin-top: 8px !important;
+}
+
+.download-btn {
+ background: linear-gradient(135deg, #11998e 0%, #38ef7d 100%) !important;
+ border: none !important;
+ color: white !important;
+ font-weight: 700 !important;
+ font-size: 1.05rem !important;
+ padding: 12px 28px !important;
+ border-radius: 12px !important;
+ box-shadow: 0 4px 15px rgba(17, 153, 142, 0.35) !important;
+ transition: all 0.3s ease !important;
+}
+
+.download-btn:hover {
+ transform: translateY(-2px) !important;
+ box-shadow: 0 8px 25px rgba(17, 153, 142, 0.5) !important;
+}
+
+.info-badge {
+ display: inline-block;
+ background: rgba(102, 126, 234, 0.12);
+ color: #667eea;
+ padding: 4px 12px;
+ border-radius: 20px;
+ font-size: 0.85rem;
+ font-weight: 600;
+ margin: 4px 0;
+}
+
+footer { display: none !important; }
+"""
+
+def build_app():
+ """Construct the Gradio Blocks app."""
+ theme = gr.themes.Soft(
+ primary_hue=gr.themes.colors.indigo,
+ secondary_hue=gr.themes.colors.purple,
+ neutral_hue=gr.themes.colors.gray,
+ font=gr.themes.GoogleFont("Inter"),
+ ).set(
+ body_background_fill="*neutral_950",
+ body_background_fill_dark="*neutral_950",
+ block_background_fill="*neutral_900",
+ block_background_fill_dark="*neutral_900",
+ block_border_width="0px",
+ block_shadow="0 2px 12px rgba(0,0,0,0.3)",
+ input_background_fill="*neutral_800",
+ input_background_fill_dark="*neutral_800",
+ )
+
+ with gr.Blocks(
+ theme=theme,
+ css=CSS,
+ title="LATO.2 — Image to 3D Mesh",
+ ) as app:
+ # ── Header ────────────────────────────────────────────────────────
+ gr.HTML(
+ '
LATO.2
'
+ 'Factorized 3D Mesh Generation with Vertex & Topology Flow
'
+ )
+ gr.Markdown(DESCRIPTION)
+
+ with gr.Row(equal_height=False):
+ # ── LEFT: Inputs ──────────────────────────────────────────────
+ with gr.Column(scale=1, min_width=380):
+ gr.Markdown("### 📥 Inputs")
+
+ input_image = gr.Image(
+ label="Conditioning Image (optional)",
+ type="numpy",
+ height=280,
+ sources=["upload", "clipboard"],
+ elem_id="input-image",
+ )
+ input_mesh = gr.File(
+ label="Reference Mesh (.glb / .obj / .ply / .stl)",
+ file_types=[".glb", ".gltf", ".obj", ".ply", ".stl", ".off"],
+ type="filepath",
+ elem_id="input-mesh",
+ )
+
+ with gr.Accordion("⚙️ Generation Parameters", open=True, elem_classes="param-accordion"):
+ vert_num = gr.Slider(
+ label="Target Vertex Count",
+ minimum=200,
+ maximum=5000,
+ value=2000,
+ step=100,
+ info="Number of vertices in the generated mesh (200–5000)",
+ )
+ cfg_strength = gr.Slider(
+ label="CFG Strength",
+ minimum=0.0,
+ maximum=10.0,
+ value=3.0,
+ step=0.5,
+ info="Classifier-free guidance strength",
+ )
+ with gr.Row():
+ vflow_steps = gr.Slider(
+ label="V-Flow Steps",
+ minimum=4,
+ maximum=64,
+ value=24,
+ step=4,
+ info="Euler steps for vertex flow",
+ )
+ tflow_steps = gr.Slider(
+ label="T-Flow Steps",
+ minimum=10,
+ maximum=100,
+ value=50,
+ step=5,
+ info="Euler steps for topology flow",
+ )
+ seed = gr.Number(
+ label="Random Seed",
+ value=42,
+ precision=0,
+ info="Seed for reproducibility",
+ )
+
+ generate_btn = gr.Button(
+ "🚀 Generate 3D Mesh",
+ variant="primary",
+ size="lg",
+ elem_classes="generate-btn",
+ elem_id="generate-btn",
+ )
+
+ # ── Example meshes ────────────────────────────────────────
+ if os.path.isdir(EXAMPLES_DIR):
+ example_files = sorted(
+ os.path.join(EXAMPLES_DIR, f)
+ for f in os.listdir(EXAMPLES_DIR)
+ if os.path.splitext(f)[1].lower() in MESH_EXTENSIONS
+ )
+ if example_files:
+ gr.Markdown("### 📂 Example Meshes")
+ gr.Examples(
+ examples=[[None, f, 2000, 3.0, 24, 50, 42] for f in example_files],
+ inputs=[input_image, input_mesh, vert_num, cfg_strength, vflow_steps, tflow_steps, seed],
+ label="Click to load an example",
+ cache_examples=False,
+ )
+
+ # ── RIGHT: Outputs ────────────────────────────────────────────
+ with gr.Column(scale=2, min_width=600):
+ gr.Markdown("### 🖼️ 3D Output — Rerun Viewer")
+
+ rerun_viewer = Rerun(
+ streaming=False,
+ height=560,
+ elem_id="rerun-viewer",
+ )
+
+ gr.Markdown("---")
+ gr.Markdown("### 📦 Download")
+
+ glb_output = gr.File(
+ label="Generated GLB File",
+ type="filepath",
+ elem_id="glb-output",
+ interactive=False,
+ )
+
+ download_btn = gr.DownloadButton(
+ label="⬇️ Download GLB File",
+ size="lg",
+ elem_classes="download-btn",
+ elem_id="download-btn",
+ visible=False,
+ )
+
+ # ── Event wiring ──────────────────────────────────────────────────
+
+ def on_generate(image, mesh_path, vn, cfg, vfs, tfs, s):
+ rr_data, glb_path = generate_mesh(
+ input_image=image,
+ input_mesh_path=mesh_path,
+ vert_num=int(vn),
+ cfg_strength=float(cfg),
+ vflow_steps=int(vfs),
+ tflow_steps=int(tfs),
+ seed=int(s),
+ )
+ return (
+ rr_data,
+ glb_path,
+ gr.update(value=glb_path, visible=True),
+ )
+
+ generate_btn.click(
+ fn=on_generate,
+ inputs=[input_image, input_mesh, vert_num, cfg_strength, vflow_steps, tflow_steps, seed],
+ outputs=[rerun_viewer, glb_output, download_btn],
+ )
+
+ # ── Footer ────────────────────────────────────────────────────────
+ gr.HTML(
+ ''
+ '🔬 LATO.2 — Factorized 3D Mesh Generation with Vertex & Topology Flow
'
+ '
Hang Long, Tianhao Zhao et al. • '
+ 'arXiv 2607.10623 • '
+ '🤗 Model'
+ '
'
+ )
+
+ return app
+
+
+# ═══════════════════════════════════════════════════════════════════════════════
+# Entry point
+# ═══════════════════════════════════════════════════════════════════════════════
+
+def parse_app_args():
+ p = argparse.ArgumentParser(description="LATO.2 Gradio App")
+ p.add_argument("--port", type=int, default=7860, help="Port to serve on")
+ p.add_argument("--share", action="store_true", help="Create a public Gradio link")
+ p.add_argument("--server_name", default="0.0.0.0", help="Server bind address")
+ return p.parse_args()
+
+
+if __name__ == "__main__":
+ args = parse_app_args()
+ app = build_app()
+ app.queue(max_size=4)
+ app.launch(
+ server_name=args.server_name,
+ server_port=args.port,
+ share=args.share,
+ show_error=True,
+ )
diff --git a/assets/example_mesh/crocodile.glb b/assets/example_mesh/crocodile.glb
new file mode 100644
index 0000000000000000000000000000000000000000..617df3bb4c6c0ad01f727a76c362f1080248bbed
--- /dev/null
+++ b/assets/example_mesh/crocodile.glb
@@ -0,0 +1,3 @@
+version https://git-lfs.github.com/spec/v1
+oid sha256:6f6e3db8a580db37a6e9145f68b3cc765143a5b081d6dc6ce6713949cbad21ca
+size 1721020
diff --git a/assets/example_mesh/dragon.glb b/assets/example_mesh/dragon.glb
new file mode 100644
index 0000000000000000000000000000000000000000..67a599c1334792c1e76fb672916d67ddaa05088f
--- /dev/null
+++ b/assets/example_mesh/dragon.glb
@@ -0,0 +1,3 @@
+version https://git-lfs.github.com/spec/v1
+oid sha256:347d0a6f76b3b6afa34b13b75de14e73ab2e45ef5fdeb17c62a7d8216d358fd2
+size 1846248
diff --git a/assets/example_mesh/spaceman.glb b/assets/example_mesh/spaceman.glb
new file mode 100644
index 0000000000000000000000000000000000000000..9daec7257c0f7f71f28a73fd5dd9d87ec4e3334b
--- /dev/null
+++ b/assets/example_mesh/spaceman.glb
@@ -0,0 +1,3 @@
+version https://git-lfs.github.com/spec/v1
+oid sha256:20a48887b1d682b91f8a19e9ec9ec83e9bf0c402ef062ccdd26085a626b2eca1
+size 1516348
diff --git a/assets/teaser.png b/assets/teaser.png
new file mode 100644
index 0000000000000000000000000000000000000000..dc9c239f809b44bbb8e9ca526b9af783476d2d68
--- /dev/null
+++ b/assets/teaser.png
@@ -0,0 +1,3 @@
+version https://git-lfs.github.com/spec/v1
+oid sha256:3a8551b9f4b9b0c5becfa6225ca554237d10f3eea283e1b51790288e27776bc4
+size 12922537
diff --git a/dataset/mesh_render.py b/dataset/mesh_render.py
new file mode 100644
index 0000000000000000000000000000000000000000..fb229c0bb53fc0ad336d7e7ed52ee1bb42e90a5a
--- /dev/null
+++ b/dataset/mesh_render.py
@@ -0,0 +1,265 @@
+from typing import List, Optional, Sequence, Tuple, Union
+
+import numpy as np
+
+try:
+ import open3d as o3d
+except Exception as _e:
+ o3d = None
+ _OPEN3D_IMPORT_ERROR = _e
+
+
+ColorLike = Union[Sequence[float], np.ndarray]
+
+
+def _to_rgb01(color: ColorLike) -> Tuple[float, float, float]:
+ c = np.asarray(color, dtype=np.float64).reshape(-1)[:3]
+ if c.max() > 1.0 + 1e-6:
+ c = c / 255.0
+ return float(c[0]), float(c[1]), float(c[2])
+
+
+def _axis_index(up_axis: str) -> int:
+ return {"x": 0, "y": 1, "z": 2}[up_axis.lower()]
+
+
+def _orbit_eye(
+ center: np.ndarray,
+ distance: float,
+ azimuth_deg: float,
+ elevation_deg: float,
+ up_axis: str,
+) -> Tuple[np.ndarray, np.ndarray]:
+ az = np.deg2rad(azimuth_deg)
+ el = np.deg2rad(elevation_deg)
+ ce = np.cos(el)
+
+ horiz = distance * ce
+ vert = distance * np.sin(el)
+ ai = _axis_index(up_axis)
+ offset = np.zeros(3, dtype=np.float64)
+
+ plane_axes = [i for i in range(3) if i != ai]
+ offset[plane_axes[0]] = horiz * np.cos(az)
+ offset[plane_axes[1]] = horiz * np.sin(az)
+ offset[ai] = vert
+ eye = center + offset
+ up = np.zeros(3, dtype=np.float64)
+ up[ai] = 1.0
+ return eye, up
+
+
+class WhiteModelRenderer:
+ def __init__(
+ self,
+ img_res: int = 512,
+ mesh_color: ColorLike = (0.78, 0.78, 0.82),
+ bg_color: ColorLike = (1.0, 1.0, 1.0),
+ up_axis: str = "y",
+ add_ground: bool = True,
+ shadow: bool = True,
+ elevation_range: Tuple[float, float] = (15.0, 40.0),
+ azimuth_range: Tuple[float, float] = (0.0, 360.0),
+ camera_distance: float = 1.8,
+ fov: float = 50.0,
+ ground_color: ColorLike = (0.92, 0.92, 0.92),
+ sun_intensity: float = 90000.0,
+ ambient_intensity: float = 32000.0,
+ crop_to_object: bool = False,
+ crop_padding: float = 1.2,
+ ):
+ if o3d is None:
+ raise ImportError(
+ f"open3d is required for WhiteModelRenderer but failed to import: {_OPEN3D_IMPORT_ERROR}"
+ )
+ self.img_res = int(img_res)
+ self.mesh_color = _to_rgb01(mesh_color)
+ self.bg_color = _to_rgb01(bg_color)
+ self.up_axis = up_axis.lower()
+ self.add_ground = add_ground
+ self.shadow = shadow
+ self.elevation_range = elevation_range
+ self.azimuth_range = azimuth_range
+ self.camera_distance = float(camera_distance)
+ self.fov = float(fov)
+ self.ground_color = _to_rgb01(ground_color)
+ self.sun_intensity = float(sun_intensity)
+ self.ambient_intensity = float(ambient_intensity)
+ self.crop_to_object = crop_to_object
+ self.crop_padding = float(crop_padding)
+
+ self._renderer = None
+ self._rng = np.random.default_rng()
+
+ def _ensure_renderer(self):
+ if self._renderer is None:
+ self._renderer = o3d.visualization.rendering.OffscreenRenderer(
+ self.img_res, self.img_res
+ )
+ return self._renderer
+
+ def _make_o3d_mesh(self, vertices: np.ndarray, faces: np.ndarray):
+ mesh = o3d.geometry.TriangleMesh()
+ mesh.vertices = o3d.utility.Vector3dVector(
+ np.asarray(vertices, dtype=np.float64)
+ )
+ mesh.triangles = o3d.utility.Vector3iVector(np.asarray(faces, dtype=np.int32))
+ mesh.compute_vertex_normals()
+ return mesh
+
+ def _make_ground(self, mesh_min: np.ndarray, mesh_max: np.ndarray):
+ ai = _axis_index(self.up_axis)
+ center = (mesh_min + mesh_max) / 2.0
+ extent = float(np.max(mesh_max - mesh_min))
+ size = max(extent * 6.0, 4.0)
+
+ plane_axes = [i for i in range(3) if i != ai]
+ bottom = mesh_min[ai] - extent * 0.02
+
+ corners_2d = (
+ np.array([[-0.5, -0.5], [0.5, -0.5], [0.5, 0.5], [-0.5, 0.5]]) * size
+ )
+ verts = np.zeros((4, 3), dtype=np.float64)
+ for k, (a, b) in enumerate(corners_2d):
+ verts[k, plane_axes[0]] = center[plane_axes[0]] + a
+ verts[k, plane_axes[1]] = center[plane_axes[1]] + b
+ verts[k, ai] = bottom
+ tris = np.array([[0, 1, 2], [0, 2, 3]], dtype=np.int32)
+ ground = o3d.geometry.TriangleMesh()
+ ground.vertices = o3d.utility.Vector3dVector(verts)
+ ground.triangles = o3d.utility.Vector3iVector(tris)
+ ground.compute_vertex_normals()
+ return ground
+
+ def _lit_material(self, rgb: Tuple[float, float, float], roughness: float = 0.85):
+ mat = o3d.visualization.rendering.MaterialRecord()
+ mat.shader = "defaultLit"
+ mat.base_color = [rgb[0], rgb[1], rgb[2], 1.0]
+ mat.base_roughness = roughness
+ mat.base_metallic = 0.0
+ mat.base_reflectance = 0.4
+ return mat
+
+ def _setup_scene(
+ self,
+ vertices: np.ndarray,
+ faces: np.ndarray,
+ mesh_color: Tuple[float, float, float],
+ ):
+ renderer = self._ensure_renderer()
+ scene = renderer.scene
+ scene.clear_geometry()
+ scene.set_background(
+ [self.bg_color[0], self.bg_color[1], self.bg_color[2], 1.0]
+ )
+
+ mesh = self._make_o3d_mesh(vertices, faces)
+ scene.add_geometry("mesh", mesh, self._lit_material(mesh_color))
+
+ mesh_min = np.asarray(vertices, dtype=np.float64).min(axis=0)
+ mesh_max = np.asarray(vertices, dtype=np.float64).max(axis=0)
+ if self.add_ground:
+ ground = self._make_ground(mesh_min, mesh_max)
+ scene.add_geometry(
+ "ground", ground, self._lit_material(self.ground_color, roughness=0.95)
+ )
+
+ ai = _axis_index(self.up_axis)
+ sun_dir = np.array([0.35, 0.35, 0.35])
+ sun_dir[ai] = -1.0
+ sun_dir = sun_dir / np.linalg.norm(sun_dir)
+
+ scene.scene.set_sun_light(sun_dir.tolist(), [1.0, 1.0, 1.0], self.sun_intensity)
+ scene.scene.enable_sun_light(True)
+ scene.scene.set_indirect_light_intensity(self.ambient_intensity)
+
+ center = (mesh_min + mesh_max) / 2.0
+ return center
+
+ def _object_mask_from_depth(self):
+ renderer = self._renderer
+ depth = np.asarray(renderer.render_to_depth_image(z_in_view_space=True))
+ mask = np.isfinite(depth) & (depth > 0)
+ return mask
+
+ def _crop_resize_to_object(self, rgb: np.ndarray, mask: np.ndarray) -> np.ndarray:
+ from PIL import Image as _Image
+
+ ys, xs = np.where(mask)
+ if xs.size == 0:
+ out = _Image.fromarray(rgb).resize(
+ (self.img_res, self.img_res), _Image.LANCZOS
+ )
+ return np.asarray(out)
+
+ x0, y0, x1, y1 = xs.min(), ys.min(), xs.max(), ys.max()
+ cx, cy = (x0 + x1) / 2.0, (y0 + y1) / 2.0
+ size = int(max(x1 - x0, y1 - y0) * self.crop_padding)
+ size = max(size, 1)
+ half = size // 2
+ bx0, by0, bx1, by1 = (
+ int(round(cx - half)),
+ int(round(cy - half)),
+ int(round(cx - half)) + size,
+ int(round(cy - half)) + size,
+ )
+
+ H, W = rgb.shape[:2]
+ canvas = np.zeros((size, size, 3), dtype=np.uint8)
+ sx0, sy0 = max(0, bx0), max(0, by0)
+ sx1, sy1 = min(W, bx1), min(H, by1)
+ if sx1 > sx0 and sy1 > sy0:
+ canvas[sy0 - by0 : sy1 - by0, sx0 - bx0 : sx1 - bx0] = rgb[sy0:sy1, sx0:sx1]
+
+ out = _Image.fromarray(canvas).resize(
+ (self.img_res, self.img_res), _Image.LANCZOS
+ )
+ return np.ascontiguousarray(np.asarray(out))
+
+ def render(
+ self,
+ vertices: np.ndarray,
+ faces: np.ndarray,
+ num_views: int = 1,
+ mesh_color: Optional[ColorLike] = None,
+ azimuths: Optional[Sequence[float]] = None,
+ elevations: Optional[Sequence[float]] = None,
+ seed: Optional[int] = None,
+ ) -> Tuple[List[np.ndarray], List[dict]]:
+ rng = np.random.default_rng(seed) if seed is not None else self._rng
+ rgb = self.mesh_color if mesh_color is None else _to_rgb01(mesh_color)
+
+ center = self._setup_scene(vertices, faces, rgb)
+ renderer = self._renderer
+
+ images: List[np.ndarray] = []
+ params: List[dict] = []
+ for v in range(num_views):
+ if azimuths is not None:
+ az = float(azimuths[v])
+ else:
+ az = float(rng.uniform(*self.azimuth_range))
+ if elevations is not None:
+ el = float(elevations[v])
+ else:
+ el = float(rng.uniform(*self.elevation_range))
+
+ eye, up = _orbit_eye(center, self.camera_distance, az, el, self.up_axis)
+ renderer.setup_camera(self.fov, center.tolist(), eye.tolist(), up.tolist())
+
+ img = renderer.render_to_image()
+ arr = np.asarray(img)
+ if arr.ndim == 3 and arr.shape[2] == 4:
+ arr = arr[:, :, :3]
+ arr = arr.astype(np.uint8)
+
+ if self.crop_to_object:
+ mask = self._object_mask_from_depth()
+ arr = self._crop_resize_to_object(arr, mask)
+
+ images.append(np.ascontiguousarray(arr))
+ params.append(
+ {"azimuth": az, "elevation": el, "distance": self.camera_distance}
+ )
+
+ return images, params
diff --git a/dataset/topo_dataset.py b/dataset/topo_dataset.py
new file mode 100644
index 0000000000000000000000000000000000000000..2af3604ffa0c9a7e27f91ec73becb26645c27348
--- /dev/null
+++ b/dataset/topo_dataset.py
@@ -0,0 +1,88 @@
+from __future__ import annotations
+
+import os
+import traceback
+from typing import Dict, List, Optional
+
+import numpy as np
+import torch
+
+from dataset.utils import (
+ MESH_EXTENSIONS,
+ dedup_quantized_mesh,
+ extract_active_voxels,
+ quantize_mesh_clustering,
+)
+
+
+class TopoVoxelDataset(torch.utils.data.Dataset):
+ def __init__(
+ self,
+ root_dir: str,
+ num_discrete: int = 1024,
+ voxel_res: int = 64,
+ max_vertices: Optional[int] = None,
+ num_samples: Optional[int] = None,
+ ):
+ self.root_dir = root_dir
+ self.num_discrete = int(num_discrete)
+ self.voxel_res = int(voxel_res)
+ self.max_vertices = max_vertices
+ self.files = sorted(
+ f
+ for f in os.listdir(root_dir)
+ if os.path.splitext(f)[1].lower() in MESH_EXTENSIONS
+ )
+ if num_samples is not None:
+ self.files = self.files[:num_samples]
+ if not self.files:
+ raise ValueError(f"no mesh files ({MESH_EXTENSIONS}) under {root_dir}")
+
+ def __len__(self) -> int:
+ return len(self.files)
+
+ def __getitem__(self, idx: int) -> Dict:
+ name = os.path.splitext(self.files[idx])[0]
+ path = os.path.join(self.root_dir, self.files[idx])
+ res = self.num_discrete
+ try:
+ quantized = quantize_mesh_clustering(path, resolution=res)
+ if quantized is None:
+ return {"name": name, "error": "empty mesh"}
+ v_int, offsets, faces = quantized
+ if len(faces) < 1 or len(v_int) < 3:
+ return {"name": name, "error": "degenerate mesh after quantization"}
+
+ gt_int, _, gt_faces = dedup_quantized_mesh(v_int, offsets, faces, res)
+ num_gt = len(gt_int)
+ if num_gt < 3 or len(gt_faces) < 1:
+ return {"name": name, "error": f"too few vertices/faces ({num_gt})"}
+ if self.max_vertices is not None and num_gt > self.max_vertices:
+ return {
+ "name": name,
+ "error": f"vertex count {num_gt} exceeds max_vertices={self.max_vertices}",
+ }
+
+ quant_v = gt_int.astype(np.float64) / (res - 1.0) - 0.5
+ quant_v = np.clip(quant_v, -0.5 + 1e-6, 0.5 - 1e-6).astype(np.float32)
+ voxel_coords = extract_active_voxels(quant_v, gt_faces, self.voxel_res)
+
+ return {
+ "name": name,
+ "vertices": torch.from_numpy(gt_int.astype(np.int64)),
+ "voxel_coords": voxel_coords.long(),
+ }
+ except Exception as e:
+ return {"name": name, "error": f"{e}\n{traceback.format_exc()}"}
+
+
+def collate_fn(batch: List[Dict]) -> Dict:
+ errors = [b for b in batch if "error" in b]
+ good = [b for b in batch if "error" not in b]
+ collated: Dict = {"errors": errors}
+ if not good:
+ return collated
+ collated["name"] = [b["name"] for b in good]
+ collated["vertices"] = [b["vertices"] for b in good]
+ collated["voxel_coords"] = [b["voxel_coords"] for b in good]
+ return collated
diff --git a/dataset/utils.py b/dataset/utils.py
new file mode 100644
index 0000000000000000000000000000000000000000..7861475105ad2a3b6184b56445307045c70e1e81
--- /dev/null
+++ b/dataset/utils.py
@@ -0,0 +1,283 @@
+import numpy as np
+import torch
+import trimesh
+from trimesh import grouping
+
+from o_voxel.convert import mesh_to_flexible_dual_grid
+
+MESH_EXTENSIONS = {".obj", ".glb", ".gltf", ".ply", ".stl", ".off"}
+
+
+def quantize_mesh_clustering(mesh_path: str, resolution: int = 1024):
+ mesh = trimesh.load(mesh_path, process=False, force="mesh")
+ if mesh is None or len(mesh.vertices) == 0 or len(mesh.faces) == 0:
+ return None
+
+ vertices = np.asarray(mesh.vertices, dtype=np.float64)
+ faces = np.asarray(mesh.faces, dtype=np.int64)
+
+ bbox_min, bbox_max = vertices.min(axis=0), vertices.max(axis=0)
+ center = (bbox_min + bbox_max) / 2.0
+ max_extent = max(float((bbox_max - bbox_min).max()), 1e-7)
+ normalized_v = (vertices - center) / max_extent + 0.5 # [0, 1]
+
+ v_grid = np.clip(np.floor(normalized_v * resolution), 0, resolution - 1).astype(
+ np.int64
+ )
+ v_hash = (
+ v_grid[:, 0] * resolution * resolution
+ + v_grid[:, 1] * resolution
+ + v_grid[:, 2]
+ )
+
+ unique_hashes, inverse = np.unique(v_hash, return_inverse=True)
+ num_clusters = len(unique_hashes)
+
+ v_sum = np.zeros((num_clusters, 3), dtype=np.float64)
+ counts = np.zeros(num_clusters, dtype=np.float64)
+ np.add.at(v_sum, inverse, normalized_v)
+ np.add.at(counts, inverse, 1)
+ v_mean = v_sum / counts[:, None]
+
+ v_int = np.clip(np.floor(v_mean * resolution), 0, resolution - 1).astype(np.int32)
+
+ # Offset of the cluster mean relative to the voxel center, scaled to (-1, 1).
+ voxel_center = (v_int.astype(np.float64) + 0.5) / float(resolution)
+ offsets = ((v_mean - voxel_center) * 2.0 * resolution).astype(np.float32)
+ offsets = np.clip(offsets, -1.0, 1.0)
+
+ new_faces = inverse[faces]
+ valid = (
+ (new_faces[:, 0] != new_faces[:, 1])
+ & (new_faces[:, 1] != new_faces[:, 2])
+ & (new_faces[:, 2] != new_faces[:, 0])
+ )
+ return v_int, offsets, new_faces[valid]
+
+
+def realign_offsets(
+ orig_int: np.ndarray, orig_off: np.ndarray, kept_int: np.ndarray, resolution: int
+) -> np.ndarray:
+ base = resolution * resolution
+ orig_i = orig_int.astype(np.int64)
+ kept_i = kept_int.astype(np.int64)
+ orig_h = orig_i[:, 0] * base + orig_i[:, 1] * resolution + orig_i[:, 2]
+ kept_h = kept_i[:, 0] * base + kept_i[:, 1] * resolution + kept_i[:, 2]
+
+ order = np.argsort(orig_h)
+ sorted_h = orig_h[order]
+ sorted_off = orig_off.astype(np.float64)[order]
+ cumsum = np.concatenate([np.zeros((1, 3)), np.cumsum(sorted_off, axis=0)], axis=0)
+ left = np.searchsorted(sorted_h, kept_h, side="left")
+ right = np.searchsorted(sorted_h, kept_h, side="right")
+ counts = (right - left).clip(min=1).astype(np.float64)
+ out = ((cumsum[right] - cumsum[left]) / counts[:, None]).astype(np.float32)
+ out = np.clip(out, -1.0, 1.0)
+ out[right <= left] = 0.0
+ return out
+
+
+def dedup_quantized_mesh(
+ v_int: np.ndarray, offsets: np.ndarray, faces: np.ndarray, resolution: int
+):
+ tmesh = trimesh.Trimesh(vertices=v_int, faces=faces, process=False)
+ tmesh.merge_vertices()
+ tmesh.update_faces(tmesh.nondegenerate_faces())
+ tmesh.update_faces(tmesh.unique_faces())
+ tmesh.remove_unreferenced_vertices()
+
+ gt_int = np.asarray(tmesh.vertices).astype(np.int32)
+ gt_faces = np.asarray(tmesh.faces, dtype=np.int64)
+ gt_offsets = realign_offsets(v_int, offsets, gt_int, resolution)
+ return gt_int, gt_offsets, gt_faces
+
+
+def extract_active_voxels(
+ vertices: np.ndarray, faces: np.ndarray, resolution: int
+) -> torch.Tensor:
+ coords, *_ = mesh_to_flexible_dual_grid(
+ vertices=torch.as_tensor(vertices * 0.99999, dtype=torch.float32).contiguous(),
+ faces=torch.as_tensor(faces, dtype=torch.int32).contiguous(),
+ grid_size=resolution,
+ aabb=torch.tensor([[-0.5, -0.5, -0.5], [0.5, 0.5, 0.5]], dtype=torch.float32),
+ )
+ coords = coords.cpu().long()
+ coords_1d = (
+ coords[:, 0] * resolution * resolution
+ + coords[:, 1] * resolution
+ + coords[:, 2]
+ )
+ return coords[torch.argsort(coords_1d)].int()
+
+
+def union_voxels(a: torch.Tensor, b: torch.Tensor, resolution: int) -> torch.Tensor:
+ if a.numel() == 0:
+ return b.int().clone()
+ if b.numel() == 0:
+ return a.int().clone()
+ combined = torch.cat([a.reshape(-1, 3).long(), b.reshape(-1, 3).long()], dim=0)
+ combined = combined.clamp(0, resolution - 1)
+ hashes = (
+ combined[:, 0] * resolution * resolution
+ + combined[:, 1] * resolution
+ + combined[:, 2]
+ )
+ sorted_hashes, sort_idx = torch.sort(hashes)
+ keep = torch.ones_like(sorted_hashes, dtype=torch.bool)
+ keep[1:] = sorted_hashes[1:] != sorted_hashes[:-1]
+ return combined[sort_idx[keep]].int()
+
+
+def _sample_surface_uniform(tm_mesh: trimesh.Trimesh, n_samples: int):
+ face_idx = np.random.choice(len(tm_mesh.faces), size=n_samples, replace=True)
+ tri = tm_mesh.vertices[tm_mesh.faces[face_idx]] # (N, 3, 3)
+ u = np.random.rand(n_samples, 1)
+ v = np.random.rand(n_samples, 1)
+ sqrt_u = np.sqrt(u)
+ points = (
+ (1 - sqrt_u) * tri[:, 0]
+ + (sqrt_u * (1 - v)) * tri[:, 1]
+ + (sqrt_u * v) * tri[:, 2]
+ )
+ normals = tm_mesh.face_normals[face_idx]
+ return points.astype(np.float32), normals.astype(np.float32), face_idx
+
+
+def _sample_edges_dora(
+ tm_mesh: trimesh.Trimesh, n_len_samples: int, n_uniform_samples: int
+):
+ parts_start, parts_end, parts_norm, parts_virt = [], [], [], []
+
+ adj_faces = tm_mesh.face_adjacency
+ adj_edges = tm_mesh.face_adjacency_edges
+ if len(adj_faces) > 0:
+ n0 = tm_mesh.face_normals[adj_faces[:, 0]]
+ n1 = tm_mesh.face_normals[adj_faces[:, 1]]
+ sum_normals = n0 + n1
+ norms = np.linalg.norm(sum_normals, axis=1, keepdims=True)
+ norms[norms < 1e-6] = 1.0
+
+ faces_pair = tm_mesh.faces[adj_faces]
+ unique_idx_0 = np.sum(faces_pair, axis=2)[:, 0] - np.sum(adj_edges, axis=1)
+ unique_idx_1 = np.sum(faces_pair, axis=2)[:, 1] - np.sum(adj_edges, axis=1)
+ virtual = (
+ tm_mesh.vertices[unique_idx_0] + tm_mesh.vertices[unique_idx_1]
+ ) * 0.5
+
+ parts_start.append(tm_mesh.vertices[adj_edges[:, 0]])
+ parts_end.append(tm_mesh.vertices[adj_edges[:, 1]])
+ parts_norm.append(sum_normals / norms)
+ parts_virt.append(virtual)
+
+ edges_sorted = tm_mesh.edges_sorted
+ if len(edges_sorted) > 0:
+ boundary_group = grouping.group_rows(edges_sorted, require_count=1)
+ if len(boundary_group) > 0:
+ boundary_indices = np.concatenate(
+ [np.atleast_1d(g) for g in boundary_group]
+ )
+ face_indices = boundary_indices // 3
+ edge_v = edges_sorted[boundary_indices]
+ unique_idx = np.sum(tm_mesh.faces[face_indices], axis=1) - np.sum(
+ edge_v, axis=1
+ )
+
+ parts_start.append(tm_mesh.vertices[edge_v[:, 0]])
+ parts_end.append(tm_mesh.vertices[edge_v[:, 1]])
+ parts_norm.append(tm_mesh.face_normals[face_indices])
+ parts_virt.append(tm_mesh.vertices[unique_idx])
+
+ if not parts_start:
+ return None, None, None
+
+ v_start = np.concatenate(parts_start, axis=0)
+ v_end = np.concatenate(parts_end, axis=0)
+ normals = np.concatenate(parts_norm, axis=0)
+ v_virtual = np.concatenate(parts_virt, axis=0)
+
+ lengths = np.linalg.norm(v_end - v_start, axis=1)
+ total = lengths.sum()
+ num_edges = len(lengths)
+ probs_len = (
+ lengths / total if total >= 1e-9 else np.full(num_edges, 1.0 / num_edges)
+ )
+ probs_len = probs_len / probs_len.sum()
+
+ chosen = np.concatenate(
+ [
+ np.random.choice(num_edges, size=n_len_samples, p=probs_len),
+ np.random.choice(num_edges, size=n_uniform_samples),
+ ]
+ )
+ t = np.random.rand(len(chosen), 1)
+ points = v_start[chosen] + (v_end[chosen] - v_start[chosen]) * t
+ triplets = np.stack([v_start[chosen], v_end[chosen], v_virtual[chosen]], axis=1)
+ return (
+ points.astype(np.float32),
+ normals[chosen].astype(np.float32),
+ triplets.astype(np.float32),
+ )
+
+
+def _vdf_from_triplets(
+ points: np.ndarray, triplets: np.ndarray, normalize: bool
+) -> np.ndarray:
+ view_dtype = np.dtype((np.void, triplets.dtype.itemsize * triplets.shape[-1]))
+ v_view = triplets.view(view_dtype).squeeze(-1)
+ sort_idx = np.argsort(v_view, axis=1)
+ v_sorted = triplets[np.arange(triplets.shape[0])[:, None], sort_idx]
+
+ dirs = v_sorted - points[:, None, :] # (N, 3, 3)
+ if normalize:
+ dirs = dirs / (np.linalg.norm(dirs, axis=-1, keepdims=True) + 1e-8)
+ return dirs.reshape(len(points), 9).astype(np.float32)
+
+
+def sample_point_features(
+ tm_mesh: trimesh.Trimesh,
+ n_samples: int,
+ sample_type: str = "dora",
+ normalize_vdf: bool = True,
+) -> torch.Tensor:
+ # (N, 15) float32 point features: [xyz(3), normal(3), vdf(9)].
+ vertices = np.asarray(tm_mesh.vertices, dtype=np.float64)
+ faces = np.asarray(tm_mesh.faces)
+
+ if sample_type == "dora":
+ n_surf_area = n_samples // 4
+ n_surf_uniform = n_samples // 4
+ n_edge_len = n_samples // 4
+ n_edge_uniform = n_samples - n_surf_area - n_surf_uniform - n_edge_len
+
+ p_edge, n_edge, triplets_edge = _sample_edges_dora(
+ tm_mesh, n_edge_len, n_edge_uniform
+ )
+ if p_edge is None:
+ n_surf_area += n_edge_len
+ n_surf_uniform += n_edge_uniform
+ elif sample_type == "uniform":
+ n_surf_area, n_surf_uniform = n_samples, 0
+ p_edge = None
+ else:
+ raise ValueError(f"unknown sample_type: {sample_type!r}")
+
+ p_area, idx_area = tm_mesh.sample(n_surf_area, return_index=True)
+ n_area = tm_mesh.face_normals[idx_area]
+ if n_surf_uniform > 0:
+ p_unif, n_unif, idx_unif = _sample_surface_uniform(tm_mesh, n_surf_uniform)
+ points = np.concatenate([p_area, p_unif], axis=0).astype(np.float32)
+ normals = np.concatenate([n_area, n_unif], axis=0).astype(np.float32)
+ idx_surf = np.concatenate([idx_area, idx_unif], axis=0)
+ else:
+ points = p_area.astype(np.float32)
+ normals = n_area.astype(np.float32)
+ idx_surf = idx_area
+ triplets = vertices[faces[idx_surf]].astype(np.float32)
+
+ if p_edge is not None:
+ points = np.concatenate([points, p_edge], axis=0)
+ normals = np.concatenate([normals, n_edge], axis=0)
+ triplets = np.concatenate([triplets, triplets_edge], axis=0)
+
+ vdf = _vdf_from_triplets(points, triplets, normalize=normalize_vdf)
+ return torch.from_numpy(np.concatenate([points, normals, vdf], axis=-1))
diff --git a/dataset/voxel_dataset.py b/dataset/voxel_dataset.py
new file mode 100644
index 0000000000000000000000000000000000000000..46f916d1a9ad943ea12607a020f7c3a5f128249a
--- /dev/null
+++ b/dataset/voxel_dataset.py
@@ -0,0 +1,211 @@
+import os
+from typing import Dict, List, Optional
+import numpy as np
+import torch
+import trimesh
+from torch.utils.data import Dataset
+
+from dataset.utils import (
+ MESH_EXTENSIONS,
+ dedup_quantized_mesh,
+ extract_active_voxels,
+ quantize_mesh_clustering,
+ sample_point_features,
+ union_voxels,
+)
+
+
+class VoxelVertexDataset(Dataset):
+ def __init__(
+ self,
+ root_dir: str,
+ resolution: int = 1024,
+ min_resolution: int = 64,
+ pc_sample_number: int = 819200,
+ sample_type: str = "dora",
+ normalize_vdf: bool = True,
+ need_encoder_inputs: bool = True,
+ min_vertices: int = 0,
+ max_vertices: Optional[int] = None,
+ num_samples: Optional[int] = None,
+ render: bool = False,
+ img_res: int = 518,
+ render_azimuth: float = 45.0,
+ render_elevation: float = 30.0,
+ ):
+ self.root_dir = root_dir
+ self.resolution = resolution
+ self.min_resolution = min_resolution
+ self.pc_sample_number = pc_sample_number
+ self.sample_type = sample_type
+ self.normalize_vdf = normalize_vdf
+ self.need_encoder_inputs = need_encoder_inputs
+ self.min_vertices = min_vertices
+ self.max_vertices = max_vertices
+ self.render = render
+ self.img_res = img_res
+ self.render_azimuth = render_azimuth
+ self.render_elevation = render_elevation
+
+ self.files = sorted(
+ f
+ for f in os.listdir(root_dir)
+ if os.path.splitext(f)[1].lower() in MESH_EXTENSIONS
+ )
+ if num_samples is not None:
+ self.files = self.files[:num_samples]
+ if not self.files:
+ raise ValueError(f"no mesh files ({MESH_EXTENSIONS}) under {root_dir}")
+
+ self._renderer = None # lazy: one EGL context per DataLoader worker
+
+ def __len__(self) -> int:
+ return len(self.files)
+
+ def _render_image(self, vertices: np.ndarray, faces: np.ndarray) -> np.ndarray:
+ if self._renderer is None:
+ from dataset.mesh_render import WhiteModelRenderer
+
+ self._renderer = WhiteModelRenderer(
+ img_res=self.img_res,
+ mesh_color=(0.78, 0.78, 0.82),
+ bg_color=(0.0, 0.0, 0.0),
+ up_axis="y",
+ add_ground=False,
+ shadow=True,
+ crop_to_object=True,
+ crop_padding=1.2,
+ )
+ imgs, _ = self._renderer.render(
+ np.asarray(vertices, dtype=np.float64),
+ np.asarray(faces, dtype=np.int64),
+ num_views=1,
+ azimuths=[self.render_azimuth],
+ elevations=[self.render_elevation],
+ )
+ return imgs[0] # (img_res, img_res, 3) uint8
+
+ def __getitem__(self, idx: int) -> Dict:
+ name = os.path.splitext(self.files[idx])[0]
+ path = os.path.join(self.root_dir, self.files[idx])
+ res = self.resolution
+ min_res = self.min_resolution
+ try:
+ quantized = quantize_mesh_clustering(path, resolution=res)
+ if quantized is None:
+ return {"name": name, "error": "empty mesh"}
+ v_int, offsets, faces = quantized
+ if len(faces) < 1 or len(v_int) < 3:
+ return {"name": name, "error": "degenerate mesh after quantization"}
+
+ gt_int, gt_offsets, gt_faces = dedup_quantized_mesh(
+ v_int, offsets, faces, res
+ )
+ num_gt = len(gt_int)
+ if num_gt < max(self.min_vertices, 3) or len(gt_faces) < 1:
+ return {
+ "name": name,
+ "error": f"too few vertices/faces after dedup ({num_gt})",
+ }
+ if self.max_vertices is not None and num_gt > self.max_vertices:
+ return {
+ "name": name,
+ "error": f"vertex count {num_gt} exceeds max_vertices={self.max_vertices}",
+ }
+
+ quant_v = gt_int.astype(np.float64) / (res - 1.0) - 0.5
+ quant_v = np.clip(quant_v, -0.5 + 1e-6, 0.5 - 1e-6).astype(np.float32)
+ tmesh = trimesh.Trimesh(vertices=quant_v, faces=gt_faces, process=False)
+
+ if self.need_encoder_inputs:
+ vertex_added_active = extract_active_voxels(quant_v, gt_faces, res)
+ vertex_added_active = union_voxels(
+ vertex_added_active, torch.from_numpy(gt_int), res
+ )
+ point_cloud = sample_point_features(
+ tmesh,
+ self.pc_sample_number,
+ sample_type=self.sample_type,
+ normalize_vdf=self.normalize_vdf,
+ )
+ else:
+ vertex_added_active = torch.zeros((0, 3), dtype=torch.int32)
+ point_cloud = torch.zeros((0, 15), dtype=torch.float32)
+
+ min_active = extract_active_voxels(quant_v, gt_faces, self.min_resolution)
+
+ data = {
+ "name": name,
+ f"gt_vertex_voxels_{res}": torch.from_numpy(gt_int),
+ f"gt_vertex_offsets_{res}": torch.from_numpy(gt_offsets),
+ "quantized_vertices": torch.from_numpy(quant_v),
+ "quantized_faces": torch.from_numpy(gt_faces),
+ f"vertex_added_active_voxels_{res}": vertex_added_active,
+ f"point_cloud_{res}": point_cloud,
+ f"active_voxels_{min_res}": min_active,
+ }
+
+ # Raw mesh, bbox-normalized into the same [-0.5, 0.5] frame.
+ raw = trimesh.load(path, process=False, force="mesh")
+ raw_v = np.asarray(raw.vertices, dtype=np.float64)
+ center = (raw_v.min(axis=0) + raw_v.max(axis=0)) / 2.0
+ extent = max(float((raw_v.max(axis=0) - raw_v.min(axis=0)).max()), 1e-7)
+ data["original_vertices"] = torch.from_numpy(
+ ((raw_v - center) / extent).astype(np.float32)
+ )
+ data["original_faces"] = torch.from_numpy(
+ np.asarray(raw.faces, dtype=np.int64)
+ )
+
+ if self.render:
+ render_v = v_int.astype(np.float64) / res - 0.5
+ data["image"] = self._render_image(render_v, faces)
+
+ return data
+ except Exception as e:
+ import traceback
+
+ return {"name": name, "error": f"{e}\n{traceback.format_exc()}"}
+
+
+def collate_fn(
+ batch: List[Dict], resolution: int = 1024, min_resolution: int = 64
+) -> Dict:
+ res = resolution
+ min_res = min_resolution
+ errors = [b for b in batch if "error" in b]
+ batch = [b for b in batch if "error" not in b]
+ collated: Dict = {"errors": errors}
+ if not batch:
+ return collated
+
+ collated["name"] = [b["name"] for b in batch]
+ for key in (
+ "quantized_vertices",
+ "quantized_faces",
+ "original_vertices",
+ "original_faces",
+ ):
+ collated[key] = [b[key] for b in batch]
+ if "image" in batch[0]:
+ collated["image"] = [b["image"] for b in batch]
+
+ for key in (
+ f"gt_vertex_voxels_{res}",
+ f"vertex_added_active_voxels_{res}",
+ f"active_voxels_{min_res}",
+ ):
+ rows = []
+ for i, b in enumerate(batch):
+ coords = b[key]
+ batch_idx = torch.full((coords.shape[0], 1), i, dtype=torch.int32)
+ rows.append(torch.cat([batch_idx, coords], dim=1))
+ collated[key] = torch.cat(rows, dim=0)
+
+ collated[f"gt_vertex_offsets_{res}"] = torch.cat(
+ [b[f"gt_vertex_offsets_{res}"] for b in batch], dim=0
+ )
+ collated[f"point_cloud_{res}"] = torch.stack(
+ [b[f"point_cloud_{res}"] for b in batch], dim=0
+ )
+ return collated
diff --git a/models/__init__.py b/models/__init__.py
new file mode 100644
index 0000000000000000000000000000000000000000..ea4515a1c9bb40dbcc9ff6f9bc1b6c3c8c7f9017
--- /dev/null
+++ b/models/__init__.py
@@ -0,0 +1,9 @@
+from models.offset_head import OffsetHead
+from models.vdf_encoder import VDFEncoder
+from models.vertex_autoencoder import VertexVAE
+from models.dino_encoder import DinoV2Encoder
+from models.vertex_structured_flow import VertexSLatFlowModel
+from models.flow_sampler import VertFlowEulerCfgSampler, TopoFlowEulerSampler
+from models.topo_autoencoder import TopologyVAE
+from models.topo_flow import TopologySiTFlow
+from models.voxel_encoder import VoxelFieldConditioner
\ No newline at end of file
diff --git a/models/dino_encoder.py b/models/dino_encoder.py
new file mode 100644
index 0000000000000000000000000000000000000000..91be06b4c41bb083ed563373596ddb03bdc0e7cb
--- /dev/null
+++ b/models/dino_encoder.py
@@ -0,0 +1,73 @@
+import os
+from typing import Union
+
+import numpy as np
+import torch
+import torch.nn as nn
+import torch.nn.functional as F
+
+DINO_GITHUB_REPO = "facebookresearch/dinov2"
+DINO_LOCAL_REPO_DIRNAME = "facebookresearch_dinov2_main"
+
+
+class DinoV2Encoder(nn.Module):
+ def __init__(
+ self,
+ model_name: str,
+ hub_dir: str,
+ img_res: int,
+ ):
+ super().__init__()
+ self.img_res = int(img_res)
+
+ hub_dir = os.path.abspath(os.path.expanduser(hub_dir))
+ os.makedirs(hub_dir, exist_ok=True)
+ torch.hub.set_dir(hub_dir)
+ local_repo = os.path.join(hub_dir, DINO_LOCAL_REPO_DIRNAME)
+ if os.path.isdir(local_repo):
+ self.backbone = torch.hub.load(
+ local_repo, model_name, source="local", pretrained=True
+ )
+ else:
+ self.backbone = torch.hub.load(
+ DINO_GITHUB_REPO, model_name, source="github", pretrained=True
+ )
+ self.backbone.eval()
+ for p in self.backbone.parameters():
+ p.requires_grad_(False)
+
+ self.register_buffer(
+ "img_mean",
+ torch.tensor([0.485, 0.456, 0.406]).view(1, 3, 1, 1),
+ persistent=False,
+ )
+ self.register_buffer(
+ "img_std",
+ torch.tensor([0.229, 0.224, 0.225]).view(1, 3, 1, 1),
+ persistent=False,
+ )
+
+ @torch.no_grad()
+ def forward(self, images: Union[np.ndarray, torch.Tensor]) -> torch.Tensor:
+ """Images -> layer-normed DINO-v2 patch tokens (B, L, C).
+
+ Accepts uint8 channels-last images — (H, W, 3) or (B, H, W, 3) in
+ [0, 255], the dataset render format — or float channels-first
+ (B, 3, H, W) in [0, 1]. Any resolution; resized to ``img_res``.
+ """
+ if isinstance(images, np.ndarray):
+ images = torch.from_numpy(np.ascontiguousarray(images))
+ if images.dim() == 3:
+ images = images[None]
+ if images.dtype == torch.uint8:
+ images = images.permute(0, 3, 1, 2).float() / 255.0
+
+ x = images.to(device=self.img_mean.device, dtype=torch.float32)
+ if x.shape[-2:] != (self.img_res, self.img_res):
+ x = F.interpolate(
+ x, (self.img_res, self.img_res), mode="bicubic", align_corners=False
+ )
+ x = (x - self.img_mean) / self.img_std
+
+ feats = self.backbone(x, is_training=True)["x_prenorm"]
+ return F.layer_norm(feats, feats.shape[-1:])
diff --git a/models/flow_sampler.py b/models/flow_sampler.py
new file mode 100644
index 0000000000000000000000000000000000000000..918c07c14908d8bf10d78606ed56d8091e92a35c
--- /dev/null
+++ b/models/flow_sampler.py
@@ -0,0 +1,58 @@
+import numpy as np
+import torch
+
+
+class VertFlowEulerCfgSampler:
+ def _pred_v(self, model, x_t, t, cond, **kwargs):
+ t_vec = torch.full(
+ (x_t.shape[0],), 1000.0 * t, device=x_t.device, dtype=torch.float32
+ )
+ return model(x_t, t_vec, cond, **kwargs)
+
+ @torch.no_grad()
+ def sample(
+ self,
+ model,
+ noise,
+ cond,
+ neg_cond,
+ steps=12,
+ cfg_strength=3.0,
+ rescale_t=1.0,
+ **kwargs,
+ ):
+ x = noise
+ t_seq = np.linspace(1.0, 0.0, steps + 1)
+ t_seq = rescale_t * t_seq / (1 + (rescale_t - 1) * t_seq)
+ for i in range(steps):
+ t, t_prev = float(t_seq[i]), float(t_seq[i + 1])
+ v_cond = self._pred_v(model, x, t, cond, **kwargs)
+ v_uncond = self._pred_v(model, x, t, neg_cond, **kwargs)
+ v = (1 + cfg_strength) * v_cond - cfg_strength * v_uncond
+ x = x - (t - t_prev) * v
+ return x
+
+
+class TopoFlowEulerSampler:
+ def _pred_v(self, model, x_t, t, verts, mask, cond, cond_mask):
+ t_vec = torch.full((x_t.shape[0],), t, device=x_t.device, dtype=torch.float32)
+ return model(x_t, t_vec, verts=verts, mask=mask, cond=cond, cond_mask=cond_mask)
+
+ @torch.no_grad()
+ def sample(
+ self,
+ model,
+ noise,
+ verts,
+ mask,
+ cond=None,
+ cond_mask=None,
+ steps=50,
+ ):
+ x = noise
+ t_seq = np.linspace(0.0, 1.0, steps + 1)
+ for i in range(steps):
+ t, t_next = float(t_seq[i]), float(t_seq[i + 1])
+ v = self._pred_v(model, x, t, verts, mask, cond, cond_mask)
+ x = x + (t_next - t) * v
+ return x
diff --git a/models/offset_head.py b/models/offset_head.py
new file mode 100644
index 0000000000000000000000000000000000000000..7ecf13256b0e6e082151c29ceee3706a62e7510f
--- /dev/null
+++ b/models/offset_head.py
@@ -0,0 +1,22 @@
+import torch.nn as nn
+
+
+class OffsetHead(nn.Module):
+ def __init__(self, feat_dim: int, mlp_ratio: float = 4.0):
+ super().__init__()
+ self.mlp = nn.Sequential(
+ nn.Linear(feat_dim, int(feat_dim * mlp_ratio)),
+ nn.GELU(approximate="tanh"),
+ nn.Linear(int(feat_dim * mlp_ratio), 3),
+ nn.Tanh(),
+ )
+
+ def forward(self, vtx_feats):
+ """
+ Input:
+ vtx_feats: [N, feat_dim]
+ Output:
+ offsets: [N, 3], in range (-1, 1)
+ """
+ offsets = self.mlp(vtx_feats) # [N, 3], (-1, 1)
+ return offsets
diff --git a/models/topo_autoencoder.py b/models/topo_autoencoder.py
new file mode 100644
index 0000000000000000000000000000000000000000..7c4cd27033a5940d71716dcd9e96f206505e4b8c
--- /dev/null
+++ b/models/topo_autoencoder.py
@@ -0,0 +1,297 @@
+from __future__ import annotations
+
+from typing import List, Optional
+
+import numpy as np
+import torch
+from torch import nn
+
+from modules.pointnet import Pointnet
+from modules.transformer.hybrid import (
+ HybridGraphFlashStack,
+ FlashVarlenTransformerBlock,
+)
+from modules.utils import manual_cast, str_to_dtype
+from modules.transformer.blocks import (
+ PointEmbed,
+ RotaryPositionPhasesEmbedder,
+ MaskedTransformerCrossAttnBlock,
+)
+
+
+class TopologyEncoderHybrid(nn.Module):
+ def __init__(
+ self,
+ z_dim: int = 32,
+ hidden_dim: int = 384,
+ num_heads: int = 6,
+ num_discrete: int = 256,
+ dtype: str = "float32",
+ num_hybrid_stages: int = 2,
+ num_flash_per_stage: int = 1,
+ use_gradient_checkpointing: bool = False,
+ pc_cross_attn: bool = False,
+ ):
+ super().__init__()
+ self.dtype = str_to_dtype(dtype)
+ self.num_discrete = num_discrete
+ self.pc_cross_attn = bool(pc_cross_attn)
+
+ head_dim = hidden_dim // num_heads
+ self.rope = RotaryPositionPhasesEmbedder(head_dim=head_dim, dim=3)
+
+ self.backbone = HybridGraphFlashStack(
+ hidden_size=hidden_dim,
+ num_heads=num_heads,
+ num_stages=num_hybrid_stages,
+ num_flash_per_stage=num_flash_per_stage,
+ gradient_checkpointing=use_gradient_checkpointing,
+ )
+ if self.pc_cross_attn:
+ self.pc_cross_blocks = nn.ModuleList(
+ [
+ MaskedTransformerCrossAttnBlock(
+ hidden_dim, num_heads, cond_dim=hidden_dim
+ )
+ for _ in range(num_hybrid_stages)
+ ]
+ )
+ self.z_proj = nn.Linear(hidden_dim, z_dim * 2)
+ nn.init.zeros_(self.z_proj.weight)
+ nn.init.zeros_(self.z_proj.bias)
+
+ def forward(
+ self,
+ verts: torch.Tensor,
+ pc_tokens: Optional[torch.Tensor],
+ point_embedder: PointEmbed,
+ verts_mask: Optional[torch.Tensor] = None,
+ adj_matrix: Optional[torch.Tensor] = None,
+ ):
+ rope_phases = self.rope(verts.long())
+ coords = (verts + 0.5) / self.num_discrete * 2 - 1
+ vert_tokens = point_embedder(coords)
+ vert_tokens = manual_cast(vert_tokens, self.dtype)
+
+ adj_mask = None
+ if adj_matrix is not None:
+ b, n, _ = adj_matrix.shape
+ eye = torch.eye(n, device=adj_matrix.device, dtype=torch.bool).unsqueeze(0)
+ adj_mask = adj_matrix.bool() | eye
+
+ if self.pc_cross_attn:
+ if pc_tokens is None:
+ raise ValueError(
+ "pc_tokens required when encoder pc_cross_attn is enabled"
+ )
+ for stage, ca_block in zip(self.backbone.stages, self.pc_cross_blocks):
+ vert_tokens = stage(
+ vert_tokens,
+ x_mask=verts_mask,
+ adj_matrix=adj_mask,
+ rope_phases=rope_phases,
+ )
+ vert_tokens = ca_block(
+ vert_tokens,
+ pc_tokens,
+ x_mask=verts_mask,
+ c_mask=None,
+ )
+ else:
+ vert_tokens = self.backbone(
+ vert_tokens,
+ x_mask=verts_mask,
+ adj_matrix=adj_mask,
+ rope_phases=rope_phases,
+ )
+ z = self.z_proj(vert_tokens)
+ z = manual_cast(z, self.dtype)
+ return z
+
+
+class TopologyDecoderHybrid(nn.Module):
+ def __init__(
+ self,
+ z_dim: int = 32,
+ hidden_dim: int = 384,
+ num_heads: int = 6,
+ num_discrete: int = 256,
+ dtype: str = "float32",
+ num_hybrid_stages: int = 2,
+ num_flash_per_stage: int = 1,
+ use_gradient_checkpointing: bool = False,
+ ):
+ super().__init__()
+ self.dtype = str_to_dtype(dtype)
+ self.num_discrete = num_discrete
+ self.input_proj = nn.Linear(z_dim, hidden_dim)
+ self.backbone = HybridGraphFlashStack(
+ hidden_size=hidden_dim,
+ num_heads=num_heads,
+ num_stages=num_hybrid_stages,
+ num_flash_per_stage=num_flash_per_stage,
+ gradient_checkpointing=use_gradient_checkpointing,
+ )
+
+ def forward(self, z: torch.Tensor, verts_mask: Optional[torch.Tensor] = None):
+ h = self.input_proj(z)
+ h = manual_cast(h, self.dtype)
+ h = self.backbone(
+ h,
+ x_mask=verts_mask,
+ adj_matrix=None,
+ rope_phases=None,
+ )
+ return h
+
+
+class TopologyConnectionPredictor(nn.Module):
+ def __init__(self, hidden_dim: int = 384):
+ super().__init__()
+ self.mlp = nn.Sequential(
+ nn.Linear(hidden_dim * 2, 256),
+ nn.GELU(),
+ nn.Linear(256, 1),
+ )
+
+ def forward(self, vert_feat_u: torch.Tensor, vert_feat_v: torch.Tensor):
+ pair_feat_0 = torch.cat([vert_feat_u, vert_feat_v], dim=-1)
+ pair_feat_1 = torch.cat([vert_feat_v, vert_feat_u], dim=-1)
+ h = (self.mlp(pair_feat_0) + self.mlp(pair_feat_1)) / 2.0
+ return h.squeeze(-1)
+
+
+class TopologyVAE(nn.Module):
+ def __init__(
+ self,
+ z_dim: int = 32,
+ hidden_dim: int = 384,
+ pc_dim: int = 15,
+ inner_pc_dim: int = 256,
+ num_heads: int = 6,
+ num_discrete: int = 256,
+ dtype: str = "float32",
+ num_hybrid_stages: int = 2,
+ num_flash_per_stage: int = 1,
+ num_connection_blocks: Optional[int] = None,
+ use_gradient_checkpointing: bool = False,
+ encoder_pc_cross_attn: bool = False,
+ ):
+ super().__init__()
+ self.dtype = str_to_dtype(dtype)
+ self.num_discrete = num_discrete
+ self.encoder_pc_cross_attn = bool(encoder_pc_cross_attn)
+
+ self.point_embed = PointEmbed(hidden_dim=hidden_dim, dim=hidden_dim)
+ self.point_net = Pointnet(
+ in_channels=pc_dim,
+ out_channels=inner_pc_dim,
+ hidden_dim=256,
+ n_blocks=5,
+ )
+ self.point_fusion = nn.Linear(hidden_dim + inner_pc_dim, hidden_dim)
+
+ self.encoder = TopologyEncoderHybrid(
+ z_dim=z_dim,
+ hidden_dim=hidden_dim,
+ num_heads=num_heads,
+ num_discrete=num_discrete,
+ dtype=dtype,
+ num_hybrid_stages=num_hybrid_stages,
+ num_flash_per_stage=num_flash_per_stage,
+ use_gradient_checkpointing=use_gradient_checkpointing,
+ pc_cross_attn=self.encoder_pc_cross_attn,
+ )
+ self.decoder = TopologyDecoderHybrid(
+ z_dim=z_dim,
+ hidden_dim=hidden_dim,
+ num_heads=num_heads,
+ num_discrete=num_discrete,
+ dtype=dtype,
+ num_hybrid_stages=num_hybrid_stages,
+ num_flash_per_stage=num_flash_per_stage,
+ use_gradient_checkpointing=use_gradient_checkpointing,
+ )
+ self.connection_predictor = TopologyConnectionPredictor(hidden_dim=hidden_dim)
+
+ head_dim = hidden_dim // num_heads
+ self.connection_rope = RotaryPositionPhasesEmbedder(head_dim=head_dim, dim=3)
+
+ n_conn = (
+ num_connection_blocks
+ if num_connection_blocks is not None
+ else (num_hybrid_stages * (1 + num_flash_per_stage))
+ )
+ self.connection_transformer_blocks = nn.ModuleList(
+ [
+ FlashVarlenTransformerBlock(
+ hidden_dim,
+ num_heads,
+ gradient_checkpointing=use_gradient_checkpointing,
+ )
+ for _ in range(n_conn)
+ ]
+ )
+
+ def encode(
+ self,
+ verts: torch.Tensor,
+ pc_tokens: torch.Tensor,
+ verts_mask: Optional[torch.Tensor] = None,
+ adj_matrix: Optional[torch.Tensor] = None,
+ ):
+ moments = self.encoder(
+ verts=verts,
+ pc_tokens=pc_tokens,
+ point_embedder=self.point_embed,
+ verts_mask=verts_mask,
+ adj_matrix=adj_matrix,
+ )
+ mean, logvar = moments.chunk(2, dim=-1)
+ return mean, logvar
+
+ def decode(
+ self,
+ z: torch.Tensor,
+ verts: torch.Tensor,
+ verts_mask: Optional[torch.Tensor] = None,
+ chunk_size: int = 20000,
+ threshold: float = 0.0,
+ ) -> List[np.ndarray]:
+ # return: list of [N, 2] numpy arrays of predicted edges for each batch item
+ verts_feat = self.decoder(z=z, verts_mask=verts_mask)
+
+ rope = self.connection_rope(verts.long())
+ for block in self.connection_transformer_blocks:
+ verts_feat = block(verts_feat, verts_mask, rope_phases=rope)
+
+ all_pred_edges_list = []
+ for i in range(verts_mask.shape[0]):
+ valid = verts_mask[i]
+ valid_verts = verts[i][valid]
+ valid_feats = verts_feat[i][valid]
+ num_valid = int(valid_verts.shape[0])
+
+ u_idx, v_idx = torch.triu_indices(
+ num_valid, num_valid, offset=1, device=z.device
+ )
+ pred_edges_list = []
+ for i in range(0, u_idx.numel(), chunk_size):
+ cu = u_idx[i : i + chunk_size]
+ cv = v_idx[i : i + chunk_size]
+ logits = self.connection_predictor(
+ valid_feats[cu].unsqueeze(0),
+ valid_feats[cv].unsqueeze(0),
+ ).squeeze(0)
+ take = logits > threshold
+ if bool(take.any()):
+ pred_edges_list.append(torch.stack([cu[take], cv[take]], dim=-1))
+
+ pred_edges = (
+ torch.cat(pred_edges_list, dim=0).cpu().numpy()
+ if pred_edges_list
+ else np.empty((0, 2), dtype=np.int64)
+ )
+ all_pred_edges_list.append(pred_edges)
+
+ return all_pred_edges_list
diff --git a/models/topo_flow.py b/models/topo_flow.py
new file mode 100644
index 0000000000000000000000000000000000000000..d620981f07265de354bab5117f68553303e5ec35
--- /dev/null
+++ b/models/topo_flow.py
@@ -0,0 +1,400 @@
+from __future__ import annotations
+
+import torch
+import torch.nn as nn
+import torch.nn.functional as F
+from torch.utils.checkpoint import checkpoint
+
+from modules.transformer.blocks import RotaryPositionPhasesEmbedder, TimestepEmbedder
+from modules.attention import (
+ can_flash_varlen,
+ flash_varlen_self_attention,
+ flash_varlen_cross_attention,
+ sdpa_padding_mask,
+)
+from modules.norm import RMSNorm
+from modules.utils import modulate
+
+
+class TopologySiTBlockFlashVarlen(nn.Module):
+ def __init__(
+ self,
+ hidden_size: int,
+ num_heads: int,
+ mlp_ratio: float = 4.0,
+ dropout: float = 0.0,
+ gradient_checkpointing: bool = False,
+ qk_norm_eps: float = 1e-5,
+ qk_norm_variance_in_fp32: bool = True,
+ with_cross_attn: bool = False,
+ ):
+ super().__init__()
+ if hidden_size % num_heads != 0:
+ raise ValueError(
+ f"hidden_size {hidden_size} not divisible by num_heads {num_heads}"
+ )
+ self.hidden_size = hidden_size
+ self.num_heads = num_heads
+ self.head_dim = hidden_size // num_heads
+ self.with_cross_attn = bool(with_cross_attn)
+ self.norm1 = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6)
+ self.norm2 = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6)
+ self.qkv = nn.Linear(hidden_size, hidden_size * 3, bias=True)
+ self.proj_out = nn.Linear(hidden_size, hidden_size, bias=True)
+ mlp_hidden = int(hidden_size * mlp_ratio)
+ self.mlp = nn.Sequential(
+ nn.Linear(hidden_size, mlp_hidden, bias=True),
+ nn.GELU(approximate="tanh"),
+ nn.Linear(mlp_hidden, hidden_size, bias=True),
+ )
+ self._n_adaln_chunks = 7 if self.with_cross_attn else 6
+ self.adaLN_modulation = nn.Sequential(
+ nn.SiLU(),
+ nn.Linear(hidden_size, self._n_adaln_chunks * hidden_size, bias=True),
+ )
+ self.dropout = dropout
+ self.gradient_checkpointing = bool(gradient_checkpointing)
+
+ self.norm_q, self.norm_k = (
+ RMSNorm(
+ self.head_dim,
+ eps=qk_norm_eps,
+ elementwise_affine=True,
+ variance_in_fp32=qk_norm_variance_in_fp32,
+ ),
+ RMSNorm(
+ self.head_dim,
+ eps=qk_norm_eps,
+ elementwise_affine=True,
+ variance_in_fp32=qk_norm_variance_in_fp32,
+ ),
+ )
+
+ if self.with_cross_attn:
+ self.norm_ca = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6)
+ self.q_ca = nn.Linear(hidden_size, hidden_size, bias=True)
+ self.kv_ca = nn.Linear(hidden_size, hidden_size * 2, bias=True)
+ self.proj_ca_out = nn.Linear(hidden_size, hidden_size, bias=True)
+ self.norm_q_ca, self.norm_k_ca = (
+ RMSNorm(
+ self.head_dim,
+ eps=qk_norm_eps,
+ elementwise_affine=True,
+ variance_in_fp32=qk_norm_variance_in_fp32,
+ ),
+ RMSNorm(
+ self.head_dim,
+ eps=qk_norm_eps,
+ elementwise_affine=True,
+ variance_in_fp32=qk_norm_variance_in_fp32,
+ ),
+ )
+
+ def _forward_once(
+ self,
+ x: torch.Tensor,
+ c: torch.Tensor,
+ key_padding_mask: torch.Tensor | None,
+ rope_phases: torch.Tensor,
+ cond_emb: torch.Tensor | None,
+ cond_mask: torch.Tensor | None,
+ ) -> torch.Tensor:
+ chunks = self.adaLN_modulation(c).chunk(self._n_adaln_chunks, dim=1)
+ if self.with_cross_attn:
+ shift_msa, scale_msa, gate_msa, gate_mca, shift_mlp, scale_mlp, gate_mlp = (
+ chunks
+ )
+ else:
+ shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = chunks
+
+ h = modulate(self.norm1(x), shift_msa, scale_msa)
+ b, n, d = h.shape
+ qkv = (
+ self.qkv(h)
+ .view(b, n, 3, self.num_heads, self.head_dim)
+ .permute(2, 0, 3, 1, 4)
+ )
+ q, k, v = qkv[0], qkv[1], qkv[2]
+ if self.norm_q is not None:
+ q = self.norm_q(q)
+ if self.norm_k is not None:
+ k = self.norm_k(k)
+ q = RotaryPositionPhasesEmbedder.apply_rotary_embedding(q, rope_phases)
+ k = RotaryPositionPhasesEmbedder.apply_rotary_embedding(k, rope_phases)
+
+ x_mask = None if key_padding_mask is None else ~key_padding_mask.bool()
+ if can_flash_varlen(q, x_mask):
+ attn_out = flash_varlen_self_attention(q, k, v, x_mask)
+ elif x_mask is not None:
+ attn_out = F.scaled_dot_product_attention(
+ q,
+ k,
+ v,
+ attn_mask=sdpa_padding_mask(x_mask),
+ dropout_p=self.dropout if self.training else 0.0,
+ )
+ else:
+ attn_out = F.scaled_dot_product_attention(
+ q,
+ k,
+ v,
+ attn_mask=None,
+ dropout_p=self.dropout if self.training else 0.0,
+ )
+
+ attn_out = attn_out.transpose(1, 2).reshape(b, n, d)
+ x = x + gate_msa.unsqueeze(1) * self.proj_out(attn_out)
+
+ if self.with_cross_attn and cond_emb is not None:
+ h_ca = self.norm_ca(x)
+ nk = cond_emb.shape[1]
+ q_ca = (
+ self.q_ca(h_ca)
+ .view(b, n, self.num_heads, self.head_dim)
+ .transpose(1, 2)
+ )
+ kv_ca = (
+ self.kv_ca(cond_emb)
+ .view(b, nk, 2, self.num_heads, self.head_dim)
+ .permute(2, 0, 3, 1, 4)
+ )
+ k_ca, v_ca = kv_ca[0], kv_ca[1]
+ if self.norm_q_ca is not None:
+ q_ca = self.norm_q_ca(q_ca)
+ if self.norm_k_ca is not None:
+ k_ca = self.norm_k_ca(k_ca)
+
+ q_mask_bool = (
+ torch.ones(b, n, dtype=torch.bool, device=q_ca.device)
+ if key_padding_mask is None
+ else ~key_padding_mask.bool()
+ )
+ k_mask_bool = (
+ torch.ones(b, nk, dtype=torch.bool, device=q_ca.device)
+ if cond_mask is None
+ else cond_mask.bool()
+ )
+
+ if can_flash_varlen(q_ca, q_mask_bool):
+ ca_out = flash_varlen_cross_attention(
+ q_ca, k_ca, v_ca, q_mask_bool, k_mask_bool
+ )
+ else:
+ k_attn_mask = k_mask_bool.view(b, 1, 1, nk)
+ ca_out = F.scaled_dot_product_attention(
+ q_ca, k_ca, v_ca, attn_mask=k_attn_mask, dropout_p=0.0
+ )
+ ca_out = ca_out.transpose(1, 2).reshape(b, n, d)
+ x = x + gate_mca.unsqueeze(1) * self.proj_ca_out(ca_out)
+
+ h2 = modulate(self.norm2(x), shift_mlp, scale_mlp)
+ x = x + gate_mlp.unsqueeze(1) * self.mlp(h2)
+ if key_padding_mask is not None:
+ valid = ~key_padding_mask.bool()
+ x = torch.where(valid.unsqueeze(-1), x, torch.zeros_like(x))
+ return x
+
+ def forward(
+ self,
+ x: torch.Tensor,
+ c: torch.Tensor,
+ key_padding_mask: torch.Tensor | None,
+ rope_phases: torch.Tensor,
+ cond_emb: torch.Tensor | None = None,
+ cond_mask: torch.Tensor | None = None,
+ ) -> torch.Tensor:
+ if self.training and self.gradient_checkpointing:
+ return checkpoint(
+ self._forward_once,
+ x,
+ c,
+ key_padding_mask,
+ rope_phases,
+ cond_emb,
+ cond_mask,
+ use_reentrant=False,
+ )
+ return self._forward_once(
+ x, c, key_padding_mask, rope_phases, cond_emb, cond_mask
+ )
+
+
+class TopologyFinalLayer(nn.Module):
+ def __init__(self, hidden_size: int, out_channels: int):
+ super().__init__()
+ self.norm_final = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6)
+ self.linear = nn.Linear(hidden_size, out_channels, bias=True)
+ self.adaLN_modulation = nn.Sequential(
+ nn.SiLU(),
+ nn.Linear(hidden_size, 2 * hidden_size, bias=True),
+ )
+
+ def forward(self, x: torch.Tensor, c: torch.Tensor) -> torch.Tensor:
+ shift, scale = self.adaLN_modulation(c).chunk(2, dim=1)
+ x = modulate(self.norm_final(x), shift, scale)
+ return self.linear(x)
+
+
+def _get_1d_sincos_embed(n: int, dim: int) -> torch.Tensor:
+ assert dim % 2 == 0
+ pos = torch.arange(n, dtype=torch.float32)
+ omega = torch.arange(dim // 2, dtype=torch.float32) / (dim // 2)
+ omega = 1.0 / (10000**omega)
+ out = pos[:, None] * omega[None, :]
+ return torch.cat([torch.sin(out), torch.cos(out)], dim=-1)
+
+
+class TopologySiTFlow(nn.Module):
+ def __init__(
+ self,
+ z_dim: int,
+ hidden_size: int = 768,
+ depth: int = 12,
+ num_heads: int = 12,
+ mlp_ratio: float = 4.0,
+ max_vertices: int = 8192,
+ num_discrete: int = 1024,
+ dropout: float = 0.0,
+ gradient_checkpointing: bool = False,
+ cond_in_dim: int = 0,
+ cond_dropout_prob: float = 0.0,
+ qk_norm_eps: float = 1e-5,
+ qk_norm_variance_in_fp32: bool = True,
+ ):
+ super().__init__()
+ self.z_dim = z_dim
+ self.hidden_size = hidden_size
+ self.max_vertices = max_vertices
+ self.num_discrete = int(num_discrete)
+ self.gradient_checkpointing = bool(gradient_checkpointing)
+ self.cond_in_dim = int(cond_in_dim)
+ self.cond_dropout_prob = float(cond_dropout_prob)
+ self.qk_norm_eps = float(qk_norm_eps)
+ self.qk_norm_variance_in_fp32 = bool(qk_norm_variance_in_fp32)
+
+ self.input_proj = nn.Linear(z_dim, hidden_size, bias=True)
+ self.coord_embed = nn.Sequential(
+ nn.Linear(3, hidden_size, bias=True),
+ nn.SiLU(),
+ nn.Linear(hidden_size, hidden_size, bias=True),
+ )
+ self.t_embedder = TimestepEmbedder(hidden_size)
+ self.rope = RotaryPositionPhasesEmbedder(
+ head_dim=hidden_size // num_heads, dim=3
+ )
+
+ if self.cond_in_dim > 0:
+ self.cond_proj = nn.Sequential(
+ nn.Linear(self.cond_in_dim, hidden_size, bias=True),
+ nn.SiLU(),
+ nn.Linear(hidden_size, hidden_size, bias=True),
+ )
+ self.null_token = nn.Parameter(torch.zeros(hidden_size))
+ else:
+ self.cond_proj = None
+ self.null_token = None
+
+ pe = _get_1d_sincos_embed(max_vertices, hidden_size)
+ self.register_buffer("pos_embed", pe.unsqueeze(0), persistent=False)
+
+ with_cross_attn = self.cond_in_dim > 0
+ self.blocks = nn.ModuleList(
+ [
+ TopologySiTBlockFlashVarlen(
+ hidden_size=hidden_size,
+ num_heads=num_heads,
+ mlp_ratio=mlp_ratio,
+ dropout=dropout,
+ gradient_checkpointing=self.gradient_checkpointing,
+ with_cross_attn=with_cross_attn,
+ qk_norm_eps=self.qk_norm_eps,
+ qk_norm_variance_in_fp32=self.qk_norm_variance_in_fp32,
+ )
+ for _ in range(depth)
+ ]
+ )
+ self.final_layer = TopologyFinalLayer(hidden_size, z_dim)
+
+ def _prepare_cond(
+ self,
+ b: int,
+ cond: torch.Tensor | None,
+ cond_mask: torch.Tensor | None,
+ cond_drop_override: torch.Tensor | None,
+ device: torch.device,
+ ) -> tuple[torch.Tensor | None, torch.Tensor | None]:
+ if self.cond_proj is None:
+ return None, None
+
+ if cond is None:
+ null_emb = self.null_token.view(1, 1, -1).expand(b, 1, -1).contiguous()
+ mask_out = torch.ones(b, 1, dtype=torch.bool, device=device)
+ return null_emb, mask_out
+
+ if cond.dim() != 3 or cond.shape[0] != b or cond.shape[-1] != self.cond_in_dim:
+ raise ValueError(
+ f"cond shape {tuple(cond.shape)} expected ({b}, K, {self.cond_in_dim})"
+ )
+ k = cond.shape[1]
+ cond_emb = self.cond_proj(cond)
+ null = self.null_token.view(1, 1, -1).to(dtype=cond_emb.dtype)
+
+ if cond_mask is None:
+ mask_out = torch.ones(b, k, dtype=torch.bool, device=device)
+ else:
+ mask_out = cond_mask.to(device=device, dtype=torch.bool)
+ if mask_out.shape != (b, k):
+ raise ValueError(
+ f"cond_mask shape {tuple(mask_out.shape)} expected ({b}, {k})"
+ )
+
+ drop: torch.Tensor | None = None
+ if cond_drop_override is not None:
+ drop = cond_drop_override.to(device=device, dtype=torch.bool).reshape(b)
+ elif self.training and self.cond_dropout_prob > 0.0:
+ drop = torch.rand(b, device=device) < self.cond_dropout_prob
+
+ if drop is not None:
+ null_emb = null.expand(b, k, -1)
+ cond_emb = torch.where(drop.view(b, 1, 1), null_emb, cond_emb)
+ mask_out = torch.where(drop.view(b, 1), torch.ones_like(mask_out), mask_out)
+ return cond_emb, mask_out
+
+ def forward(
+ self,
+ x: torch.Tensor,
+ t: torch.Tensor,
+ verts: torch.Tensor,
+ mask: torch.Tensor,
+ cond: torch.Tensor | None = None,
+ cond_mask: torch.Tensor | None = None,
+ cond_drop_override: torch.Tensor | None = None,
+ ) -> torch.Tensor:
+ b, n, _ = x.shape
+ if n > self.max_vertices:
+ raise ValueError(f"Sequence length {n} > max_vertices {self.max_vertices}")
+ if self.cond_in_dim == 0:
+ if cond is not None:
+ raise ValueError("TopologySiT(cond_in_dim=0): pass cond=None")
+ if cond_mask is not None:
+ raise ValueError("TopologySiT(cond_in_dim=0): cond_mask is unused")
+ if cond_drop_override is not None:
+ raise ValueError(
+ "TopologySiT(cond_in_dim=0): cond_drop_override is unused"
+ )
+
+ coords = ((verts.float() + 0.5) / self.num_discrete) * 2.0 - 1.0
+ h = self.input_proj(x) + self.pos_embed[:, :n, :] + self.coord_embed(coords)
+ c = self.t_embedder(t)
+
+ cond_emb, cond_mask_eff = self._prepare_cond(
+ b, cond, cond_mask, cond_drop_override, x.device
+ )
+
+ key_padding_mask = ~mask
+ rope_phases = self.rope(verts.long())
+ for block in self.blocks:
+ h = block(h, c, key_padding_mask, rope_phases, cond_emb, cond_mask_eff)
+ out = self.final_layer(h, c)
+ out = torch.where(mask.unsqueeze(-1), out, torch.zeros_like(out))
+ return out
diff --git a/models/vdf_encoder.py b/models/vdf_encoder.py
new file mode 100644
index 0000000000000000000000000000000000000000..40df530073691478cfddca5d54798e508ed9264b
--- /dev/null
+++ b/models/vdf_encoder.py
@@ -0,0 +1,55 @@
+import torch
+import torch.nn as nn
+from torch.utils.checkpoint import checkpoint
+
+from modules.pointnet import LocalPoolPointnet
+
+
+class VDFEncoder(nn.Module):
+ def __init__(
+ self,
+ in_channels,
+ hidden_dim,
+ out_channels,
+ scatter_type,
+ n_blocks,
+ resolution=64,
+ use_checkpoint=False,
+ ):
+ super().__init__()
+ self.pointnet = LocalPoolPointnet(
+ in_channels=in_channels,
+ out_channels=out_channels,
+ hidden_dim=hidden_dim,
+ n_blocks=n_blocks,
+ scatter_type=scatter_type,
+ )
+
+ self.resolution = resolution
+ self.use_checkpoint = use_checkpoint
+
+ def forward(
+ self,
+ p,
+ sparse_coords,
+ res=None,
+ bbox_size=(-0.5, 0.5),
+ ):
+ """
+ Input:
+ p: [N, in_channels]
+ sparse_coords: [M, 4], (b, z, y, x)
+ Output:
+ geo_feats: [N, out_channels]
+ """
+ if res is None:
+ res = self.resolution
+
+ if self.use_checkpoint and self.training:
+ geo_feats = checkpoint(
+ self.pointnet, p, sparse_coords, res, bbox_size, use_reentrant=False
+ )
+ else:
+ geo_feats = self.pointnet(p, sparse_coords, res=res, bbox_size=bbox_size)
+
+ return geo_feats
diff --git a/models/vertex_autoencoder.py b/models/vertex_autoencoder.py
new file mode 100644
index 0000000000000000000000000000000000000000..1a05d453beed9175711219156e36847f1e1d1123
--- /dev/null
+++ b/models/vertex_autoencoder.py
@@ -0,0 +1,595 @@
+import torch
+import torch.nn as nn
+from typing import *
+import torch.nn.functional as F
+
+from modules import sparse as sp
+from modules.sparse import SparseTensor
+from modules.sparse.linear import SparseLinear
+from modules.sparse.nonlinearity import SparseGELU
+from modules.utils import (
+ zero_module,
+ convert_module_to_f16,
+ convert_module_to_f32,
+ flatten_coords,
+ per_batch_counts,
+)
+from modules.sparse.transformer import SparseTransformerBase, SparseTransformerCrossBase
+from modules.sparse.blocks import SparseResBlock3d
+from modules.utils import DiagonalGaussianDistribution
+
+
+class SparseOccHead(nn.Module):
+ def __init__(self, channels: int, out_channels: int, mlp_ratio: float = 4.0):
+ super().__init__()
+ self.mlp = nn.Sequential(
+ SparseLinear(channels, int(channels * mlp_ratio)),
+ SparseGELU(approximate="tanh"),
+ SparseLinear(int(channels * mlp_ratio), out_channels),
+ )
+
+ def forward(self, x: sp.SparseTensor) -> sp.SparseTensor:
+ return self.mlp(x)
+
+
+class SparseEncoderBlock(nn.Module):
+ def __init__(
+ self,
+ resolution: int,
+ in_channels: int,
+ model_channels: int,
+ num_blocks: int,
+ num_downsample: int = 4,
+ num_heads: Optional[int] = None,
+ num_head_channels: Optional[int] = 64,
+ mlp_ratio: float = 4,
+ attn_mode: Literal[
+ "full", "shift_window", "shift_sequence", "shift_order", "swin"
+ ] = "swin",
+ window_size: int = 8,
+ pe_mode: Literal["ape", "rope"] = "ape",
+ use_fp16: bool = False,
+ use_checkpoint: bool = False,
+ qk_rms_norm: bool = False,
+ ):
+ super().__init__()
+ self.resolution = resolution
+
+ self.self_attn = SparseTransformerBase(
+ in_channels=model_channels,
+ model_channels=model_channels,
+ num_blocks=num_blocks,
+ num_heads=num_heads,
+ num_head_channels=num_head_channels,
+ attn_mode=attn_mode,
+ window_size=window_size,
+ pe_mode=pe_mode,
+ mlp_ratio=mlp_ratio,
+ use_fp16=use_fp16,
+ use_checkpoint=use_checkpoint,
+ qk_rms_norm=qk_rms_norm,
+ )
+
+ self.input_layer1 = sp.SparseLinear(
+ in_channels, model_channels >> num_downsample
+ )
+
+ self.downsample = nn.ModuleList(
+ [
+ SparseResBlock3d(
+ channels=model_channels >> (i + 1),
+ out_channels=model_channels >> i,
+ downsample=True,
+ upsample=False,
+ use_checkpoint=use_checkpoint,
+ )
+ for i in range(num_downsample - 1, -1, -1)
+ ]
+ )
+
+ def forward(
+ self,
+ x: SparseTensor,
+ ):
+ """
+ Input:
+ x: SparseTensor in N resolution, with feats of in_channels
+ Output:
+ h: SparseTensor in N>>num_downsample resolution, with feats of model_channels
+ """
+ x = self.input_layer1(x)
+ for block in self.downsample:
+ x = block(x)
+ h = self.self_attn(x)
+ return h
+
+
+class SparseDecoderUpsampleBlock(nn.Module):
+ def __init__(
+ self,
+ channels: int,
+ resolution: int,
+ out_channels: int,
+ model_channels: int = 512,
+ num_blocks: int = 4,
+ num_heads: int = 8,
+ mlp_ratio: float = 4.0,
+ num_groups: int = 32,
+ ):
+ super().__init__()
+ self.channels = channels
+ self.resolution = resolution
+ self.out_resolution = resolution * 2
+ self.model_channels = model_channels
+ self.out_channels = out_channels
+
+ self.act_layers = nn.Sequential(
+ sp.SparseGroupNorm32(num_groups, channels), sp.SparseSiLU()
+ )
+
+ self.sub = sp.SparseSubdivide()
+
+ self.out_layers = nn.Sequential(
+ sp.SparseConv3d(
+ channels, self.out_channels, 3, indice_key=f"res_{self.out_resolution}"
+ ),
+ sp.SparseGroupNorm32(num_groups, self.out_channels),
+ sp.SparseSiLU(),
+ zero_module(
+ sp.SparseConv3d(
+ self.out_channels,
+ self.out_channels,
+ 3,
+ indice_key=f"res_{self.out_resolution}",
+ )
+ ),
+ )
+
+ if self.out_channels == channels:
+ self.skip_connection = nn.Identity()
+ else:
+ self.skip_connection = sp.SparseConv3d(
+ channels, self.out_channels, 1, indice_key=f"res_{self.out_resolution}"
+ )
+
+ self.pruning_head = SparseOccHead(self.out_channels, out_channels=1)
+
+ self.ca = SparseTransformerCrossBase(
+ in_channels=self.out_channels,
+ model_channels=self.model_channels,
+ context_channels=self.model_channels,
+ num_blocks=num_blocks,
+ num_heads=num_heads,
+ mlp_ratio=mlp_ratio,
+ attn_mode="full",
+ pe_mode="ape",
+ use_checkpoint=True,
+ qk_rms_norm=False,
+ )
+
+ self.proj_ctx = sp.SparseLinear(self.out_channels, self.model_channels)
+ self.proj_out = sp.SparseLinear(self.model_channels, self.out_channels)
+
+ def forward(
+ self,
+ x: sp.SparseTensor,
+ training=False,
+ threshold=0.5,
+ ) -> sp.SparseTensor:
+ h = self.act_layers(x)
+ h = self.sub(h)
+ x_sub = self.sub(x)
+ h = self.out_layers(h)
+ h = h + self.skip_connection(x_sub)
+ h = self.proj_out(self.ca(x=h, context=self.proj_ctx(h)))
+
+ occ_prob_q = self.pruning_head(h)
+
+ if training:
+ return h, occ_prob_q, [0]
+
+ scores_q = torch.sigmoid(occ_prob_q.feats).squeeze(-1)
+ N_full = h.feats.shape[0]
+ if N_full % 8 != 0:
+ raise ValueError(f"Number of nodes({N_full}) is not divisible by 8.")
+
+ # ensure at least one point is kept in each group of 8
+ n_parents = N_full // 8
+
+ scores_q_grouped = scores_q.view(n_parents, 8)
+
+ mask_grouped = scores_q_grouped >= threshold
+
+ none_survived = mask_grouped.sum(dim=1) == 0
+
+ # per-batch rescue counts; all 8 children of a parent share one batch index
+ if n_parents > 0:
+ parent_batch = h.coords[:, 0].view(n_parents, 8)[:, 0]
+ num_rescue = per_batch_counts(
+ parent_batch[none_survived], int(parent_batch.max().item()) + 1
+ )
+ else:
+ num_rescue = [0]
+ if none_survived.any():
+ failed_scores = scores_q_grouped[none_survived]
+ _, topk_indices = torch.topk(failed_scores, k=1, dim=1)
+
+ failed_row_idxs = torch.nonzero(none_survived, as_tuple=True)[0]
+ rows_expanded = failed_row_idxs.unsqueeze(1).expand(-1, 1)
+
+ mask_grouped[rows_expanded, topk_indices] = True
+
+ sub_mask = mask_grouped.view(-1)
+
+ h = sp.SparseTensor(feats=h.feats[sub_mask], coords=h.coords[sub_mask])
+ occ_prob_final = sp.SparseTensor(
+ feats=occ_prob_q.feats[sub_mask], coords=occ_prob_q.coords[sub_mask]
+ )
+
+ return h, occ_prob_final, num_rescue
+
+
+class SparseDecoderBlock(nn.Module):
+ def __init__(
+ self,
+ resolution: int,
+ in_channels: int,
+ out_channels: int,
+ model_channels: int = 512,
+ num_blocks: int = 4,
+ num_heads: int = 8,
+ mlp_ratio: float = 4.0,
+ use_fp16: bool = False,
+ ):
+ super().__init__()
+ self.resolution = resolution
+
+ self.upsample = SparseDecoderUpsampleBlock(
+ channels=in_channels,
+ resolution=resolution,
+ out_channels=out_channels,
+ num_blocks=num_blocks,
+ num_heads=num_heads,
+ mlp_ratio=mlp_ratio,
+ model_channels=model_channels,
+ num_groups=32,
+ )
+
+ if use_fp16:
+ self.convert_to_fp16()
+
+ def forward(
+ self,
+ x: sp.SparseTensor,
+ training: bool = False,
+ threshold: float = 0.5,
+ ):
+ h = x
+ h = h.type(x.dtype)
+ h = h.replace(F.layer_norm(h.feats, h.feats.shape[-1:]))
+ h, occ_prob, num_rescue = self.upsample(
+ h,
+ training=training,
+ threshold=threshold,
+ )
+ return h, occ_prob, num_rescue
+
+ def convert_to_fp16(self):
+ """Convert all components to float16"""
+ convert_module_to_f16(self.upsample)
+
+ def convert_to_fp32(self):
+ """Convert all components to float32"""
+ convert_module_to_f32(self.upsample)
+
+
+class VertexVAE(nn.Module):
+ def __init__(
+ self,
+ # Core architecture parameters
+ encoder_cfg: Dict = {},
+ expander_cfg: Dict = {},
+ decoder_cfg: List[Dict] = [],
+ # Shared transformer parameters
+ resolution: int = 1024,
+ num_head_channels: Optional[int] = 64,
+ mlp_ratio: float = 4.0,
+ attn_mode: str = "swin",
+ window_size: int = 8,
+ pe_mode: str = "ape",
+ use_fp16: bool = False,
+ use_checkpoint: bool = True,
+ qk_rms_norm: bool = False,
+ latent_dim: int = 8,
+ ):
+ super().__init__()
+ self.latent_dim = latent_dim
+ self.decoder_cfg = decoder_cfg
+
+ self.encoder = SparseEncoderBlock(
+ resolution=resolution,
+ in_channels=encoder_cfg["in_channels"],
+ model_channels=encoder_cfg["model_channels"],
+ num_blocks=encoder_cfg["num_blocks"],
+ num_heads=encoder_cfg["num_heads"],
+ num_downsample=len(decoder_cfg),
+ num_head_channels=num_head_channels,
+ attn_mode=attn_mode,
+ window_size=window_size,
+ pe_mode=pe_mode,
+ mlp_ratio=mlp_ratio,
+ use_fp16=use_fp16,
+ use_checkpoint=use_checkpoint,
+ qk_rms_norm=qk_rms_norm,
+ )
+
+ self.latent_expander = SparseTransformerBase(
+ in_channels=latent_dim,
+ model_channels=expander_cfg["model_channels"],
+ num_blocks=expander_cfg["num_blocks"],
+ num_heads=expander_cfg["num_heads"],
+ num_head_channels=num_head_channels,
+ attn_mode=attn_mode,
+ window_size=window_size,
+ pe_mode=pe_mode,
+ mlp_ratio=mlp_ratio,
+ use_fp16=use_fp16,
+ use_checkpoint=use_checkpoint,
+ qk_rms_norm=qk_rms_norm,
+ )
+
+ self.vtx_proj = sp.SparseLinear(
+ expander_cfg["model_channels"], decoder_cfg[0]["in_channels"]
+ )
+
+ self.vtx_pruning_head = SparseOccHead(
+ expander_cfg["model_channels"], out_channels=1
+ )
+
+ self.out_layer = sp.SparseLinear(expander_cfg["model_channels"], latent_dim * 2)
+
+ self.decoder_vtx = nn.ModuleList()
+ self.decoder_vtx_ca = nn.ModuleList()
+ self.latent_proj = nn.ModuleList()
+ for config in decoder_cfg:
+ self.decoder_vtx.append(
+ # using default parameters to init the upsample block
+ SparseDecoderBlock(
+ resolution=config["resolution"],
+ in_channels=config["in_channels"],
+ out_channels=config["out_channels"],
+ num_blocks=config["num_blocks"],
+ num_heads=config["num_heads"],
+ use_fp16=use_fp16,
+ )
+ )
+ self.latent_proj.append(
+ sp.SparseLinear(latent_dim, config["context_channels"])
+ )
+ self.decoder_vtx_ca.append(
+ SparseTransformerCrossBase(
+ in_channels=config["out_channels"],
+ model_channels=config["model_channels"],
+ context_channels=config["context_channels"],
+ num_blocks=config["num_blocks"],
+ num_heads=config["num_heads"],
+ num_head_channels=num_head_channels,
+ mlp_ratio=mlp_ratio,
+ attn_mode="full",
+ window_size=window_size,
+ pe_mode=pe_mode,
+ use_fp16=use_fp16,
+ use_checkpoint=use_checkpoint,
+ qk_rms_norm=qk_rms_norm,
+ )
+ )
+
+ if use_fp16:
+ self.convert_to_fp16()
+
+ def encode(
+ self,
+ x: sp.SparseTensor,
+ sample_posterior=True,
+ ) -> Tuple[torch.Tensor, torch.Tensor]:
+ h = self.encoder(x)
+ h = h.type(x.dtype)
+ h = h.replace(F.layer_norm(h.feats, h.feats.shape[-1:]))
+ h = self.out_layer(h)
+
+ posterior = DiagonalGaussianDistribution(h.feats, feat_dim=-1)
+ if sample_posterior:
+ z = posterior.sample()
+ else:
+ z = posterior.mode()
+ z = h.replace(z)
+ return z, posterior
+
+ def decode(
+ self,
+ latent_: sp.SparseTensor,
+ gt_vertex_voxels_list: List[sp.SparseTensor],
+ training=True,
+ inference_threshold=0.5,
+ verbose=False,
+ ) -> List[Dict]:
+ """
+ Args:
+ latent: Initial SparseTensor from encoder at 64-resolution.
+ gt_vertex_voxels_list: Ground-truth vertex SparseTensors at [64, 128, 256, 512, 1024]
+ training: Whether to apply pruning during training
+
+ Returns:
+ List[Dict] with separate vertex and edge predictions at each level
+ """
+ latent = self.latent_expander(latent_)
+
+ results = []
+
+ # step0: shell voxels to vertex voxels
+ vtx_probs = self.vtx_pruning_head(latent) # (N, 1)
+ if not training:
+ # Inference path: use predicted vertex mask to split vertex
+
+ scores = torch.sigmoid(vtx_probs.feats).squeeze(-1) # (N,)
+
+ vertex_mask = scores >= inference_threshold # (N,)
+ batch_indices = latent.coords[:, 0]
+ for b in batch_indices.unique():
+ batch_sel = batch_indices == b
+ if vertex_mask[batch_sel].any():
+ continue
+ batch_scores = scores[batch_sel]
+ k = min(2, batch_scores.numel())
+ print(
+ f"[VertexVAE] Warning: No points passed threshold {inference_threshold} in batch {b.item()}. Forcing top {k} points."
+ )
+
+ _, top_local = torch.topk(batch_scores, k=k)
+
+ vertex_mask[batch_sel.nonzero(as_tuple=True)[0][top_local]] = True
+
+ vertex_x = sp.SparseTensor(
+ feats=latent.feats[vertex_mask],
+ coords=latent.coords[vertex_mask],
+ )
+
+ if verbose:
+ num_batches = int(latent.coords[:, 0].max().item()) + 1
+ print(
+ f"[VertexVAE] Shell2Vertex: "
+ f"num_vertex={per_batch_counts(vertex_x.coords[:, 0], num_batches)}, "
+ f"num_shell={per_batch_counts(latent.coords[:, 0], num_batches)}"
+ )
+
+ results.append(
+ {
+ "coords": vtx_probs.coords,
+ "occ_probs": vtx_probs.feats,
+ "vertex_mask": vertex_mask,
+ }
+ )
+ else:
+ # Training path: using gt voxels to split vertex
+ gt_vertex_coords = gt_vertex_voxels_list[0].coords
+
+ pred_flat = flatten_coords(latent.coords)
+ vertex_gt_flat = flatten_coords(gt_vertex_coords)
+
+ vertex_mask = torch.isin(pred_flat, vertex_gt_flat)
+
+ vertex_x = sp.SparseTensor(
+ feats=latent.feats[vertex_mask],
+ coords=latent.coords[vertex_mask],
+ )
+
+ results.append(
+ {
+ "coords": vtx_probs.coords,
+ "occ_probs": vtx_probs.feats,
+ "vertex_mask": vertex_mask,
+ "vertex_gt_coords": gt_vertex_coords,
+ }
+ )
+
+ vertex_x = self.vtx_proj(vertex_x)
+
+ # step1: upsample
+ for i, _ in enumerate(self.decoder_vtx):
+ vertex_x, vertex_occ_probs, num_rescue = self.decoder_vtx[i](
+ vertex_x,
+ training=training,
+ threshold=inference_threshold,
+ )
+ vertex_x = self.decoder_vtx_ca[i](
+ x=vertex_x,
+ context=self.latent_proj[i](latent_),
+ )
+
+ if not training:
+ # Inference path
+ if verbose:
+ num_batches = int(latent_.coords[:, 0].max().item()) + 1
+ print(
+ f"[VertexVAE] Layer{i}: "
+ f"num_vertex={per_batch_counts(vertex_x.coords[:, 0], num_batches)}, "
+ f"num_rescue={num_rescue}"
+ )
+
+ results.append(
+ {
+ "coords": vertex_x.coords,
+ "feats": vertex_x.feats,
+ "occ_probs": vertex_occ_probs.feats,
+ "occ_coords": vertex_occ_probs.coords,
+ }
+ )
+ else:
+ # Training path
+ vertex_pred_coords = vertex_x.coords
+ gt_vertex_coords = gt_vertex_voxels_list[i + 1].coords
+
+ vertex_pred_flat = flatten_coords(vertex_pred_coords)
+ vertex_gt_flat = flatten_coords(gt_vertex_coords)
+ vertex_mask = torch.isin(vertex_pred_flat, vertex_gt_flat)
+ vertex_prune_labels = vertex_mask.float()
+
+ vertex_x = sp.SparseTensor(
+ feats=vertex_x.feats[vertex_mask],
+ coords=vertex_x.coords[vertex_mask],
+ )
+
+ results.append(
+ {
+ "coords": vertex_x.coords,
+ "feats": vertex_x.feats,
+ "occ_probs": vertex_occ_probs.feats,
+ "occ_coords": vertex_occ_probs.coords,
+ "prune_labels": vertex_prune_labels,
+ "sp_tensor": vertex_x,
+ "gt_coords": gt_vertex_coords,
+ "pred_mask": vertex_mask,
+ },
+ )
+
+ return results
+
+ def forward(
+ self,
+ sparse_input,
+ gt_vertex_voxels_list=None,
+ training=True,
+ sample_posterior=True,
+ ):
+ latent_64, posterior = self.encode(sparse_input, sample_posterior)
+ results = self.decode(
+ latent_64,
+ gt_vertex_voxels_list=gt_vertex_voxels_list,
+ training=training,
+ )
+
+ return results, posterior, latent_64
+
+ def convert_to_fp16(self):
+ """Convert all components to float16"""
+ self.encoder.apply(
+ lambda m: m.convert_to_fp16() if hasattr(m, "convert_to_fp16") else None
+ )
+ self.decoder_vtx.apply(
+ lambda m: m.convert_to_fp16() if hasattr(m, "convert_to_fp16") else None
+ )
+ self.decoder_vtx_ca.apply(
+ lambda m: m.convert_to_fp16() if hasattr(m, "convert_to_fp16") else None
+ )
+
+ def convert_to_fp32(self):
+ """Convert all components to float32"""
+ self.encoder.apply(
+ lambda m: m.convert_to_fp32() if hasattr(m, "convert_to_fp32") else None
+ )
+ self.decoder_vtx.apply(
+ lambda m: m.convert_to_fp32() if hasattr(m, "convert_to_fp32") else None
+ )
+ self.decoder_vtx_ca.apply(
+ lambda m: m.convert_to_fp32() if hasattr(m, "convert_to_fp32") else None
+ )
diff --git a/models/vertex_structured_flow.py b/models/vertex_structured_flow.py
new file mode 100644
index 0000000000000000000000000000000000000000..d9b8fa429b2cf6debe2873bd6ccde5fa7d8c2ded
--- /dev/null
+++ b/models/vertex_structured_flow.py
@@ -0,0 +1,147 @@
+from typing import *
+import torch
+import torch.nn as nn
+import torch.nn.functional as F
+
+from modules.utils import convert_module_to_f16, convert_module_to_f32
+from modules.transformer import (
+ AbsolutePositionEmbedder,
+ TimestepEmbedder,
+)
+from modules import sparse as sp
+from modules.sparse.transformer import ModulatedSparseTransformerCrossBlock
+
+
+class VertexSLatFlowModel(nn.Module):
+ def __init__(
+ self,
+ resolution: int,
+ in_channels: int,
+ model_channels: int,
+ cond_channels: int,
+ out_channels: int,
+ num_blocks: int,
+ num_heads: Optional[int] = None,
+ num_head_channels: Optional[int] = 64,
+ mlp_ratio: float = 4,
+ pe_mode: Literal["ape", "rope"] = "ape",
+ use_fp16: bool = False,
+ use_checkpoint: bool = False,
+ share_mod: bool = False,
+ qk_rms_norm: bool = False,
+ qk_rms_norm_cross: bool = False,
+ use_density: bool = False,
+ **kwargs
+ ):
+ if kwargs:
+ print(f"[SLatFlowModel] Found unused arguments: {kwargs}")
+ super().__init__()
+ self.resolution = resolution
+ self.in_channels = in_channels
+ self.model_channels = model_channels
+ self.cond_channels = cond_channels
+ self.out_channels = out_channels
+ self.num_blocks = num_blocks
+ self.num_heads = num_heads or model_channels // num_head_channels
+ self.mlp_ratio = mlp_ratio
+ self.pe_mode = pe_mode
+ self.use_fp16 = use_fp16
+ self.use_checkpoint = use_checkpoint
+ self.share_mod = share_mod
+ self.qk_rms_norm = qk_rms_norm
+ self.qk_rms_norm_cross = qk_rms_norm_cross
+ self.use_density = use_density
+ self.dtype = torch.float16 if use_fp16 else torch.float32
+
+ self.t_embedder = TimestepEmbedder(model_channels)
+
+ if self.use_density:
+ self.density_embedder = TimestepEmbedder(model_channels)
+
+ if share_mod:
+ self.adaLN_modulation = nn.Sequential(
+ nn.SiLU(), nn.Linear(model_channels, 6 * model_channels, bias=True)
+ )
+
+ if pe_mode == "ape":
+ self.pos_embedder = AbsolutePositionEmbedder(model_channels)
+
+ self.input_layer = sp.SparseLinear(
+ in_channels,
+ model_channels,
+ )
+
+ self.blocks = nn.ModuleList(
+ [
+ ModulatedSparseTransformerCrossBlock(
+ model_channels,
+ cond_channels,
+ num_heads=self.num_heads,
+ mlp_ratio=self.mlp_ratio,
+ attn_mode="full",
+ use_checkpoint=self.use_checkpoint,
+ use_rope=(pe_mode == "rope"),
+ share_mod=self.share_mod,
+ qk_rms_norm=self.qk_rms_norm,
+ qk_rms_norm_cross=self.qk_rms_norm_cross,
+ )
+ for _ in range(self.num_blocks)
+ ]
+ )
+
+ self.out_layer = sp.SparseLinear(
+ model_channels,
+ out_channels,
+ )
+
+ if use_fp16:
+ self.convert_to_fp16()
+ else:
+ self.convert_to_fp32()
+
+ @property
+ def device(self) -> torch.device:
+ """
+ Return the device of the model.
+ """
+ return next(self.parameters()).device
+
+ def convert_to_fp16(self) -> None:
+ """
+ Convert the torso of the model to float16.
+ """
+ self.blocks.apply(convert_module_to_f16)
+
+ def convert_to_fp32(self) -> None:
+ """
+ Convert the torso of the model to float32.
+ """
+ self.blocks.apply(convert_module_to_f32)
+
+ def forward(
+ self,
+ x: sp.SparseTensor,
+ t: torch.Tensor,
+ cond: torch.Tensor,
+ density: Optional[torch.Tensor] = None,
+ ) -> sp.SparseTensor:
+ h = self.input_layer(x).type(self.dtype)
+ t_emb = self.t_embedder(t)
+ if self.use_density:
+ assert (
+ density is not None
+ ), "Density tensor must be provided when use_density is True"
+ t_emb = t_emb + self.density_embedder(density.reshape(-1).float())
+ if self.share_mod:
+ t_emb = self.adaLN_modulation(t_emb)
+ t_emb = t_emb.type(self.dtype)
+ cond = cond.type(self.dtype)
+
+ if self.pe_mode == "ape":
+ h = h + self.pos_embedder(h.coords[:, 1:]).type(self.dtype)
+ for block in self.blocks:
+ h = block(h, t_emb, cond)
+
+ h = h.replace(F.layer_norm(h.feats, h.feats.shape[-1:]))
+ h = self.out_layer(h.type(x.dtype))
+ return h
diff --git a/models/voxel_encoder.py b/models/voxel_encoder.py
new file mode 100644
index 0000000000000000000000000000000000000000..2ee703f19159887e0e191975d7a9361c226c5eac
--- /dev/null
+++ b/models/voxel_encoder.py
@@ -0,0 +1,183 @@
+from __future__ import annotations
+
+import torch
+import torch.nn as nn
+
+
+def _safe_group_norm(num_channels: int, max_groups: int = 8) -> nn.GroupNorm:
+ g = min(max_groups, num_channels)
+ while g > 1 and num_channels % g != 0:
+ g -= 1
+ return nn.GroupNorm(g, num_channels)
+
+
+def _sincos_1d(n: int, dim: int) -> torch.Tensor:
+ assert dim % 2 == 0 and dim > 0, f"sincos dim must be positive even, got {dim}"
+ pos = torch.arange(n, dtype=torch.float32)
+ omega = torch.arange(dim // 2, dtype=torch.float32) / (dim // 2)
+ omega = 1.0 / (10000**omega)
+ out = pos[:, None] * omega[None, :]
+ return torch.cat([torch.sin(out), torch.cos(out)], dim=-1)
+
+
+def _get_3d_sincos_embed(n: int, dim: int) -> torch.Tensor:
+ axis_dim = (dim // 3) // 2 * 2 # split across 3 axes, round to even
+ if axis_dim <= 0:
+ raise ValueError(
+ f"cond_in_dim={dim} too small for 3D sincos PE (need >= 6 so each axis gets a positive even slice)"
+ )
+ e = _sincos_1d(n, axis_dim) # (n, axis_dim)ß
+ pe_d = e[:, None, None, :].expand(n, n, n, axis_dim)
+ pe_h = e[None, :, None, :].expand(n, n, n, axis_dim)
+ pe_w = e[None, None, :, :].expand(n, n, n, axis_dim)
+ pe = torch.cat([pe_d, pe_h, pe_w], dim=-1).reshape(n * n * n, 3 * axis_dim)
+ if pe.shape[-1] < dim:
+ pad = torch.zeros(pe.shape[0], dim - pe.shape[-1])
+ pe = torch.cat([pe, pad], dim=-1)
+ return pe
+
+
+class _ResBlock3d(nn.Module):
+ def __init__(self, channels: int) -> None:
+ super().__init__()
+ self.block = nn.Sequential(
+ _safe_group_norm(channels),
+ nn.SiLU(),
+ nn.Conv3d(channels, channels, kernel_size=3, padding=1, bias=False),
+ _safe_group_norm(channels),
+ nn.SiLU(),
+ nn.Conv3d(channels, channels, kernel_size=3, padding=1, bias=False),
+ )
+
+ def forward(self, x: torch.Tensor) -> torch.Tensor:
+ return x + self.block(x)
+
+
+class _DownBlock3d(nn.Module):
+ def __init__(self, in_ch: int, out_ch: int, blocks_per_level: int) -> None:
+ super().__init__()
+ res_blocks: list[nn.Module] = [
+ _ResBlock3d(in_ch) for _ in range(blocks_per_level)
+ ]
+ res_blocks.append(
+ nn.Conv3d(in_ch, out_ch, kernel_size=3, stride=2, padding=1, bias=False)
+ )
+ res_blocks.append(_safe_group_norm(out_ch))
+ res_blocks.append(nn.SiLU())
+ self.net = nn.Sequential(*res_blocks)
+
+ def forward(self, x: torch.Tensor) -> torch.Tensor:
+ return self.net(x)
+
+
+class VoxelFieldConditioner(nn.Module):
+ def __init__(
+ self,
+ in_channels: int,
+ cond_in_dim: int,
+ *,
+ num_downsamples: int = 2,
+ base_channels: int = 32,
+ channel_mult: int = 2,
+ blocks_per_level: int = 2,
+ pos_embed: str = "sincos",
+ ) -> None:
+ super().__init__()
+ if num_downsamples < 0:
+ raise ValueError(f"num_downsamples must be >= 0, got {num_downsamples}")
+ if in_channels <= 0:
+ raise ValueError(f"in_channels must be > 0, got {in_channels}")
+ if cond_in_dim <= 0:
+ raise ValueError(f"cond_in_dim must be > 0, got {cond_in_dim}")
+ if base_channels <= 0:
+ raise ValueError(f"base_channels must be > 0, got {base_channels}")
+ if channel_mult < 1:
+ raise ValueError(f"channel_mult must be >= 1, got {channel_mult}")
+ if blocks_per_level < 0:
+ raise ValueError(f"blocks_per_level must be >= 0, got {blocks_per_level}")
+ pos_embed = str(pos_embed).lower()
+ if pos_embed not in ("sincos", "none"):
+ raise ValueError(
+ f"pos_embed={pos_embed!r} unsupported (use 'sincos' or 'none')"
+ )
+
+ self.in_channels = int(in_channels)
+ self.cond_in_dim = int(cond_in_dim)
+ self.num_downsamples = int(num_downsamples)
+ self.base_channels = int(base_channels)
+ self.channel_mult = int(channel_mult)
+ self.blocks_per_level = int(blocks_per_level)
+ self.pos_embed = pos_embed
+
+ self.stem = nn.Sequential(
+ nn.Conv3d(
+ self.in_channels,
+ self.base_channels,
+ kernel_size=3,
+ padding=1,
+ bias=False,
+ ),
+ _safe_group_norm(self.base_channels),
+ nn.SiLU(),
+ )
+
+ down_blocks: list[nn.Module] = []
+ ch = self.base_channels
+ for _ in range(self.num_downsamples):
+ out_ch = ch * self.channel_mult
+ down_blocks.append(_DownBlock3d(ch, out_ch, self.blocks_per_level))
+ ch = out_ch
+ self.down_blocks = nn.ModuleList(down_blocks)
+ self._final_channels = ch # base_channels * channel_mult ** num_downsamples
+
+ self.tail_blocks = nn.Sequential(
+ *[_ResBlock3d(ch) for _ in range(blocks_per_level)]
+ )
+
+ self.proj = nn.Conv3d(
+ self._final_channels, self.cond_in_dim, kernel_size=1, bias=True
+ )
+
+ def _get_pe(
+ self, n_out: int, device: torch.device, dtype: torch.dtype
+ ) -> torch.Tensor:
+ buf_name = f"_sincos_pe_{n_out}"
+ if not hasattr(self, buf_name):
+ pe = _get_3d_sincos_embed(n_out, self.cond_in_dim)
+ self.register_buffer(buf_name, pe, persistent=False)
+ return getattr(self, buf_name).to(device=device, dtype=dtype)
+
+ def forward(self, field: torch.Tensor) -> torch.Tensor:
+ """
+ Input:
+ field: (B, R, R, R) or (B, C_in, R, R, R)
+ Output:
+ (B, R'^3, cond_in_dim) token sequence with 3D PE added.
+ """
+ if field.dim() == 4:
+ field = field.unsqueeze(1)
+ elif field.dim() != 5:
+ raise ValueError(
+ f"field must be 4D (B,R,R,R) or 5D (B,C,R,R,R), got {tuple(field.shape)}"
+ )
+ if field.shape[1] != self.in_channels:
+ raise ValueError(
+ f"field channel dim {field.shape[1]} != in_channels {self.in_channels}"
+ )
+ if not (field.shape[2] == field.shape[3] == field.shape[4]):
+ raise ValueError(
+ f"field must be cubic (R,R,R), got spatial {tuple(field.shape[2:])}"
+ )
+
+ x = self.stem(field)
+ for blk in self.down_blocks:
+ x = blk(x)
+ x = self.tail_blocks(x)
+ feat = self.proj(x)
+
+ n_out = feat.shape[-1]
+ tokens = feat.flatten(2).transpose(1, 2).contiguous()
+ if self.pos_embed == "sincos":
+ pe = self._get_pe(n_out, tokens.device, tokens.dtype)
+ tokens = tokens + pe.unsqueeze(0)
+ return tokens
diff --git a/modules/attention.py b/modules/attention.py
new file mode 100644
index 0000000000000000000000000000000000000000..871417ceb94338284ecf6e9d54c54aa5cbad72f3
--- /dev/null
+++ b/modules/attention.py
@@ -0,0 +1,160 @@
+from __future__ import annotations
+
+from typing import Optional
+import torch
+import torch.nn.functional as F
+
+try:
+ from flash_attn import flash_attn_varlen_func
+
+ _FLASH_ATTN_AVAILABLE = True
+except Exception:
+ flash_attn_varlen_func = None
+ _FLASH_ATTN_AVAILABLE = False
+
+
+def flash_varlen_self_attention(
+ q: torch.Tensor,
+ k: torch.Tensor,
+ v: torch.Tensor,
+ x_mask: torch.Tensor,
+) -> torch.Tensor:
+ """q,k,v: (B, H, N, Dh); x_mask: (B, N) bool."""
+ bsz, nheads, seqlen, head_dim = q.shape
+ mask = x_mask.bool()
+ lengths = mask.sum(dim=-1, dtype=torch.int32)
+ if int(lengths.max().item()) <= 0:
+ return torch.zeros_like(q)
+
+ cu_seqlens = torch.zeros((bsz + 1,), dtype=torch.int32, device=q.device)
+ cu_seqlens[1:] = torch.cumsum(lengths, dim=0)
+ max_seqlen = int(lengths.max().item())
+
+ q_flat = q.permute(0, 2, 1, 3).reshape(bsz * seqlen, nheads, head_dim)
+ k_flat = k.permute(0, 2, 1, 3).reshape(bsz * seqlen, nheads, head_dim)
+ v_flat = v.permute(0, 2, 1, 3).reshape(bsz * seqlen, nheads, head_dim)
+ valid_token_indices = torch.nonzero(mask.reshape(-1), as_tuple=False).squeeze(-1)
+
+ q_unpad = q_flat.index_select(0, valid_token_indices)
+ k_unpad = k_flat.index_select(0, valid_token_indices)
+ v_unpad = v_flat.index_select(0, valid_token_indices)
+
+ attn_unpad = flash_attn_varlen_func(
+ q_unpad,
+ k_unpad,
+ v_unpad,
+ cu_seqlens_q=cu_seqlens,
+ cu_seqlens_k=cu_seqlens,
+ max_seqlen_q=max_seqlen,
+ max_seqlen_k=max_seqlen,
+ dropout_p=0.0,
+ causal=False,
+ )
+ out_flat = torch.zeros_like(q_flat)
+ out_flat.index_copy_(0, valid_token_indices, attn_unpad)
+ out = out_flat.reshape(bsz, seqlen, nheads, head_dim).permute(0, 2, 1, 3)
+ return out
+
+
+def flash_varlen_cross_attention(
+ q: torch.Tensor,
+ k: torch.Tensor,
+ v: torch.Tensor,
+ q_mask: torch.Tensor,
+ k_mask: torch.Tensor,
+) -> torch.Tensor:
+ """Varlen cross-attn. q: (B,H,Nq,Dh), k/v: (B,H,Nk,Dh), masks (B,Nq)/(B,Nk) bool."""
+ bsz, nheads, nq, head_dim = q.shape
+ nk = k.shape[2]
+ q_mask_b = q_mask.bool()
+ k_mask_b = k_mask.bool()
+ q_lengths = q_mask_b.sum(dim=-1, dtype=torch.int32)
+ k_lengths = k_mask_b.sum(dim=-1, dtype=torch.int32)
+ if int(q_lengths.max().item()) <= 0:
+ return torch.zeros_like(q)
+
+ cu_q = torch.zeros((bsz + 1,), dtype=torch.int32, device=q.device)
+ cu_q[1:] = torch.cumsum(q_lengths, dim=0)
+ cu_k = torch.zeros((bsz + 1,), dtype=torch.int32, device=q.device)
+ cu_k[1:] = torch.cumsum(k_lengths, dim=0)
+ max_q = int(q_lengths.max().item())
+ max_k = int(k_lengths.max().item())
+
+ q_flat = q.permute(0, 2, 1, 3).reshape(bsz * nq, nheads, head_dim)
+ k_flat = k.permute(0, 2, 1, 3).reshape(bsz * nk, nheads, head_dim)
+ v_flat = v.permute(0, 2, 1, 3).reshape(bsz * nk, nheads, head_dim)
+
+ q_idx = torch.nonzero(q_mask_b.reshape(-1), as_tuple=False).squeeze(-1)
+ k_idx = torch.nonzero(k_mask_b.reshape(-1), as_tuple=False).squeeze(-1)
+
+ q_unpad = q_flat.index_select(0, q_idx)
+ k_unpad = k_flat.index_select(0, k_idx)
+ v_unpad = v_flat.index_select(0, k_idx)
+
+ attn_unpad = flash_attn_varlen_func(
+ q_unpad,
+ k_unpad,
+ v_unpad,
+ cu_seqlens_q=cu_q,
+ cu_seqlens_k=cu_k,
+ max_seqlen_q=max_q,
+ max_seqlen_k=max_k,
+ dropout_p=0.0,
+ causal=False,
+ )
+ out_flat = torch.zeros_like(q_flat)
+ out_flat.index_copy_(0, q_idx, attn_unpad)
+ out = out_flat.reshape(bsz, nq, nheads, head_dim).permute(0, 2, 1, 3)
+ return out
+
+
+def can_flash_varlen(q: torch.Tensor, x_mask: Optional[torch.Tensor]) -> bool:
+ if not _FLASH_ATTN_AVAILABLE or x_mask is None:
+ return False
+ if not q.is_cuda:
+ return False
+ if q.dtype not in (torch.float16, torch.bfloat16):
+ return False
+ return True
+
+
+def sdpa_padding_mask(x_mask: torch.Tensor) -> torch.Tensor:
+ """(B, 1, 1, N) bool: keys valid."""
+ return x_mask.bool().view(x_mask.shape[0], 1, 1, x_mask.shape[1])
+
+
+def graph_adj_varlen_attention(
+ q: torch.Tensor,
+ k: torch.Tensor,
+ v: torch.Tensor,
+ x_mask: torch.Tensor,
+ adj_matrix: Optional[torch.Tensor],
+) -> torch.Tensor:
+ bsz, nheads, seqlen, _ = q.shape
+ out = torch.zeros_like(q)
+ x_mask = x_mask.bool()
+
+ for b in range(bsz):
+ valid_indices = torch.nonzero(x_mask[b], as_tuple=False).squeeze(-1)
+ if valid_indices.numel() == 0:
+ continue
+ q_b = q[b].index_select(1, valid_indices).unsqueeze(0)
+ k_b = k[b].index_select(1, valid_indices).unsqueeze(0)
+ v_b = v[b].index_select(1, valid_indices).unsqueeze(0)
+ l_now = valid_indices.numel()
+
+ if adj_matrix is not None:
+ sub = (
+ adj_matrix[b]
+ .bool()
+ .index_select(0, valid_indices)
+ .index_select(1, valid_indices)
+ )
+ eye = torch.eye(l_now, dtype=torch.bool, device=q.device)
+ attn_mask_b = (sub | eye).view(1, 1, l_now, l_now)
+ else:
+ attn_mask_b = None
+
+ out_b = F.scaled_dot_product_attention(q_b, k_b, v_b, attn_mask=attn_mask_b)
+ out[b].index_copy_(1, valid_indices, out_b.squeeze(0))
+ return out
diff --git a/modules/norm.py b/modules/norm.py
new file mode 100644
index 0000000000000000000000000000000000000000..4900f6d075f04678514ee0b0a8ac2b070057ad6d
--- /dev/null
+++ b/modules/norm.py
@@ -0,0 +1,41 @@
+import torch
+import torch.nn as nn
+
+
+class LayerNorm32(nn.LayerNorm):
+ def forward(self, x: torch.Tensor) -> torch.Tensor:
+ return super().forward(x.float()).type(x.dtype)
+
+
+class RMSNorm(nn.Module):
+ def __init__(
+ self,
+ dim: int,
+ eps: float = 1e-5,
+ elementwise_affine: bool = True,
+ variance_in_fp32: bool = True,
+ ):
+ super().__init__()
+ self.eps = eps
+ self.elementwise_affine = elementwise_affine
+ self.variance_in_fp32 = bool(variance_in_fp32)
+ self.weight = nn.Parameter(torch.ones(dim)) if elementwise_affine else None
+
+ def forward(self, x: torch.Tensor) -> torch.Tensor:
+ input_dtype = x.dtype
+ if self.variance_in_fp32:
+ variance = x.to(torch.float32).pow(2).mean(-1, keepdim=True)
+ inv_rms = torch.rsqrt(variance + self.eps).to(input_dtype)
+ else:
+ variance = x.pow(2).mean(-1, keepdim=True)
+ inv_rms = torch.rsqrt(variance + self.eps)
+ x = x * inv_rms
+ if self.weight is not None:
+ w = self.weight
+ if w.dtype in (torch.float16, torch.bfloat16):
+ x = (x.to(w.dtype) * w).to(input_dtype)
+ else:
+ x = (x * w).to(input_dtype)
+ else:
+ x = x.to(input_dtype)
+ return x
diff --git a/modules/pointnet.py b/modules/pointnet.py
new file mode 100644
index 0000000000000000000000000000000000000000..ebc939d3e983d47e4f0c56648351441080335d76
--- /dev/null
+++ b/modules/pointnet.py
@@ -0,0 +1,330 @@
+# MIT License
+
+# Copyright (c) 2020 Songyou Peng, Michael Niemeyer, Lars Mescheder, Marc Pollefeys, Andreas Geiger.
+# Copyright (c) 2025 VAST-AI-Research and contributors.
+
+# Permission is hereby granted, free of charge, to any person obtaining a copy
+# of this software and associated documentation files (the "Software"), to deal
+# in the Software without restriction, including without limitation the rights
+# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
+# copies of the Software, and to permit persons to whom the Software is
+# furnished to do so, subject to the following conditions:
+
+# The above copyright notice and this permission notice shall be included in all
+# copies or substantial portions of the Software.
+
+# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
+# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
+# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
+# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
+# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
+# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
+# SOFTWARE
+
+# modified from https://github.com/autonomousvision/convolutional_occupancy_networks/blob/master/src/encoder/pointnet.py
+# modified from https://github.com/VAST-AI-Research/TripoSF/blob/main/triposf/modules/pointclouds/pointnet.py
+
+import torch
+import torch.nn as nn
+import copy
+from torch import Tensor
+from torch_scatter import scatter_mean
+from torch.utils.checkpoint import checkpoint
+
+
+def scale_tensor(dat, inp_scale=None, tgt_scale=None):
+ if inp_scale is None:
+ inp_scale = (-0.5, 0.5)
+ if tgt_scale is None:
+ tgt_scale = (0, 1)
+ assert tgt_scale[1] > tgt_scale[0] and inp_scale[1] > inp_scale[0]
+ if isinstance(tgt_scale, Tensor):
+ assert dat.shape[-1] == tgt_scale.shape[-1]
+ dat = (dat - inp_scale[0]) / (inp_scale[1] - inp_scale[0])
+ dat = dat * (tgt_scale[1] - tgt_scale[0]) + tgt_scale[0]
+ return dat.clamp(tgt_scale[0] + 1e-6, tgt_scale[1] - 1e-6)
+
+
+# Resnet Blocks for pointnet
+class ResnetBlockFC(nn.Module):
+ """Fully connected ResNet Block class.
+
+ Args:
+ size_in (int): input dimension
+ size_out (int): output dimension
+ size_h (int): hidden dimension
+ """
+
+ def __init__(self, size_in, size_out=None, size_h=None):
+ super().__init__()
+ # Attributes
+ if size_out is None:
+ size_out = size_in
+
+ if size_h is None:
+ size_h = min(size_in, size_out)
+
+ self.size_in = size_in
+ self.size_h = size_h
+ self.size_out = size_out
+ # Submodules
+ self.fc_0 = nn.Linear(size_in, size_h)
+ self.fc_1 = nn.Linear(size_h, size_out)
+ self.actvn = nn.GELU(approximate="tanh")
+
+ if size_in == size_out:
+ self.shortcut = None
+ else:
+ self.shortcut = nn.Linear(size_in, size_out, bias=False)
+ # Initialization
+ nn.init.xavier_uniform_(self.fc_0.weight)
+ if self.fc_0.bias is not None:
+ nn.init.constant_(self.fc_0.bias, 0)
+ if self.shortcut is not None:
+ nn.init.xavier_uniform_(self.shortcut.weight)
+ if self.shortcut.bias is not None:
+ nn.init.constant_(self.shortcut.bias, 0)
+
+ nn.init.xavier_uniform_(self.fc_1.weight)
+ if self.fc_1.bias is not None:
+ nn.init.constant_(self.fc_1.bias, 0)
+
+ def forward(self, x):
+ net = self.fc_0(self.actvn(x))
+ dx = self.fc_1(self.actvn(net))
+
+ if self.shortcut is not None:
+ x_s = self.shortcut(x)
+ else:
+ x_s = x
+
+ return x_s + dx
+
+
+class LocalPoolPointnet(nn.Module):
+ def __init__(
+ self,
+ in_channels=3,
+ out_channels=128,
+ hidden_dim=128,
+ scatter_type="mean",
+ n_blocks=5,
+ ):
+ super().__init__()
+ self.scatter_type = scatter_type
+ self.in_channels = in_channels
+ self.hidden_dim = hidden_dim
+ self.out_channels = out_channels
+ self.fc_pos = nn.Linear(in_channels, 2 * hidden_dim)
+ self.blocks = nn.ModuleList(
+ [ResnetBlockFC(2 * hidden_dim, hidden_dim) for i in range(n_blocks)]
+ )
+ self.fc_c = nn.Linear(hidden_dim, out_channels)
+ self.in_channels = in_channels
+ if self.scatter_type == "mean":
+ self.scatter = scatter_mean
+ else:
+ raise ValueError("Incorrect scatter type")
+ self.initialize_weights()
+
+ def initialize_weights(self):
+
+ nn.init.xavier_uniform_(self.fc_pos.weight)
+ if self.fc_pos.bias is not None:
+ nn.init.constant_(self.fc_pos.bias, 0)
+
+ nn.init.xavier_uniform_(self.fc_c.weight)
+ if self.fc_c.bias is not None:
+ nn.init.constant_(self.fc_c.bias, 0)
+
+ def convert_to_sparse_feats(self, c, sparse_coords):
+ """
+ Input:
+ sparse_coords: Tensor [Nx, 4], point to sparse indices
+ c: Tensor [B, res, C], input feats of each grid
+ Output:
+ c_out: Tensor [B, Np, C], aggregated grid feats of each point
+ """
+ feats_new = torch.zeros(
+ (sparse_coords.shape[0], c.shape[-1]), device=c.device, dtype=c.dtype
+ )
+ offsets = 0
+
+ batch_nums = copy.deepcopy(sparse_coords[..., 0])
+ for i in range(len(c)):
+ coords_num_i = (batch_nums == i).sum()
+ feats_new[offsets : offsets + coords_num_i] = c[i, :coords_num_i]
+ offsets += coords_num_i
+ return feats_new
+
+ def generate_sparse_grid_features(self, index, c, max_coord_num):
+ # scatter grid features from points
+ bs, fea_dim = c.size(0), c.size(2)
+ res = max_coord_num
+ c_out = c.new_zeros(bs, self.out_channels, res)
+ c_out = scatter_mean(c.permute(0, 2, 1), index, out=c_out).permute(
+ 0, 2, 1
+ ) # B x res X C
+ return c_out
+
+ def pool_sparse_local(self, index, c, max_coord_num):
+ """
+ Input:
+ index: Tensor [B, 1, Np], sparse indices of each point
+ c: Tensor [B, Np, C], input feats of each point
+ Output:
+ c_out: Tensor [B, Np, C], aggregated grid feats of each point
+ """
+
+ bs, fea_dim = c.size(0), c.size(2)
+ res = max_coord_num
+ c_out = c.new_zeros(bs, fea_dim, res)
+ c_out = self.scatter(c.permute(0, 2, 1), index, out=c_out)
+
+ # gather feature back to points
+ c_out = c_out.gather(dim=2, index=index.expand(-1, fea_dim, -1))
+ return c_out.permute(0, 2, 1)
+
+ @torch.no_grad()
+ def coordinate2sparseindex(self, x, sparse_coords, res):
+ """
+ Input:
+ x: Tensor [B, Np, 3], points scaled at ([0, 1] * res)
+ sparse_coords: Tensor [Nx, 4] ([batch_number, x, y, z])
+ res: Int, resolution of the grid index
+ Output:
+ sparse_index: Tensor [B, 1, Np], sparse indices of each point
+ """
+ B = x.shape[0]
+ sparse_index = torch.zeros((B, x.shape[1]), device=x.device, dtype=torch.int64)
+
+ index = (x[..., 0] * res + x[..., 1]) * res + x[..., 2]
+ sparse_indices = copy.deepcopy(sparse_coords)
+ sparse_indices[..., 1] = (
+ sparse_indices[..., 1] * res + sparse_indices[..., 2]
+ ) * res + sparse_indices[..., 3]
+ sparse_indices = sparse_indices[..., :2]
+
+ for i in range(B):
+ mask_i = sparse_indices[..., 0] == i
+ coords_i = sparse_indices[mask_i, 1]
+ coords_num_i = len(coords_i)
+ sparse_index[i] = torch.searchsorted(coords_i, index[i])
+
+ return sparse_index[:, None, :]
+
+ def forward(self, p, sparse_coords, res=64, bbox_size=(-0.5, 0.5)):
+ """
+ Input:
+ p : Tensor [B, Np(819_200), 3]
+ sparse_coords: Tensor [Nx, 4] ([batch_number, x, y, z])
+
+ Output:
+ sparse_pc_feats: [Nx, self.out_channels]
+ """
+ batch_size, T, D = p.size()
+ max_coord_num = 0
+ for i in range(batch_size):
+ max_coord_num = max(
+ max_coord_num, (sparse_coords[..., 0] == i).sum().item() + 5
+ )
+
+ if D == self.in_channels:
+ p, normals = p[..., :3], p[..., 3:]
+
+ coord = scale_tensor(p, inp_scale=bbox_size) * res
+ p = 2 * (coord - (coord.floor() + 0.5)) # dist to the centrios, [-1., 1.]
+ index = self.coordinate2sparseindex(coord.long(), sparse_coords, res)
+
+ if D == self.in_channels:
+ p = torch.cat((p, normals), dim=-1)
+ net = self.fc_pos(p)
+ net = self.blocks[0](net)
+ for block in self.blocks[1:]:
+ pooled = self.pool_sparse_local(index, net, max_coord_num=max_coord_num)
+
+ net = torch.cat([net, pooled], dim=2)
+ net = block(net)
+ c = self.fc_c(net)
+ feats = self.generate_sparse_grid_features(
+ index, c, max_coord_num=max_coord_num
+ )
+ feats = self.convert_to_sparse_feats(feats, sparse_coords)
+
+ # torch.cuda.empty_cache()
+ return feats
+
+
+class Pointnet(nn.Module):
+ def __init__(
+ self,
+ in_channels=16,
+ out_channels=32,
+ hidden_dim=32,
+ n_blocks=5,
+ use_checkpoint=True,
+ ):
+ super().__init__()
+ self.in_channels = in_channels
+ self.out_channels = out_channels
+ self.hidden_dim = hidden_dim
+ self.use_checkpoint = use_checkpoint
+
+ self.fc_pos = nn.Linear(in_channels, 2 * hidden_dim)
+
+ self.blocks = nn.ModuleList(
+ [ResnetBlockFC(2 * hidden_dim, hidden_dim) for i in range(n_blocks)]
+ )
+
+ self.fc_c = nn.Linear(hidden_dim, out_channels)
+
+ self.initialize_weights()
+
+ def initialize_weights(self):
+ nn.init.xavier_uniform_(self.fc_pos.weight)
+ if self.fc_pos.bias is not None:
+ nn.init.constant_(self.fc_pos.bias, 0)
+
+ nn.init.xavier_uniform_(self.fc_c.weight)
+ if self.fc_c.bias is not None:
+ nn.init.constant_(self.fc_c.bias, 0)
+
+ @staticmethod
+ def _forward_block_concat(module, x):
+ return module(torch.cat([x, x], dim=-1))
+
+ def forward(self, p, res=64, bbox_size=(-0.5, 0.5)):
+ """
+ Input:
+ p : Tensor [M, in_channels]
+ Output:
+ feats: Tensor [M, out_channels]
+ """
+
+ pos_world = p[..., 0:3] # [M, 3]
+ other_feats = p[..., 3:] # [M, in_channels - 3]
+
+ scaled_pos = scale_tensor(pos_world, inp_scale=bbox_size) * res
+ local_pos = 2 * (scaled_pos - (scaled_pos.floor() + 0.5))
+
+ net_input = torch.cat([local_pos, other_feats], dim=-1)
+
+ net = self.fc_pos(net_input)
+
+ if self.use_checkpoint and net.requires_grad:
+ net = checkpoint(self.blocks[0], net, use_reentrant=False)
+ else:
+ net = self.blocks[0](net)
+
+ for block in self.blocks[1:]:
+ if self.use_checkpoint and net.requires_grad:
+ net = checkpoint(
+ self._forward_block_concat, block, net, use_reentrant=False
+ )
+ else:
+ net_concat = torch.cat([net, net], dim=-1)
+ net = block(net_concat)
+
+ feats = self.fc_c(net)
+
+ return feats
diff --git a/modules/sparse/__init__.py b/modules/sparse/__init__.py
new file mode 100644
index 0000000000000000000000000000000000000000..fd47cc26a05db5305bdc86141f3e8b12cc16f3eb
--- /dev/null
+++ b/modules/sparse/__init__.py
@@ -0,0 +1,130 @@
+# MIT License
+
+# Copyright (c) Microsoft Corporation.
+# Copyright (c) 2025 VAST-AI-Research and contributors.
+
+# Permission is hereby granted, free of charge, to any person obtaining a copy
+# of this software and associated documentation files (the "Software"), to deal
+# in the Software without restriction, including without limitation the rights
+# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
+# copies of the Software, and to permit persons to whom the Software is
+# furnished to do so, subject to the following conditions:
+
+# The above copyright notice and this permission notice shall be included in all
+# copies or substantial portions of the Software.
+
+# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
+# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
+# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
+# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
+# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
+# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
+# SOFTWARE
+
+from typing import *
+
+BACKEND = 'spconv'
+DEBUG = False
+ATTN = 'flash_attn'
+
+def __from_env():
+ import os
+
+ global BACKEND
+ global DEBUG
+ global ATTN
+
+ env_sparse_backend = os.environ.get('SPARSE_BACKEND')
+ env_sparse_debug = os.environ.get('SPARSE_DEBUG')
+ env_sparse_attn = os.environ.get('SPARSE_ATTN_BACKEND')
+ if env_sparse_attn is None:
+ env_sparse_attn = os.environ.get('ATTN_BACKEND')
+
+ if env_sparse_backend is not None and env_sparse_backend in ['spconv', 'torchsparse']:
+ BACKEND = env_sparse_backend
+ if env_sparse_debug is not None:
+ DEBUG = env_sparse_debug == '1'
+ if env_sparse_attn is not None and env_sparse_attn in ['xformers', 'flash_attn']:
+ ATTN = env_sparse_attn
+
+ print(f"[SPARSE] Backend: {BACKEND}, Attention: {ATTN}")
+
+
+__from_env()
+
+
+def set_backend(backend: Literal['spconv', 'torchsparse']):
+ global BACKEND
+ BACKEND = backend
+
+def set_debug(debug: bool):
+ global DEBUG
+ DEBUG = debug
+
+def set_attn(attn: Literal['xformers', 'flash_attn']):
+ global ATTN
+ ATTN = attn
+
+
+import importlib
+
+__attributes = {
+ 'SparseTensor': 'basic',
+ 'sparse_batch_broadcast': 'basic',
+ 'sparse_batch_op': 'basic',
+ 'sparse_cat': 'basic',
+ 'sparse_unbind': 'basic',
+ 'SparseGroupNorm': 'norm',
+ 'SparseLayerNorm': 'norm',
+ 'SparseGroupNorm32': 'norm',
+ 'SparseLayerNorm32': 'norm',
+ 'SparseReLU': 'nonlinearity',
+ 'SparseSiLU': 'nonlinearity',
+ 'SparseGELU': 'nonlinearity',
+ 'SparseActivation': 'nonlinearity',
+ 'SparseLinear': 'linear',
+ 'sparse_scaled_dot_product_attention': 'attention',
+ 'SerializeMode': 'attention',
+ 'SerializeModes': 'attention',
+ 'sparse_serialized_scaled_dot_product_self_attention': 'attention',
+ 'sparse_windowed_scaled_dot_product_self_attention': 'attention',
+ 'SparseMultiHeadAttention': 'attention',
+ 'SparseConv3d': 'conv',
+ 'SparseInverseConv3d': 'conv',
+ 'SparseDownsample': 'spatial',
+ 'SparseUpsample': 'spatial',
+ 'SparseSubdivide' : 'spatial',
+
+ 'SparseSubdivide_attn' : 'spatial',
+ 'SparseSpatial2Channel': 'spatial',
+ 'SparseChannel2Spatial': 'spatial',
+}
+
+__submodules = ['transformer']
+
+__all__ = list(__attributes.keys()) + __submodules
+
+def __getattr__(name):
+ if name not in globals():
+ if name in __attributes:
+ module_name = __attributes[name]
+ module = importlib.import_module(f".{module_name}", __name__)
+ globals()[name] = getattr(module, name)
+ elif name in __submodules:
+ module = importlib.import_module(f".{name}", __name__)
+ globals()[name] = module
+ else:
+ raise AttributeError(f"module {__name__} has no attribute {name}")
+ return globals()[name]
+
+
+# For Pylance
+if __name__ == '__main__':
+ from .basic import *
+ from .norm import *
+ from .nonlinearity import *
+ from .linear import *
+ from .attention import *
+ from .conv import *
+ from .spatial import *
+ import transformer
diff --git a/modules/sparse/attention/__init__.py b/modules/sparse/attention/__init__.py
new file mode 100644
index 0000000000000000000000000000000000000000..ada9cdd0e3e3739eb72299ec16111d7c56d4b57e
--- /dev/null
+++ b/modules/sparse/attention/__init__.py
@@ -0,0 +1,27 @@
+# MIT License
+
+# Copyright (c) Microsoft Corporation.
+# Copyright (c) 2025 VAST-AI-Research and contributors.
+
+# Permission is hereby granted, free of charge, to any person obtaining a copy
+# of this software and associated documentation files (the "Software"), to deal
+# in the Software without restriction, including without limitation the rights
+# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
+# copies of the Software, and to permit persons to whom the Software is
+# furnished to do so, subject to the following conditions:
+
+# The above copyright notice and this permission notice shall be included in all
+# copies or substantial portions of the Software.
+
+# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
+# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
+# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
+# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
+# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
+# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
+# SOFTWARE
+
+from .full_attn import *
+from .serialized_attn import *
+from .windowed_attn import *
+from .modules import *
diff --git a/modules/sparse/attention/full_attn.py b/modules/sparse/attention/full_attn.py
new file mode 100644
index 0000000000000000000000000000000000000000..6238b0c6d77c599ac89dbf82c594fdcebfa83b5e
--- /dev/null
+++ b/modules/sparse/attention/full_attn.py
@@ -0,0 +1,238 @@
+# MIT License
+
+# Copyright (c) Microsoft Corporation.
+# Copyright (c) 2025 VAST-AI-Research and contributors.
+
+# Permission is hereby granted, free of charge, to any person obtaining a copy
+# of this software and associated documentation files (the "Software"), to deal
+# in the Software without restriction, including without limitation the rights
+# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
+# copies of the Software, and to permit persons to whom the Software is
+# furnished to do so, subject to the following conditions:
+
+# The above copyright notice and this permission notice shall be included in all
+# copies or substantial portions of the Software.
+
+# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
+# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
+# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
+# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
+# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
+# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
+# SOFTWARE
+
+from typing import *
+import torch
+from .. import SparseTensor
+from .. import DEBUG, ATTN
+
+if ATTN == 'xformers':
+ import xformers.ops as xops
+elif ATTN == 'flash_attn':
+ import flash_attn
+else:
+ raise ValueError(f"Unknown attention module: {ATTN}")
+
+
+__all__ = [
+ 'sparse_scaled_dot_product_attention',
+]
+
+
+@overload
+def sparse_scaled_dot_product_attention(qkv: SparseTensor) -> SparseTensor:
+ """
+ Apply scaled dot product attention to a sparse tensor.
+
+ Args:
+ qkv (SparseTensor): A [N, *, 3, H, C] sparse tensor containing Qs, Ks, and Vs.
+ """
+ ...
+
+@overload
+def sparse_scaled_dot_product_attention(q: SparseTensor, kv: Union[SparseTensor, torch.Tensor]) -> SparseTensor:
+ """
+ Apply scaled dot product attention to a sparse tensor.
+
+ Args:
+ q (SparseTensor): A [N, *, H, C] sparse tensor containing Qs.
+ 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.
+ """
+ ...
+
+@overload
+def sparse_scaled_dot_product_attention(q: torch.Tensor, kv: SparseTensor) -> torch.Tensor:
+ """
+ Apply scaled dot product attention to a sparse tensor.
+
+ Args:
+ q (SparseTensor): A [N, L, H, C] dense tensor containing Qs.
+ kv (SparseTensor or torch.Tensor): A [N, *, 2, H, C] sparse tensor containing Ks and Vs.
+ """
+ ...
+
+@overload
+def sparse_scaled_dot_product_attention(q: SparseTensor, k: SparseTensor, v: SparseTensor) -> SparseTensor:
+ """
+ Apply scaled dot product attention to a sparse tensor.
+
+ Args:
+ q (SparseTensor): A [N, *, H, Ci] sparse tensor containing Qs.
+ k (SparseTensor): A [N, *, H, Ci] sparse tensor containing Ks.
+ v (SparseTensor): A [N, *, H, Co] sparse tensor containing Vs.
+
+ Note:
+ k and v are assumed to have the same coordinate map.
+ """
+ ...
+
+@overload
+def sparse_scaled_dot_product_attention(q: SparseTensor, k: torch.Tensor, v: torch.Tensor) -> SparseTensor:
+ """
+ Apply scaled dot product attention to a sparse tensor.
+
+ Args:
+ q (SparseTensor): A [N, *, H, Ci] sparse tensor containing Qs.
+ k (torch.Tensor): A [N, L, H, Ci] dense tensor containing Ks.
+ v (torch.Tensor): A [N, L, H, Co] dense tensor containing Vs.
+ """
+ ...
+
+@overload
+def sparse_scaled_dot_product_attention(q: torch.Tensor, k: SparseTensor, v: SparseTensor) -> torch.Tensor:
+ """
+ Apply scaled dot product attention to a sparse tensor.
+
+ Args:
+ q (torch.Tensor): A [N, L, H, Ci] dense tensor containing Qs.
+ k (SparseTensor): A [N, *, H, Ci] sparse tensor containing Ks.
+ v (SparseTensor): A [N, *, H, Co] sparse tensor containing Vs.
+ """
+ ...
+
+def sparse_scaled_dot_product_attention(*args, **kwargs):
+ arg_names_dict = {
+ 1: ['qkv'],
+ 2: ['q', 'kv'],
+ 3: ['q', 'k', 'v']
+ }
+ num_all_args = len(args) + len(kwargs)
+ assert num_all_args in arg_names_dict, f"Invalid number of arguments, got {num_all_args}, expected 1, 2, or 3"
+ for key in arg_names_dict[num_all_args][len(args):]:
+ assert key in kwargs, f"Missing argument {key}"
+
+ if num_all_args == 1:
+ qkv = args[0] if len(args) > 0 else kwargs['qkv']
+ assert isinstance(qkv, SparseTensor), f"qkv must be a SparseTensor, got {type(qkv)}"
+ assert len(qkv.shape) == 4 and qkv.shape[1] == 3, f"Invalid shape for qkv, got {qkv.shape}, expected [N, *, 3, H, C]"
+ device = qkv.device
+
+ s = qkv
+ q_seqlen = [qkv.layout[i].stop - qkv.layout[i].start for i in range(qkv.shape[0])]
+ kv_seqlen = q_seqlen
+ qkv = qkv.feats # [T, 3, H, C]
+
+ elif num_all_args == 2:
+ q = args[0] if len(args) > 0 else kwargs['q']
+ kv = args[1] if len(args) > 1 else kwargs['kv']
+ assert isinstance(q, SparseTensor) and isinstance(kv, (SparseTensor, torch.Tensor)) or \
+ isinstance(q, torch.Tensor) and isinstance(kv, SparseTensor), \
+ f"Invalid types, got {type(q)} and {type(kv)}"
+ assert q.shape[0] == kv.shape[0], f"Batch size mismatch, got {q.shape[0]} and {kv.shape[0]}"
+ device = q.device
+
+ if isinstance(q, SparseTensor):
+ assert len(q.shape) == 3, f"Invalid shape for q, got {q.shape}, expected [N, *, H, C]"
+ s = q
+ q_seqlen = [q.layout[i].stop - q.layout[i].start for i in range(q.shape[0])]
+ q = q.feats # [T_Q, H, C]
+ else:
+ assert len(q.shape) == 4, f"Invalid shape for q, got {q.shape}, expected [N, L, H, C]"
+ s = None
+ N, L, H, C = q.shape
+ q_seqlen = [L] * N
+ q = q.reshape(N * L, H, C) # [T_Q, H, C]
+
+ if isinstance(kv, SparseTensor):
+ assert len(kv.shape) == 4 and kv.shape[1] == 2, f"Invalid shape for kv, got {kv.shape}, expected [N, *, 2, H, C]"
+ kv_seqlen = [kv.layout[i].stop - kv.layout[i].start for i in range(kv.shape[0])]
+ kv = kv.feats # [T_KV, 2, H, C]
+ else:
+ assert len(kv.shape) == 5, f"Invalid shape for kv, got {kv.shape}, expected [N, L, 2, H, C]"
+ N, L, _, H, C = kv.shape
+ kv_seqlen = [L] * N
+ kv = kv.reshape(N * L, 2, H, C) # [T_KV, 2, H, C]
+
+ elif num_all_args == 3:
+ q = args[0] if len(args) > 0 else kwargs['q']
+ k = args[1] if len(args) > 1 else kwargs['k']
+ v = args[2] if len(args) > 2 else kwargs['v']
+ assert isinstance(q, SparseTensor) and isinstance(k, (SparseTensor, torch.Tensor)) and type(k) == type(v) or \
+ isinstance(q, torch.Tensor) and isinstance(k, SparseTensor) and isinstance(v, SparseTensor), \
+ f"Invalid types, got {type(q)}, {type(k)}, and {type(v)}"
+ 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]}"
+ device = q.device
+
+ if isinstance(q, SparseTensor):
+ assert len(q.shape) == 3, f"Invalid shape for q, got {q.shape}, expected [N, *, H, Ci]"
+ s = q
+ q_seqlen = [q.layout[i].stop - q.layout[i].start for i in range(q.shape[0])]
+ q = q.feats # [T_Q, H, Ci]
+ else:
+ assert len(q.shape) == 4, f"Invalid shape for q, got {q.shape}, expected [N, L, H, Ci]"
+ s = None
+ N, L, H, CI = q.shape
+ q_seqlen = [L] * N
+ q = q.reshape(N * L, H, CI) # [T_Q, H, Ci]
+
+ if isinstance(k, SparseTensor):
+ assert len(k.shape) == 3, f"Invalid shape for k, got {k.shape}, expected [N, *, H, Ci]"
+ assert len(v.shape) == 3, f"Invalid shape for v, got {v.shape}, expected [N, *, H, Co]"
+ kv_seqlen = [k.layout[i].stop - k.layout[i].start for i in range(k.shape[0])]
+ k = k.feats # [T_KV, H, Ci]
+ v = v.feats # [T_KV, H, Co]
+ else:
+ assert len(k.shape) == 4, f"Invalid shape for k, got {k.shape}, expected [N, L, H, Ci]"
+ assert len(v.shape) == 4, f"Invalid shape for v, got {v.shape}, expected [N, L, H, Co]"
+ N, L, H, CI, CO = *k.shape, v.shape[-1]
+ kv_seqlen = [L] * N
+ k = k.reshape(N * L, H, CI) # [T_KV, H, Ci]
+ v = v.reshape(N * L, H, CO) # [T_KV, H, Co]
+
+ if DEBUG:
+ if s is not None:
+ for i in range(s.shape[0]):
+ assert (s.coords[s.layout[i]] == i).all(), f"SparseScaledDotProductSelfAttention: batch index mismatch"
+ if num_all_args in [2, 3]:
+ assert q.shape[:2] == [1, sum(q_seqlen)], f"SparseScaledDotProductSelfAttention: q shape mismatch"
+ if num_all_args == 3:
+ assert k.shape[:2] == [1, sum(kv_seqlen)], f"SparseScaledDotProductSelfAttention: k shape mismatch"
+ assert v.shape[:2] == [1, sum(kv_seqlen)], f"SparseScaledDotProductSelfAttention: v shape mismatch"
+
+ if ATTN == 'xformers':
+ if num_all_args == 1:
+ q, k, v = qkv.unbind(dim=1)
+ elif num_all_args == 2:
+ k, v = kv.unbind(dim=1)
+ q = q.unsqueeze(0)
+ k = k.unsqueeze(0)
+ v = v.unsqueeze(0)
+ mask = xops.fmha.BlockDiagonalMask.from_seqlens(q_seqlen, kv_seqlen)
+ out = xops.memory_efficient_attention(q, k, v, mask)[0]
+ elif ATTN == 'flash_attn':
+ cu_seqlens_q = torch.cat([torch.tensor([0]), torch.cumsum(torch.tensor(q_seqlen), dim=0)]).int().to(device)
+ if num_all_args in [2, 3]:
+ cu_seqlens_kv = torch.cat([torch.tensor([0]), torch.cumsum(torch.tensor(kv_seqlen), dim=0)]).int().to(device)
+ if num_all_args == 1:
+ out = flash_attn.flash_attn_varlen_qkvpacked_func(qkv, cu_seqlens_q, max(q_seqlen))
+ elif num_all_args == 2:
+ out = flash_attn.flash_attn_varlen_kvpacked_func(q, kv, cu_seqlens_q, cu_seqlens_kv, max(q_seqlen), max(kv_seqlen))
+ elif num_all_args == 3:
+ out = flash_attn.flash_attn_varlen_func(q, k, v, cu_seqlens_q, cu_seqlens_kv, max(q_seqlen), max(kv_seqlen))
+ else:
+ raise ValueError(f"Unknown attention module: {ATTN}")
+
+ if s is not None:
+ return s.replace(out)
+ else:
+ return out.reshape(N, L, H, -1)
diff --git a/modules/sparse/attention/modules.py b/modules/sparse/attention/modules.py
new file mode 100644
index 0000000000000000000000000000000000000000..4e48a354e99b2c78bb41c0aad3e08a7eddc08782
--- /dev/null
+++ b/modules/sparse/attention/modules.py
@@ -0,0 +1,214 @@
+# MIT License
+
+# Copyright (c) Microsoft Corporation.
+# Copyright (c) 2025 VAST-AI-Research and contributors.
+
+# Permission is hereby granted, free of charge, to any person obtaining a copy
+# of this software and associated documentation files (the "Software"), to deal
+# in the Software without restriction, including without limitation the rights
+# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
+# copies of the Software, and to permit persons to whom the Software is
+# furnished to do so, subject to the following conditions:
+
+# The above copyright notice and this permission notice shall be included in all
+# copies or substantial portions of the Software.
+
+# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
+# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
+# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
+# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
+# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
+# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
+# SOFTWARE
+
+from typing import *
+import torch
+import torch.nn as nn
+import torch.nn.functional as F
+from .. import SparseTensor
+from .full_attn import sparse_scaled_dot_product_attention
+from .serialized_attn import SerializeMode, sparse_serialized_scaled_dot_product_self_attention
+from .windowed_attn import sparse_windowed_scaled_dot_product_self_attention
+
+
+class RotaryPositionEmbedder(nn.Module):
+ def __init__(self, hidden_size: int, in_channels: int = 3):
+ super().__init__()
+ assert hidden_size % 2 == 0, "Hidden size must be divisible by 2"
+ self.hidden_size = hidden_size
+ self.in_channels = in_channels
+ self.freq_dim = hidden_size // in_channels // 2
+ self.freqs = torch.arange(self.freq_dim, dtype=torch.float32) / self.freq_dim
+ self.freqs = 1.0 / (10000 ** self.freqs)
+
+ def _get_phases(self, indices: torch.Tensor) -> torch.Tensor:
+ self.freqs = self.freqs.to(indices.device)
+ phases = torch.outer(indices, self.freqs)
+ phases = torch.polar(torch.ones_like(phases), phases)
+ return phases
+
+ def _rotary_embedding(self, x: torch.Tensor, phases: torch.Tensor) -> torch.Tensor:
+ x_complex = torch.view_as_complex(x.float().reshape(*x.shape[:-1], -1, 2))
+
+ if phases.dim() == x_complex.dim() - 1:
+ phases = phases.unsqueeze(-2)
+
+ x_rotated = x_complex * phases
+ x_embed = torch.view_as_real(x_rotated).reshape(*x_rotated.shape[:-1], -1).to(x.dtype)
+ return x_embed
+
+ def forward(self, q: torch.Tensor, k: torch.Tensor, indices: Optional[torch.Tensor] = None) -> Tuple[torch.Tensor, torch.Tensor]:
+ """
+ Args:
+ q (torch.Tensor): [..., N, D] tensor of queries
+ k (torch.Tensor): [..., N, D] tensor of keys
+ indices (torch.Tensor): [..., N, C] tensor of spatial positions
+ """
+ if indices is None:
+ indices = torch.arange(q.shape[-2], device=q.device)
+ if len(q.shape) > 2:
+ indices = indices.unsqueeze(0).expand(q.shape[:-2] + (-1,))
+
+ phases = self._get_phases(indices.reshape(-1)).reshape(*indices.shape[:-1], -1)
+ if phases.shape[1] < self.hidden_size // 2:
+ phases = torch.cat([phases, torch.polar(
+ torch.ones(*phases.shape[:-1], self.hidden_size // 2 - phases.shape[1], device=phases.device),
+ torch.zeros(*phases.shape[:-1], self.hidden_size // 2 - phases.shape[1], device=phases.device)
+ )], dim=-1)
+ q_embed = self._rotary_embedding(q, phases)
+ k_embed = self._rotary_embedding(k, phases)
+ return q_embed, k_embed
+
+
+class SparseMultiHeadRMSNorm(nn.Module):
+ def __init__(self, dim: int, heads: int):
+ super().__init__()
+ self.scale = dim ** 0.5
+ self.gamma = nn.Parameter(torch.ones(heads, dim))
+
+ def forward(self, x: Union[SparseTensor, torch.Tensor]) -> Union[SparseTensor, torch.Tensor]:
+ x_type = x.dtype
+ x = x.float()
+ if isinstance(x, SparseTensor):
+ x = x.replace(F.normalize(x.feats, dim=-1))
+ else:
+ x = F.normalize(x, dim=-1)
+ return (x * self.gamma * self.scale).to(x_type)
+
+
+class SparseMultiHeadAttention(nn.Module):
+ def __init__(
+ self,
+ channels: int,
+ num_heads: int,
+ ctx_channels: Optional[int] = None,
+ type: Literal["self", "cross"] = "self",
+ attn_mode: Literal["full", "serialized", "windowed"] = "full",
+ window_size: Optional[int] = None,
+ shift_sequence: Optional[int] = None,
+ shift_window: Optional[Tuple[int, int, int]] = None,
+ serialize_mode: Optional[SerializeMode] = None,
+ qkv_bias: bool = True,
+ use_rope: bool = False,
+ qk_rms_norm: bool = False,
+ ):
+ super().__init__()
+ assert channels % num_heads == 0
+ assert type in ["self", "cross"], f"Invalid attention type: {type}"
+ assert attn_mode in ["full", "serialized", "windowed"], f"Invalid attention mode: {attn_mode}"
+ assert type == "self" or attn_mode == "full", "Cross-attention only supports full attention"
+ assert type == "self" or use_rope is False, "Rotary position embeddings only supported for self-attention"
+ self.channels = channels
+ self.ctx_channels = ctx_channels if ctx_channels is not None else channels
+ self.num_heads = num_heads
+ self._type = type
+ self.attn_mode = attn_mode
+ self.window_size = window_size
+ self.shift_sequence = shift_sequence
+ self.shift_window = shift_window
+ self.serialize_mode = serialize_mode
+ self.use_rope = use_rope
+ self.qk_rms_norm = qk_rms_norm
+
+ if self._type == "self":
+ self.to_qkv = nn.Linear(channels, channels * 3, bias=qkv_bias)
+ else:
+ self.to_q = nn.Linear(channels, channels, bias=qkv_bias)
+ self.to_kv = nn.Linear(self.ctx_channels, channels * 2, bias=qkv_bias)
+
+ if self.qk_rms_norm:
+ self.q_rms_norm = SparseMultiHeadRMSNorm(channels // num_heads, num_heads)
+ self.k_rms_norm = SparseMultiHeadRMSNorm(channels // num_heads, num_heads)
+
+ self.to_out = nn.Linear(channels, channels)
+
+ if use_rope:
+ # self.rope = RotaryPositionEmbedder(channels)
+
+ head_dim = channels // self.num_heads
+ self.rope = RotaryPositionEmbedder(head_dim)
+
+
+ @staticmethod
+ def _linear(module: nn.Linear, x: Union[SparseTensor, torch.Tensor]) -> Union[SparseTensor, torch.Tensor]:
+ if isinstance(x, SparseTensor):
+ return x.replace(module(x.feats))
+ else:
+ return module(x)
+
+ @staticmethod
+ def _reshape_chs(x: Union[SparseTensor, torch.Tensor], shape: Tuple[int, ...]) -> Union[SparseTensor, torch.Tensor]:
+ if isinstance(x, SparseTensor):
+ return x.reshape(*shape)
+ else:
+ return x.reshape(*x.shape[:2], *shape)
+
+ def _fused_pre(self, x: Union[SparseTensor, torch.Tensor], num_fused: int) -> Union[SparseTensor, torch.Tensor]:
+ if isinstance(x, SparseTensor):
+ x_feats = x.feats.unsqueeze(0)
+ else:
+ x_feats = x
+ x_feats = x_feats.reshape(*x_feats.shape[:2], num_fused, self.num_heads, -1)
+ return x.replace(x_feats.squeeze(0)) if isinstance(x, SparseTensor) else x_feats
+
+ def _rope(self, qkv: SparseTensor) -> SparseTensor:
+ q, k, v = qkv.feats.unbind(dim=1) # [T, H, C]
+ q, k = self.rope(q, k, qkv.coords[:, 1:])
+ qkv = qkv.replace(torch.stack([q, k, v], dim=1))
+ return qkv
+
+ def forward(self, x: Union[SparseTensor, torch.Tensor], context: Optional[Union[SparseTensor, torch.Tensor]] = None) -> Union[SparseTensor, torch.Tensor]:
+ if self._type == "self": # self-attn, default
+ qkv = self._linear(self.to_qkv, x)
+ qkv = self._fused_pre(qkv, num_fused=3) # to reshape
+ if self.use_rope: # False, default
+ qkv = self._rope(qkv)
+ if self.qk_rms_norm:
+ q, k, v = qkv.unbind(dim=1)
+ q = self.q_rms_norm(q)
+ k = self.k_rms_norm(k)
+ qkv = qkv.replace(torch.stack([q.feats, k.feats, v.feats], dim=1))
+ if self.attn_mode == "full":
+ h = sparse_scaled_dot_product_attention(qkv)
+ elif self.attn_mode == "serialized":
+ h = sparse_serialized_scaled_dot_product_self_attention(
+ qkv, self.window_size, serialize_mode=self.serialize_mode, shift_sequence=self.shift_sequence, shift_window=self.shift_window
+ )
+ elif self.attn_mode == "windowed":
+ h = sparse_windowed_scaled_dot_product_self_attention(
+ qkv, self.window_size, shift_window=self.shift_window
+ )
+ else: # cross attn, default False
+ q = self._linear(self.to_q, x)
+ q = self._reshape_chs(q, (self.num_heads, -1))
+ kv = self._linear(self.to_kv, context)
+ kv = self._fused_pre(kv, num_fused=2)
+ if self.qk_rms_norm:
+ q = self.q_rms_norm(q)
+ k, v = kv.unbind(dim=1)
+ k = self.k_rms_norm(k)
+ kv = kv.replace(torch.stack([k.feats, v.feats], dim=1))
+ h = sparse_scaled_dot_product_attention(q, kv)
+ h = self._reshape_chs(h, (-1,))
+ h = self._linear(self.to_out, h)
+ return h
diff --git a/modules/sparse/attention/serialized_attn.py b/modules/sparse/attention/serialized_attn.py
new file mode 100644
index 0000000000000000000000000000000000000000..d97987464a4efaced20401bc66b4ccf13251263e
--- /dev/null
+++ b/modules/sparse/attention/serialized_attn.py
@@ -0,0 +1,217 @@
+# MIT License
+
+# Copyright (c) Microsoft Corporation.
+# Copyright (c) 2025 VAST-AI-Research and contributors.
+
+# Permission is hereby granted, free of charge, to any person obtaining a copy
+# of this software and associated documentation files (the "Software"), to deal
+# in the Software without restriction, including without limitation the rights
+# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
+# copies of the Software, and to permit persons to whom the Software is
+# furnished to do so, subject to the following conditions:
+
+# The above copyright notice and this permission notice shall be included in all
+# copies or substantial portions of the Software.
+
+# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
+# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
+# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
+# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
+# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
+# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
+# SOFTWARE
+
+from typing import *
+from enum import Enum
+import torch
+import math
+from .. import SparseTensor
+from .. import DEBUG, ATTN
+
+if ATTN == 'xformers':
+ import xformers.ops as xops
+elif ATTN == 'flash_attn':
+ import flash_attn
+else:
+ raise ValueError(f"Unknown attention module: {ATTN}")
+
+
+__all__ = [
+ 'sparse_serialized_scaled_dot_product_self_attention',
+ 'SerializeModes',
+]
+
+
+class SerializeMode(Enum):
+ Z_ORDER = 0
+ Z_ORDER_TRANSPOSED = 1
+ HILBERT = 2
+ HILBERT_TRANSPOSED = 3
+
+
+SerializeModes = [
+ SerializeMode.Z_ORDER,
+ SerializeMode.Z_ORDER_TRANSPOSED,
+ SerializeMode.HILBERT,
+ SerializeMode.HILBERT_TRANSPOSED
+]
+
+
+def calc_serialization(
+ tensor: SparseTensor,
+ window_size: int,
+ serialize_mode: SerializeMode = SerializeMode.Z_ORDER,
+ shift_sequence: int = 0,
+ shift_window: Tuple[int, int, int] = (0, 0, 0)
+) -> Tuple[torch.Tensor, torch.Tensor, List[int]]:
+ """
+ Calculate serialization and partitioning for a set of coordinates.
+
+ Args:
+ tensor (SparseTensor): The input tensor.
+ window_size (int): The window size to use.
+ serialize_mode (SerializeMode): The serialization mode to use.
+ shift_sequence (int): The shift of serialized sequence.
+ shift_window (Tuple[int, int, int]): The shift of serialized coordinates.
+
+ Returns:
+ (torch.Tensor, torch.Tensor): Forwards and backwards indices.
+ """
+ fwd_indices = []
+ bwd_indices = []
+ seq_lens = []
+ seq_batch_indices = []
+ offsets = [0]
+
+ if 'vox2seq' not in globals():
+ import vox2seq
+
+ # Serialize the input
+ serialize_coords = tensor.coords[:, 1:].clone()
+ serialize_coords += torch.tensor(shift_window, dtype=torch.int32, device=tensor.device).reshape(1, 3)
+ if serialize_mode == SerializeMode.Z_ORDER:
+ code = vox2seq.encode(serialize_coords, mode='z_order', permute=[0, 1, 2])
+ elif serialize_mode == SerializeMode.Z_ORDER_TRANSPOSED:
+ code = vox2seq.encode(serialize_coords, mode='z_order', permute=[1, 0, 2])
+ elif serialize_mode == SerializeMode.HILBERT:
+ code = vox2seq.encode(serialize_coords, mode='hilbert', permute=[0, 1, 2])
+ elif serialize_mode == SerializeMode.HILBERT_TRANSPOSED:
+ code = vox2seq.encode(serialize_coords, mode='hilbert', permute=[1, 0, 2])
+ else:
+ raise ValueError(f"Unknown serialize mode: {serialize_mode}")
+
+ for bi, s in enumerate(tensor.layout):
+ num_points = s.stop - s.start
+ num_windows = (num_points + window_size - 1) // window_size
+ valid_window_size = num_points / num_windows
+ to_ordered = torch.argsort(code[s.start:s.stop])
+ if num_windows == 1:
+ fwd_indices.append(to_ordered)
+ bwd_indices.append(torch.zeros_like(to_ordered).scatter_(0, to_ordered, torch.arange(num_points, device=tensor.device)))
+ fwd_indices[-1] += s.start
+ bwd_indices[-1] += offsets[-1]
+ seq_lens.append(num_points)
+ seq_batch_indices.append(bi)
+ offsets.append(offsets[-1] + seq_lens[-1])
+ else:
+ # Partition the input
+ offset = 0
+ mids = [(i + 0.5) * valid_window_size + shift_sequence for i in range(num_windows)]
+ split = [math.floor(i * valid_window_size + shift_sequence) for i in range(num_windows + 1)]
+ bwd_index = torch.zeros((num_points,), dtype=torch.int64, device=tensor.device)
+ for i in range(num_windows):
+ mid = mids[i]
+ valid_start = split[i]
+ valid_end = split[i + 1]
+ padded_start = math.floor(mid - 0.5 * window_size)
+ padded_end = padded_start + window_size
+ fwd_indices.append(to_ordered[torch.arange(padded_start, padded_end, device=tensor.device) % num_points])
+ offset += valid_start - padded_start
+ 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))
+ offset += padded_end - valid_start
+ fwd_indices[-1] += s.start
+ seq_lens.extend([window_size] * num_windows)
+ seq_batch_indices.extend([bi] * num_windows)
+ bwd_indices.append(bwd_index + offsets[-1])
+ offsets.append(offsets[-1] + num_windows * window_size)
+
+ fwd_indices = torch.cat(fwd_indices)
+ bwd_indices = torch.cat(bwd_indices)
+
+ return fwd_indices, bwd_indices, seq_lens, seq_batch_indices
+
+
+def sparse_serialized_scaled_dot_product_self_attention(
+ qkv: SparseTensor,
+ window_size: int,
+ serialize_mode: SerializeMode = SerializeMode.Z_ORDER,
+ shift_sequence: int = 0,
+ shift_window: Tuple[int, int, int] = (0, 0, 0)
+) -> SparseTensor:
+ """
+ Apply serialized scaled dot product self attention to a sparse tensor.
+
+ Args:
+ qkv (SparseTensor): [N, *, 3, H, C] sparse tensor containing Qs, Ks, and Vs.
+ window_size (int): The window size to use.
+ serialize_mode (SerializeMode): The serialization mode to use.
+ shift_sequence (int): The shift of serialized sequence.
+ shift_window (Tuple[int, int, int]): The shift of serialized coordinates.
+ shift (int): The shift to use.
+ """
+ assert len(qkv.shape) == 4 and qkv.shape[1] == 3, f"Invalid shape for qkv, got {qkv.shape}, expected [N, *, 3, H, C]"
+
+ serialization_spatial_cache_name = f'serialization_{serialize_mode}_{window_size}_{shift_sequence}_{shift_window}'
+ serialization_spatial_cache = qkv.get_spatial_cache(serialization_spatial_cache_name)
+ if serialization_spatial_cache is None:
+ fwd_indices, bwd_indices, seq_lens, seq_batch_indices = calc_serialization(qkv, window_size, serialize_mode, shift_sequence, shift_window)
+ qkv.register_spatial_cache(serialization_spatial_cache_name, (fwd_indices, bwd_indices, seq_lens, seq_batch_indices))
+ else:
+ fwd_indices, bwd_indices, seq_lens, seq_batch_indices = serialization_spatial_cache
+
+ M = fwd_indices.shape[0]
+ T = qkv.feats.shape[0]
+ H = qkv.feats.shape[2]
+ C = qkv.feats.shape[3]
+
+ qkv_feats = qkv.feats[fwd_indices] # [M, 3, H, C]
+
+ if DEBUG:
+ start = 0
+ qkv_coords = qkv.coords[fwd_indices]
+ for i in range(len(seq_lens)):
+ assert (qkv_coords[start:start+seq_lens[i], 0] == seq_batch_indices[i]).all(), f"SparseWindowedScaledDotProductSelfAttention: batch index mismatch"
+ start += seq_lens[i]
+
+ if all([seq_len == window_size for seq_len in seq_lens]):
+ B = len(seq_lens)
+ N = window_size
+ qkv_feats = qkv_feats.reshape(B, N, 3, H, C)
+ if ATTN == 'xformers':
+ q, k, v = qkv_feats.unbind(dim=2) # [B, N, H, C]
+ out = xops.memory_efficient_attention(q, k, v) # [B, N, H, C]
+ elif ATTN == 'flash_attn':
+ out = flash_attn.flash_attn_qkvpacked_func(qkv_feats) # [B, N, H, C]
+ else:
+ raise ValueError(f"Unknown attention module: {ATTN}")
+ out = out.reshape(B * N, H, C) # [M, H, C]
+ else:
+ if ATTN == 'xformers':
+ q, k, v = qkv_feats.unbind(dim=1) # [M, H, C]
+ q = q.unsqueeze(0) # [1, M, H, C]
+ k = k.unsqueeze(0) # [1, M, H, C]
+ v = v.unsqueeze(0) # [1, M, H, C]
+ mask = xops.fmha.BlockDiagonalMask.from_seqlens(seq_lens)
+ out = xops.memory_efficient_attention(q, k, v, mask)[0] # [M, H, C]
+ elif ATTN == 'flash_attn':
+ cu_seqlens = torch.cat([torch.tensor([0]), torch.cumsum(torch.tensor(seq_lens), dim=0)], dim=0) \
+ .to(qkv.device).int()
+ out = flash_attn.flash_attn_varlen_qkvpacked_func(qkv_feats, cu_seqlens, max(seq_lens)) # [M, H, C]
+
+ out = out[bwd_indices] # [T, H, C]
+
+ if DEBUG:
+ qkv_coords = qkv_coords[bwd_indices]
+ assert torch.equal(qkv_coords, qkv.coords), "SparseWindowedScaledDotProductSelfAttention: coordinate mismatch"
+
+ return qkv.replace(out)
diff --git a/modules/sparse/attention/windowed_attn.py b/modules/sparse/attention/windowed_attn.py
new file mode 100644
index 0000000000000000000000000000000000000000..ff33c7da42fc8272a2c4b21036c442a84751b025
--- /dev/null
+++ b/modules/sparse/attention/windowed_attn.py
@@ -0,0 +1,158 @@
+# MIT License
+
+# Copyright (c) Microsoft Corporation.
+# Copyright (c) 2025 VAST-AI-Research and contributors.
+
+# Permission is hereby granted, free of charge, to any person obtaining a copy
+# of this software and associated documentation files (the "Software"), to deal
+# in the Software without restriction, including without limitation the rights
+# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
+# copies of the Software, and to permit persons to whom the Software is
+# furnished to do so, subject to the following conditions:
+
+# The above copyright notice and this permission notice shall be included in all
+# copies or substantial portions of the Software.
+
+# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
+# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
+# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
+# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
+# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
+# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
+# SOFTWARE
+
+from typing import *
+import torch
+import math
+from .. import SparseTensor
+from .. import DEBUG, ATTN
+
+if ATTN == 'xformers':
+ import xformers.ops as xops
+elif ATTN == 'flash_attn':
+ import flash_attn
+else:
+ raise ValueError(f"Unknown attention module: {ATTN}")
+
+
+__all__ = [
+ 'sparse_windowed_scaled_dot_product_self_attention',
+]
+
+
+def calc_window_partition(
+ tensor: SparseTensor,
+ window_size: Union[int, Tuple[int, ...]],
+ shift_window: Union[int, Tuple[int, ...]] = 0
+) -> Tuple[torch.Tensor, torch.Tensor, List[int], List[int]]:
+ """
+ Calculate serialization and partitioning for a set of coordinates.
+
+ Args:
+ tensor (SparseTensor): The input tensor.
+ window_size (int): The window size to use.
+ shift_window (Tuple[int, ...]): The shift of serialized coordinates.
+
+ Returns:
+ (torch.Tensor): Forwards indices.
+ (torch.Tensor): Backwards indices.
+ (List[int]): Sequence lengths.
+ (List[int]): Sequence batch indices.
+ """
+ DIM = tensor.coords.shape[1] - 1
+ shift_window = (shift_window,) * DIM if isinstance(shift_window, int) else shift_window
+ window_size = (window_size,) * DIM if isinstance(window_size, int) else window_size
+ shifted_coords = tensor.coords.clone().detach()
+ shifted_coords[:, 1:] += torch.tensor(shift_window, device=tensor.device, dtype=torch.int32).unsqueeze(0)
+
+ MAX_COORDS = shifted_coords[:, 1:].max(dim=0).values.tolist()
+ NUM_WINDOWS = [math.ceil((mc + 1) / ws) for mc, ws in zip(MAX_COORDS, window_size)]
+ OFFSET = torch.cumprod(torch.tensor([1] + NUM_WINDOWS[::-1]), dim=0).tolist()[::-1]
+
+ shifted_coords[:, 1:] //= torch.tensor(window_size, device=tensor.device, dtype=torch.int32).unsqueeze(0)
+ shifted_indices = (shifted_coords * torch.tensor(OFFSET, device=tensor.device, dtype=torch.int32).unsqueeze(0)).sum(dim=1)
+ fwd_indices = torch.argsort(shifted_indices)
+ bwd_indices = torch.empty_like(fwd_indices)
+ bwd_indices[fwd_indices] = torch.arange(fwd_indices.shape[0], device=tensor.device)
+ seq_lens = torch.bincount(shifted_indices)
+ seq_batch_indices = torch.arange(seq_lens.shape[0], device=tensor.device, dtype=torch.int32) // OFFSET[0]
+ mask = seq_lens != 0
+ seq_lens = seq_lens[mask].tolist()
+ seq_batch_indices = seq_batch_indices[mask].tolist()
+
+ return fwd_indices, bwd_indices, seq_lens, seq_batch_indices
+
+
+def sparse_windowed_scaled_dot_product_self_attention(
+ qkv: SparseTensor,
+ window_size: int,
+ shift_window: Tuple[int, int, int] = (0, 0, 0)
+) -> SparseTensor:
+ """
+ Apply windowed scaled dot product self attention to a sparse tensor.
+
+ Args:
+ qkv (SparseTensor): [N, *, 3, H, C] sparse tensor containing Qs, Ks, and Vs.
+ window_size (int): The window size to use.
+ shift_window (Tuple[int, int, int]): The shift of serialized coordinates.
+ shift (int): The shift to use.
+ """
+ assert len(qkv.shape) == 4 and qkv.shape[1] == 3, f"Invalid shape for qkv, got {qkv.shape}, expected [N, *, 3, H, C]"
+
+ serialization_spatial_cache_name = f'window_partition_{window_size}_{shift_window}_{qkv.feats.shape[0]}'
+ serialization_spatial_cache = qkv.get_spatial_cache(serialization_spatial_cache_name)
+ if serialization_spatial_cache is None:
+ fwd_indices, bwd_indices, seq_lens, seq_batch_indices = calc_window_partition(qkv, window_size, shift_window)
+ qkv.register_spatial_cache(serialization_spatial_cache_name, (fwd_indices, bwd_indices, seq_lens, seq_batch_indices))
+ else:
+ fwd_indices, bwd_indices, seq_lens, seq_batch_indices = serialization_spatial_cache
+
+ M = fwd_indices.shape[0]
+ T = qkv.feats.shape[0]
+ H = qkv.feats.shape[2]
+ C = qkv.feats.shape[3]
+
+ qkv_feats = qkv.feats[fwd_indices] # [M, 3, H, C]
+
+ if DEBUG:
+ start = 0
+ qkv_coords = qkv.coords[fwd_indices]
+ for i in range(len(seq_lens)):
+ seq_coords = qkv_coords[start:start+seq_lens[i]]
+ assert (seq_coords[:, 0] == seq_batch_indices[i]).all(), f"SparseWindowedScaledDotProductSelfAttention: batch index mismatch"
+ assert (seq_coords[:, 1:].max(dim=0).values - seq_coords[:, 1:].min(dim=0).values < window_size).all(), \
+ f"SparseWindowedScaledDotProductSelfAttention: window size exceeded"
+ start += seq_lens[i]
+
+ if all([seq_len == window_size for seq_len in seq_lens]):
+ B = len(seq_lens)
+ N = window_size
+ qkv_feats = qkv_feats.reshape(B, N, 3, H, C)
+ if ATTN == 'xformers':
+ q, k, v = qkv_feats.unbind(dim=2) # [B, N, H, C]
+ out = xops.memory_efficient_attention(q, k, v) # [B, N, H, C]
+ elif ATTN == 'flash_attn':
+ out = flash_attn.flash_attn_qkvpacked_func(qkv_feats) # [B, N, H, C]
+ else:
+ raise ValueError(f"Unknown attention module: {ATTN}")
+ out = out.reshape(B * N, H, C) # [M, H, C]
+ else:
+ if ATTN == 'xformers':
+ q, k, v = qkv_feats.unbind(dim=1) # [M, H, C]
+ q = q.unsqueeze(0) # [1, M, H, C]
+ k = k.unsqueeze(0) # [1, M, H, C]
+ v = v.unsqueeze(0) # [1, M, H, C]
+ mask = xops.fmha.BlockDiagonalMask.from_seqlens(seq_lens)
+ out = xops.memory_efficient_attention(q, k, v, mask)[0] # [M, H, C]
+ elif ATTN == 'flash_attn':
+ cu_seqlens = torch.cat([torch.tensor([0]), torch.cumsum(torch.tensor(seq_lens), dim=0)], dim=0) \
+ .to(qkv.device).int()
+ out = flash_attn.flash_attn_varlen_qkvpacked_func(qkv_feats, cu_seqlens, max(seq_lens)) # [M, H, C]
+
+ out = out[bwd_indices] # [T, H, C]
+
+ if DEBUG:
+ qkv_coords = qkv_coords[bwd_indices]
+ assert torch.equal(qkv_coords, qkv.coords), "SparseWindowedScaledDotProductSelfAttention: coordinate mismatch"
+
+ return qkv.replace(out)
diff --git a/modules/sparse/basic.py b/modules/sparse/basic.py
new file mode 100644
index 0000000000000000000000000000000000000000..a9a410e8820b1350b6a48694c790460023efca09
--- /dev/null
+++ b/modules/sparse/basic.py
@@ -0,0 +1,482 @@
+# MIT License
+
+# Copyright (c) Microsoft Corporation.
+# Copyright (c) 2025 VAST-AI-Research and contributors.
+
+# Permission is hereby granted, free of charge, to any person obtaining a copy
+# of this software and associated documentation files (the "Software"), to deal
+# in the Software without restriction, including without limitation the rights
+# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
+# copies of the Software, and to permit persons to whom the Software is
+# furnished to do so, subject to the following conditions:
+
+# The above copyright notice and this permission notice shall be included in all
+# copies or substantial portions of the Software.
+
+# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
+# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
+# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
+# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
+# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
+# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
+# SOFTWARE
+
+from typing import *
+import torch
+import torch.nn as nn
+from . import BACKEND, DEBUG
+SparseTensorData = None # Lazy import
+
+
+__all__ = [
+ 'SparseTensor',
+ 'sparse_batch_broadcast',
+ 'sparse_batch_op',
+ 'sparse_cat',
+ 'sparse_unbind',
+]
+
+
+class SparseTensor:
+ """
+ Sparse tensor with support for both torchsparse and spconv backends.
+
+ Parameters:
+ - feats (torch.Tensor): Features of the sparse tensor.
+ - coords (torch.Tensor): Coordinates of the sparse tensor.
+ - shape (torch.Size): Shape of the sparse tensor.
+ - layout (List[slice]): Layout of the sparse tensor for each batch
+ - data (SparseTensorData): Sparse tensor data used for convolusion
+
+ NOTE:
+ - Data corresponding to a same batch should be contiguous.
+ - Coords should be in [0, 1023]
+ """
+ @overload
+ def __init__(self, feats: torch.Tensor, coords: torch.Tensor, shape: Optional[torch.Size] = None, layout: Optional[List[slice]] = None, **kwargs): ...
+
+ @overload
+ def __init__(self, data, shape: Optional[torch.Size] = None, layout: Optional[List[slice]] = None, **kwargs): ...
+
+ def __init__(self, *args, **kwargs):
+ # Lazy import of sparse tensor backend
+ global SparseTensorData
+ if SparseTensorData is None:
+ import importlib
+ if BACKEND == 'torchsparse':
+ SparseTensorData = importlib.import_module('torchsparse').SparseTensor
+ elif BACKEND == 'spconv':
+ SparseTensorData = importlib.import_module('spconv.pytorch').SparseConvTensor
+
+ method_id = 0
+ if len(args) != 0:
+ method_id = 0 if isinstance(args[0], torch.Tensor) else 1
+ else:
+ method_id = 1 if 'data' in kwargs else 0
+
+ if method_id == 0:
+ feats, coords, shape, layout = args + (None,) * (4 - len(args))
+ if 'feats' in kwargs:
+ feats = kwargs['feats']
+ del kwargs['feats']
+ if 'coords' in kwargs:
+ coords = kwargs['coords']
+ del kwargs['coords']
+ if 'shape' in kwargs:
+ shape = kwargs['shape']
+ del kwargs['shape']
+ if 'layout' in kwargs:
+ layout = kwargs['layout']
+ del kwargs['layout']
+
+ if shape is None:
+ shape = self.__cal_shape(feats, coords)
+ if layout is None:
+ layout = self.__cal_layout(coords, shape[0])
+ if BACKEND == 'torchsparse':
+ self.data = SparseTensorData(feats, coords, **kwargs)
+ elif BACKEND == 'spconv':
+ spatial_shape = list(coords.max(0)[0] + 1)[1:]
+ self.data = SparseTensorData(feats.reshape(feats.shape[0], -1), coords, spatial_shape, shape[0], **kwargs)
+ self.data._features = feats
+ elif method_id == 1:
+ data, shape, layout = args + (None,) * (3 - len(args))
+ if 'data' in kwargs:
+ data = kwargs['data']
+ del kwargs['data']
+ if 'shape' in kwargs:
+ shape = kwargs['shape']
+ del kwargs['shape']
+ if 'layout' in kwargs:
+ layout = kwargs['layout']
+ del kwargs['layout']
+
+ self.data = data
+ if shape is None:
+ shape = self.__cal_shape(self.feats, self.coords)
+ if layout is None:
+ layout = self.__cal_layout(self.coords, shape[0])
+
+ self._shape = shape
+ self._layout = layout
+ self._scale = kwargs.get('scale', (1, 1, 1))
+ self._spatial_cache = kwargs.get('spatial_cache', {})
+
+ if DEBUG:
+ try:
+ assert self.feats.shape[0] == self.coords.shape[0], f"Invalid feats shape: {self.feats.shape}, coords shape: {self.coords.shape}"
+ assert self.shape == self.__cal_shape(self.feats, self.coords), f"Invalid shape: {self.shape}"
+ assert self.layout == self.__cal_layout(self.coords, self.shape[0]), f"Invalid layout: {self.layout}"
+ for i in range(self.shape[0]):
+ assert torch.all(self.coords[self.layout[i], 0] == i), f"The data of batch {i} is not contiguous"
+ except Exception as e:
+ print('Debugging information:')
+ print(f"- Shape: {self.shape}")
+ print(f"- Layout: {self.layout}")
+ print(f"- Scale: {self._scale}")
+ print(f"- Coords: {self.coords}")
+ raise e
+
+ def __cal_shape(self, feats, coords):
+ shape = []
+ shape.append(coords[:, 0].max().item() + 1)
+ shape.extend([*feats.shape[1:]])
+ return torch.Size(shape)
+
+ def __cal_layout(self, coords, batch_size):
+ seq_len = torch.bincount(coords[:, 0], minlength=batch_size)
+ offset = torch.cumsum(seq_len, dim=0)
+ layout = [slice((offset[i] - seq_len[i]).item(), offset[i].item()) for i in range(batch_size)]
+ return layout
+
+ @property
+ def shape(self) -> torch.Size:
+ return self._shape
+
+ def dim(self) -> int:
+ return len(self.shape)
+
+ @property
+ def layout(self) -> List[slice]:
+ return self._layout
+
+ @property
+ def feats(self) -> torch.Tensor:
+ if BACKEND == 'torchsparse':
+ return self.data.F
+ elif BACKEND == 'spconv':
+ return self.data.features
+
+ @feats.setter
+ def feats(self, value: torch.Tensor):
+ if BACKEND == 'torchsparse':
+ self.data.F = value
+ elif BACKEND == 'spconv':
+ self.data.features = value
+
+ @property
+ def coords(self) -> torch.Tensor:
+ if BACKEND == 'torchsparse':
+ return self.data.C
+ elif BACKEND == 'spconv':
+ return self.data.indices
+
+ @coords.setter
+ def coords(self, value: torch.Tensor):
+ if BACKEND == 'torchsparse':
+ self.data.C = value
+ elif BACKEND == 'spconv':
+ self.data.indices = value
+
+ @property
+ def dtype(self):
+ return self.feats.dtype
+
+ @property
+ def device(self):
+ return self.feats.device
+
+ @overload
+ def to(self, dtype: torch.dtype) -> 'SparseTensor': ...
+
+ @overload
+ def to(self, device: Optional[Union[str, torch.device]] = None, dtype: Optional[torch.dtype] = None) -> 'SparseTensor': ...
+
+ def to(self, *args, **kwargs) -> 'SparseTensor':
+ device = None
+ dtype = None
+ if len(args) == 2:
+ device, dtype = args
+ elif len(args) == 1:
+ if isinstance(args[0], torch.dtype):
+ dtype = args[0]
+ else:
+ device = args[0]
+ if 'dtype' in kwargs:
+ assert dtype is None, "to() received multiple values for argument 'dtype'"
+ dtype = kwargs['dtype']
+ if 'device' in kwargs:
+ assert device is None, "to() received multiple values for argument 'device'"
+ device = kwargs['device']
+
+ new_feats = self.feats.to(device=device, dtype=dtype)
+ new_coords = self.coords.to(device=device)
+ return self.replace(new_feats, new_coords)
+
+ def type(self, dtype):
+ new_feats = self.feats.type(dtype)
+ return self.replace(new_feats)
+
+ def cpu(self) -> 'SparseTensor':
+ new_feats = self.feats.cpu()
+ new_coords = self.coords.cpu()
+ return self.replace(new_feats, new_coords)
+
+ def cuda(self) -> 'SparseTensor':
+ new_feats = self.feats.cuda()
+ new_coords = self.coords.cuda()
+ return self.replace(new_feats, new_coords)
+
+ def half(self) -> 'SparseTensor':
+ new_feats = self.feats.half()
+ return self.replace(new_feats)
+
+ def float(self) -> 'SparseTensor':
+ new_feats = self.feats.float()
+ return self.replace(new_feats)
+
+ def detach(self) -> 'SparseTensor':
+ new_coords = self.coords.detach()
+ new_feats = self.feats.detach()
+ return self.replace(new_feats, new_coords)
+
+ def dense(self) -> torch.Tensor:
+ if BACKEND == 'torchsparse':
+ return self.data.dense()
+ elif BACKEND == 'spconv':
+ return self.data.dense()
+
+ def reshape(self, *shape) -> 'SparseTensor':
+ new_feats = self.feats.reshape(self.feats.shape[0], *shape)
+ return self.replace(new_feats)
+
+ def unbind(self, dim: int) -> List['SparseTensor']:
+ return sparse_unbind(self, dim)
+
+ def replace(self, feats: torch.Tensor, coords: Optional[torch.Tensor] = None) -> 'SparseTensor':
+ new_shape = [self.shape[0]]
+ new_shape.extend(feats.shape[1:])
+ if BACKEND == 'torchsparse':
+ new_data = SparseTensorData(
+ feats=feats,
+ coords=self.data.coords if coords is None else coords,
+ stride=self.data.stride,
+ spatial_range=self.data.spatial_range,
+ )
+ new_data._caches = self.data._caches
+ elif BACKEND == 'spconv':
+ new_data = SparseTensorData(
+ self.data.features.reshape(self.data.features.shape[0], -1),
+ self.data.indices,
+ self.data.spatial_shape,
+ self.data.batch_size,
+ self.data.grid,
+ self.data.voxel_num,
+ self.data.indice_dict
+ )
+ new_data._features = feats
+ new_data.benchmark = self.data.benchmark
+ new_data.benchmark_record = self.data.benchmark_record
+ new_data.thrust_allocator = self.data.thrust_allocator
+ new_data._timer = self.data._timer
+ new_data.force_algo = self.data.force_algo
+ new_data.int8_scale = self.data.int8_scale
+ if coords is not None:
+ new_data.indices = coords
+ new_tensor = SparseTensor(new_data, shape=torch.Size(new_shape), layout=self.layout, scale=self._scale, spatial_cache=self._spatial_cache)
+ return new_tensor
+
+ @staticmethod
+ def full(aabb, dim, value, dtype=torch.float32, device=None) -> 'SparseTensor':
+ N, C = dim
+ x = torch.arange(aabb[0], aabb[3] + 1)
+ y = torch.arange(aabb[1], aabb[4] + 1)
+ z = torch.arange(aabb[2], aabb[5] + 1)
+ coords = torch.stack(torch.meshgrid(x, y, z, indexing='ij'), dim=-1).reshape(-1, 3)
+ coords = torch.cat([
+ torch.arange(N).view(-1, 1).repeat(1, coords.shape[0]).view(-1, 1),
+ coords.repeat(N, 1),
+ ], dim=1).to(dtype=torch.int32, device=device)
+ feats = torch.full((coords.shape[0], C), value, dtype=dtype, device=device)
+ return SparseTensor(feats=feats, coords=coords)
+
+ def __merge_sparse_cache(self, other: 'SparseTensor') -> dict:
+ new_cache = {}
+ for k in set(list(self._spatial_cache.keys()) + list(other._spatial_cache.keys())):
+ if k in self._spatial_cache:
+ new_cache[k] = self._spatial_cache[k]
+ if k in other._spatial_cache:
+ if k not in new_cache:
+ new_cache[k] = other._spatial_cache[k]
+ else:
+ new_cache[k].update(other._spatial_cache[k])
+ return new_cache
+
+ def __neg__(self) -> 'SparseTensor':
+ return self.replace(-self.feats)
+
+ def __elemwise__(self, other: Union[torch.Tensor, 'SparseTensor'], op: callable) -> 'SparseTensor':
+ if isinstance(other, torch.Tensor):
+ try:
+ other = torch.broadcast_to(other, self.shape)
+ other = sparse_batch_broadcast(self, other)
+ except:
+ pass
+ if isinstance(other, SparseTensor):
+ other = other.feats
+ new_feats = op(self.feats, other)
+ new_tensor = self.replace(new_feats)
+ if isinstance(other, SparseTensor):
+ new_tensor._spatial_cache = self.__merge_sparse_cache(other)
+ return new_tensor
+
+ def __add__(self, other: Union[torch.Tensor, 'SparseTensor', float]) -> 'SparseTensor':
+ return self.__elemwise__(other, torch.add)
+
+ def __radd__(self, other: Union[torch.Tensor, 'SparseTensor', float]) -> 'SparseTensor':
+ return self.__elemwise__(other, torch.add)
+
+ def __sub__(self, other: Union[torch.Tensor, 'SparseTensor', float]) -> 'SparseTensor':
+ return self.__elemwise__(other, torch.sub)
+
+ def __rsub__(self, other: Union[torch.Tensor, 'SparseTensor', float]) -> 'SparseTensor':
+ return self.__elemwise__(other, lambda x, y: torch.sub(y, x))
+
+ def __mul__(self, other: Union[torch.Tensor, 'SparseTensor', float]) -> 'SparseTensor':
+ return self.__elemwise__(other, torch.mul)
+
+ def __rmul__(self, other: Union[torch.Tensor, 'SparseTensor', float]) -> 'SparseTensor':
+ return self.__elemwise__(other, torch.mul)
+
+ def __truediv__(self, other: Union[torch.Tensor, 'SparseTensor', float]) -> 'SparseTensor':
+ return self.__elemwise__(other, torch.div)
+
+ def __rtruediv__(self, other: Union[torch.Tensor, 'SparseTensor', float]) -> 'SparseTensor':
+ return self.__elemwise__(other, lambda x, y: torch.div(y, x))
+
+ def __getitem__(self, idx):
+ if isinstance(idx, int):
+ idx = [idx]
+ elif isinstance(idx, slice):
+ idx = range(*idx.indices(self.shape[0]))
+ elif isinstance(idx, torch.Tensor):
+ if idx.dtype == torch.bool:
+ assert idx.shape == (self.shape[0],), f"Invalid index shape: {idx.shape}"
+ idx = idx.nonzero().squeeze(1)
+ elif idx.dtype in [torch.int32, torch.int64]:
+ assert len(idx.shape) == 1, f"Invalid index shape: {idx.shape}"
+ else:
+ raise ValueError(f"Unknown index type: {idx.dtype}")
+ else:
+ raise ValueError(f"Unknown index type: {type(idx)}")
+
+ coords = []
+ feats = []
+ for new_idx, old_idx in enumerate(idx):
+ coords.append(self.coords[self.layout[old_idx]].clone())
+ coords[-1][:, 0] = new_idx
+ feats.append(self.feats[self.layout[old_idx]])
+ coords = torch.cat(coords, dim=0).contiguous()
+ feats = torch.cat(feats, dim=0).contiguous()
+ return SparseTensor(feats=feats, coords=coords)
+
+ def register_spatial_cache(self, key, value) -> None:
+ """
+ Register a spatial cache.
+ The spatial cache can be any thing you want to cache.
+ The registery and retrieval of the cache is based on current scale.
+ """
+ scale_key = str(self._scale)
+ if scale_key not in self._spatial_cache:
+ self._spatial_cache[scale_key] = {}
+ self._spatial_cache[scale_key][key] = value
+
+ def get_spatial_cache(self, key=None):
+ """
+ Get a spatial cache.
+ """
+ scale_key = str(self._scale)
+ cur_scale_cache = self._spatial_cache.get(scale_key, {})
+ if key is None:
+ return cur_scale_cache
+ return cur_scale_cache.get(key, None)
+
+
+def sparse_batch_broadcast(input: SparseTensor, other: torch.Tensor) -> torch.Tensor:
+ """
+ Broadcast a 1D tensor to a sparse tensor along the batch dimension then perform an operation.
+
+ Args:
+ input (torch.Tensor): 1D tensor to broadcast.
+ target (SparseTensor): Sparse tensor to broadcast to.
+ op (callable): Operation to perform after broadcasting. Defaults to torch.add.
+ """
+ coords, feats = input.coords, input.feats
+ broadcasted = torch.zeros_like(feats)
+ for k in range(input.shape[0]):
+ broadcasted[input.layout[k]] = other[k]
+ return broadcasted
+
+
+def sparse_batch_op(input: SparseTensor, other: torch.Tensor, op: callable = torch.add) -> SparseTensor:
+ """
+ Broadcast a 1D tensor to a sparse tensor along the batch dimension then perform an operation.
+
+ Args:
+ input (torch.Tensor): 1D tensor to broadcast.
+ target (SparseTensor): Sparse tensor to broadcast to.
+ op (callable): Operation to perform after broadcasting. Defaults to torch.add.
+ """
+ return input.replace(op(input.feats, sparse_batch_broadcast(input, other)))
+
+
+def sparse_cat(inputs: List[SparseTensor], dim: int = 0) -> SparseTensor:
+ """
+ Concatenate a list of sparse tensors.
+
+ Args:
+ inputs (List[SparseTensor]): List of sparse tensors to concatenate.
+ """
+ if dim == 0:
+ start = 0
+ coords = []
+ for input in inputs:
+ coords.append(input.coords.clone())
+ coords[-1][:, 0] += start
+ start += input.shape[0]
+ coords = torch.cat(coords, dim=0)
+ feats = torch.cat([input.feats for input in inputs], dim=0)
+ output = SparseTensor(
+ coords=coords,
+ feats=feats,
+ )
+ else:
+ feats = torch.cat([input.feats for input in inputs], dim=dim)
+ output = inputs[0].replace(feats)
+
+ return output
+
+
+def sparse_unbind(input: SparseTensor, dim: int) -> List[SparseTensor]:
+ """
+ Unbind a sparse tensor along a dimension.
+
+ Args:
+ input (SparseTensor): Sparse tensor to unbind.
+ dim (int): Dimension to unbind.
+ """
+ if dim == 0:
+ return [input[i] for i in range(input.shape[0])]
+ else:
+ feats = input.feats.unbind(dim)
+ return [input.replace(f) for f in feats]
diff --git a/modules/sparse/blocks.py b/modules/sparse/blocks.py
new file mode 100644
index 0000000000000000000000000000000000000000..6676e3245158e38cec275e97465a141b8d709653
--- /dev/null
+++ b/modules/sparse/blocks.py
@@ -0,0 +1,71 @@
+from typing import *
+import torch
+import torch.nn as nn
+import torch.nn.functional as F
+from ..utils import zero_module
+from ..norm import LayerNorm32
+from .. import sparse as sp
+
+
+class SparseResBlock3d(nn.Module):
+ def __init__(
+ self,
+ channels: int,
+ out_channels: Optional[int] = None,
+ downsample: bool = False,
+ upsample: bool = False,
+ use_checkpoint: bool = False,
+ ):
+ super().__init__()
+ self.channels = channels
+ self.out_channels = out_channels or channels
+ self.downsample = downsample
+ self.upsample = upsample
+ self.use_checkpoint = use_checkpoint
+
+ assert not (
+ downsample and upsample
+ ), "Cannot downsample and upsample at the same time"
+
+ self.norm1 = LayerNorm32(channels, elementwise_affine=True, eps=1e-6)
+ self.norm2 = LayerNorm32(self.out_channels, elementwise_affine=False, eps=1e-6)
+ self.conv1 = sp.SparseConv3d(channels, self.out_channels, 3)
+ self.conv2 = zero_module(
+ sp.SparseConv3d(self.out_channels, self.out_channels, 3)
+ )
+
+ self.skip_connection = (
+ sp.SparseLinear(channels, self.out_channels)
+ if channels != self.out_channels
+ else nn.Identity()
+ )
+ self.updown = None
+ if self.downsample:
+ self.updown = sp.SparseDownsample(2)
+ elif self.upsample:
+ self.updown = sp.SparseUpsample(2)
+
+ def _updown(self, x: sp.SparseTensor) -> sp.SparseTensor:
+ if self.updown is not None:
+ x = self.updown(x)
+ return x
+
+ def _forward(self, x: sp.SparseTensor) -> sp.SparseTensor:
+ x = self._updown(x)
+ h = x.replace(self.norm1(x.feats))
+ h = h.replace(F.silu(h.feats))
+ h = self.conv1(h)
+ h = h.replace(self.norm2(h.feats))
+ h = h.replace(F.silu(h.feats))
+ h = self.conv2(h)
+ h = h + self.skip_connection(x)
+
+ return h
+
+ def forward(self, x: torch.Tensor):
+ if self.use_checkpoint:
+ return torch.utils.checkpoint.checkpoint(
+ self._forward, x, use_reentrant=False
+ )
+ else:
+ return self._forward(x)
diff --git a/modules/sparse/conv/__init__.py b/modules/sparse/conv/__init__.py
new file mode 100644
index 0000000000000000000000000000000000000000..a0cff61228a79f856ffa13788ef876345757058b
--- /dev/null
+++ b/modules/sparse/conv/__init__.py
@@ -0,0 +1,44 @@
+# MIT License
+
+# Copyright (c) Microsoft Corporation.
+# Copyright (c) 2025 VAST-AI-Research and contributors.
+
+# Permission is hereby granted, free of charge, to any person obtaining a copy
+# of this software and associated documentation files (the "Software"), to deal
+# in the Software without restriction, including without limitation the rights
+# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
+# copies of the Software, and to permit persons to whom the Software is
+# furnished to do so, subject to the following conditions:
+
+# The above copyright notice and this permission notice shall be included in all
+# copies or substantial portions of the Software.
+
+# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
+# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
+# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
+# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
+# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
+# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
+# SOFTWARE
+
+from .. import BACKEND
+
+
+SPCONV_ALGO = 'auto' # 'auto', 'implicit_gemm', 'native'
+
+def __from_env():
+ import os
+
+ global SPCONV_ALGO
+ env_spconv_algo = os.environ.get('SPCONV_ALGO')
+ if env_spconv_algo is not None and env_spconv_algo in ['auto', 'implicit_gemm', 'native']:
+ SPCONV_ALGO = env_spconv_algo
+ print(f"[SPARSE][CONV] spconv algo: {SPCONV_ALGO}")
+
+
+__from_env()
+
+if BACKEND == 'torchsparse':
+ from .conv_torchsparse import *
+elif BACKEND == 'spconv':
+ from .conv_spconv import *
diff --git a/modules/sparse/conv/conv_spconv.py b/modules/sparse/conv/conv_spconv.py
new file mode 100644
index 0000000000000000000000000000000000000000..91dc48f5302b03160d7795e2033fcc5c11e43916
--- /dev/null
+++ b/modules/sparse/conv/conv_spconv.py
@@ -0,0 +1,107 @@
+# MIT License
+
+# Copyright (c) Microsoft Corporation.
+# Copyright (c) 2025 VAST-AI-Research and contributors.
+
+# Permission is hereby granted, free of charge, to any person obtaining a copy
+# of this software and associated documentation files (the "Software"), to deal
+# in the Software without restriction, including without limitation the rights
+# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
+# copies of the Software, and to permit persons to whom the Software is
+# furnished to do so, subject to the following conditions:
+
+# The above copyright notice and this permission notice shall be included in all
+# copies or substantial portions of the Software.
+
+# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
+# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
+# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
+# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
+# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
+# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
+# SOFTWARE
+
+import torch
+import torch.nn as nn
+from .. import SparseTensor
+from .. import DEBUG
+from . import SPCONV_ALGO
+import spconv.pytorch as spconv
+
+class SparseConv3d(nn.Module):
+ def __init__(self, in_channels, out_channels, kernel_size, stride=1, dilation=1, padding=None, bias=True, indice_key=None):
+ super(SparseConv3d, self).__init__()
+ # if 'spconv' not in globals():
+ # import spconv.pytorch as spconv
+ algo = None
+ if SPCONV_ALGO == 'native':
+ algo = spconv.ConvAlgo.Native
+ elif SPCONV_ALGO == 'implicit_gemm':
+ algo = spconv.ConvAlgo.MaskImplicitGemm
+ if stride == 1 and (padding is None):
+ self.conv = spconv.SubMConv3d(in_channels, out_channels, kernel_size, dilation=dilation, bias=bias, indice_key=indice_key, algo=algo)
+ else:
+ self.conv = spconv.SparseConv3d(in_channels, out_channels, kernel_size, stride=stride, dilation=dilation, padding=padding, bias=bias, indice_key=indice_key, algo=algo)
+ self.stride = tuple(stride) if isinstance(stride, (list, tuple)) else (stride, stride, stride)
+ self.padding = padding
+
+ def forward(self, x: SparseTensor) -> SparseTensor:
+ spatial_changed = any(s != 1 for s in self.stride) or (self.padding is not None)
+
+ dtype_ = x.feats.dtype
+ x = x.replace(x.feats.type(torch.float32))
+ new_data = self.conv(x.data)
+ new_shape = [x.shape[0], self.conv.out_channels]
+ new_layout = None if spatial_changed else x.layout
+
+ if spatial_changed and (x.shape[0] != 1):
+ # spconv was non-1 stride will break the contiguous of the output tensor, sort by the coords
+ fwd = new_data.indices[:, 0].argsort()
+ bwd = torch.zeros_like(fwd).scatter_(0, fwd, torch.arange(fwd.shape[0], device=fwd.device))
+ sorted_feats = new_data.features[fwd]
+ sorted_coords = new_data.indices[fwd]
+ unsorted_data = new_data
+ new_data = spconv.SparseConvTensor(sorted_feats, sorted_coords, unsorted_data.spatial_shape, unsorted_data.batch_size) # type: ignore
+
+ out = SparseTensor(
+ new_data, shape=torch.Size(new_shape), layout=new_layout,
+ scale=tuple([s * stride for s, stride in zip(x._scale, self.stride)]),
+ spatial_cache=x._spatial_cache,
+ )
+ out = out.replace(out.feats.type(dtype_))
+
+ if spatial_changed and (x.shape[0] != 1):
+ out.register_spatial_cache(f'conv_{self.stride}_unsorted_data', unsorted_data)
+ out.register_spatial_cache(f'conv_{self.stride}_sort_bwd', bwd)
+
+ return out
+
+class SparseInverseConv3d(nn.Module):
+ def __init__(self, in_channels, out_channels, kernel_size, stride=1, dilation=1, bias=True, indice_key=None):
+ super(SparseInverseConv3d, self).__init__()
+ if 'spconv' not in globals():
+ import spconv.pytorch as spconv
+ self.conv = spconv.SparseInverseConv3d(in_channels, out_channels, kernel_size, bias=bias, indice_key=indice_key)
+ self.stride = tuple(stride) if isinstance(stride, (list, tuple)) else (stride, stride, stride)
+
+ def forward(self, x: SparseTensor) -> SparseTensor:
+ spatial_changed = any(s != 1 for s in self.stride)
+ if spatial_changed:
+ # recover the original spconv order
+ data = x.get_spatial_cache(f'conv_{self.stride}_unsorted_data')
+ bwd = x.get_spatial_cache(f'conv_{self.stride}_sort_bwd')
+ data = data.replace_feature(x.feats[bwd])
+ if DEBUG:
+ assert torch.equal(data.indices, x.coords[bwd]), 'Recover the original order failed'
+ else:
+ data = x.data
+
+ new_data = self.conv(data)
+ new_shape = [x.shape[0], self.conv.out_channels]
+ new_layout = None if spatial_changed else x.layout
+ out = SparseTensor(
+ new_data, shape=torch.Size(new_shape), layout=new_layout,
+ scale=tuple([s // stride for s, stride in zip(x._scale, self.stride)]),
+ spatial_cache=x._spatial_cache,
+ )
+ return out
diff --git a/modules/sparse/conv/conv_torchsparse.py b/modules/sparse/conv/conv_torchsparse.py
new file mode 100644
index 0000000000000000000000000000000000000000..3140a69277176e7342db45f6a4c68c8da27419ec
--- /dev/null
+++ b/modules/sparse/conv/conv_torchsparse.py
@@ -0,0 +1,60 @@
+# MIT License
+
+# Copyright (c) Microsoft Corporation.
+# Copyright (c) 2025 VAST-AI-Research and contributors.
+
+# Permission is hereby granted, free of charge, to any person obtaining a copy
+# of this software and associated documentation files (the "Software"), to deal
+# in the Software without restriction, including without limitation the rights
+# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
+# copies of the Software, and to permit persons to whom the Software is
+# furnished to do so, subject to the following conditions:
+
+# The above copyright notice and this permission notice shall be included in all
+# copies or substantial portions of the Software.
+
+# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
+# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
+# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
+# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
+# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
+# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
+# SOFTWARE
+
+import torch
+import torch.nn as nn
+from .. import SparseTensor
+
+class SparseConv3d(nn.Module):
+ def __init__(self, in_channels, out_channels, kernel_size, stride=1, dilation=1, bias=True, indice_key=None):
+ super(SparseConv3d, self).__init__()
+ if 'torchsparse' not in globals():
+ import torchsparse
+ self.conv = torchsparse.nn.Conv3d(in_channels, out_channels, kernel_size, stride, 0, dilation, bias)
+
+ def forward(self, x: SparseTensor) -> SparseTensor:
+ out = self.conv(x.data)
+ new_shape = [x.shape[0], self.conv.out_channels]
+ out = SparseTensor(out, shape=torch.Size(new_shape), layout=x.layout if all(s == 1 for s in self.conv.stride) else None)
+ out._spatial_cache = x._spatial_cache
+ out._scale = tuple([s * stride for s, stride in zip(x._scale, self.conv.stride)])
+ return out
+
+
+class SparseInverseConv3d(nn.Module):
+ def __init__(self, in_channels, out_channels, kernel_size, stride=1, dilation=1, bias=True, indice_key=None):
+ super(SparseInverseConv3d, self).__init__()
+ if 'torchsparse' not in globals():
+ import torchsparse
+ self.conv = torchsparse.nn.Conv3d(in_channels, out_channels, kernel_size, stride, 0, dilation, bias, transposed=True)
+
+ def forward(self, x: SparseTensor) -> SparseTensor:
+ out = self.conv(x.data)
+ new_shape = [x.shape[0], self.conv.out_channels]
+ out = SparseTensor(out, shape=torch.Size(new_shape), layout=x.layout if all(s == 1 for s in self.conv.stride) else None)
+ out._spatial_cache = x._spatial_cache
+ out._scale = tuple([s // stride for s, stride in zip(x._scale, self.conv.stride)])
+ return out
+
+
+
diff --git a/modules/sparse/linear.py b/modules/sparse/linear.py
new file mode 100644
index 0000000000000000000000000000000000000000..4b4cd0b52a2fcca7b6fce450afc63b08ff756303
--- /dev/null
+++ b/modules/sparse/linear.py
@@ -0,0 +1,38 @@
+
+# MIT License
+
+# Copyright (c) Microsoft Corporation.
+# Copyright (c) 2025 VAST-AI-Research and contributors.
+
+# Permission is hereby granted, free of charge, to any person obtaining a copy
+# of this software and associated documentation files (the "Software"), to deal
+# in the Software without restriction, including without limitation the rights
+# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
+# copies of the Software, and to permit persons to whom the Software is
+# furnished to do so, subject to the following conditions:
+
+# The above copyright notice and this permission notice shall be included in all
+# copies or substantial portions of the Software.
+
+# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
+# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
+# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
+# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
+# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
+# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
+# SOFTWARE
+
+import torch
+import torch.nn as nn
+from . import SparseTensor
+
+__all__ = [
+ 'SparseLinear'
+]
+
+class SparseLinear(nn.Linear):
+ def __init__(self, in_features, out_features, bias=True):
+ super(SparseLinear, self).__init__(in_features, out_features, bias)
+
+ def forward(self, input: SparseTensor) -> SparseTensor:
+ return input.replace(super().forward(input.feats))
diff --git a/modules/sparse/nonlinearity.py b/modules/sparse/nonlinearity.py
new file mode 100644
index 0000000000000000000000000000000000000000..53b45be4d0d9167784f1141ef5d4dc6c610002d9
--- /dev/null
+++ b/modules/sparse/nonlinearity.py
@@ -0,0 +1,58 @@
+# MIT License
+
+# Copyright (c) Microsoft Corporation.
+# Copyright (c) 2025 VAST-AI-Research and contributors.
+
+# Permission is hereby granted, free of charge, to any person obtaining a copy
+# of this software and associated documentation files (the "Software"), to deal
+# in the Software without restriction, including without limitation the rights
+# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
+# copies of the Software, and to permit persons to whom the Software is
+# furnished to do so, subject to the following conditions:
+
+# The above copyright notice and this permission notice shall be included in all
+# copies or substantial portions of the Software.
+
+# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
+# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
+# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
+# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
+# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
+# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
+# SOFTWARE
+
+import torch
+import torch.nn as nn
+from . import SparseTensor
+
+__all__ = [
+ 'SparseReLU',
+ 'SparseSiLU',
+ 'SparseGELU',
+ 'SparseActivation'
+]
+
+
+class SparseReLU(nn.ReLU):
+ def forward(self, input: SparseTensor) -> SparseTensor:
+ return input.replace(super().forward(input.feats))
+
+
+class SparseSiLU(nn.SiLU):
+ def forward(self, input: SparseTensor) -> SparseTensor:
+ return input.replace(super().forward(input.feats))
+
+
+class SparseGELU(nn.GELU):
+ def forward(self, input: SparseTensor) -> SparseTensor:
+ return input.replace(super().forward(input.feats))
+
+
+class SparseActivation(nn.Module):
+ def __init__(self, activation: nn.Module):
+ super().__init__()
+ self.activation = activation
+
+ def forward(self, input: SparseTensor) -> SparseTensor:
+ return input.replace(self.activation(input.feats))
+
diff --git a/modules/sparse/norm.py b/modules/sparse/norm.py
new file mode 100644
index 0000000000000000000000000000000000000000..b7cd9d86b55427f26ec129b9626d95b3307b368b
--- /dev/null
+++ b/modules/sparse/norm.py
@@ -0,0 +1,81 @@
+# MIT License
+
+# Copyright (c) Microsoft Corporation.
+# Copyright (c) 2025 VAST-AI-Research and contributors.
+
+# Permission is hereby granted, free of charge, to any person obtaining a copy
+# of this software and associated documentation files (the "Software"), to deal
+# in the Software without restriction, including without limitation the rights
+# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
+# copies of the Software, and to permit persons to whom the Software is
+# furnished to do so, subject to the following conditions:
+
+# The above copyright notice and this permission notice shall be included in all
+# copies or substantial portions of the Software.
+
+# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
+# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
+# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
+# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
+# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
+# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
+# SOFTWARE
+
+import torch
+import torch.nn as nn
+from . import SparseTensor
+from . import DEBUG
+
+__all__ = [
+ 'SparseGroupNorm',
+ 'SparseLayerNorm',
+ 'SparseGroupNorm32',
+ 'SparseLayerNorm32',
+]
+
+
+class SparseGroupNorm(nn.GroupNorm):
+ def __init__(self, num_groups, num_channels, eps=1e-5, affine=True):
+ super(SparseGroupNorm, self).__init__(num_groups, num_channels, eps, affine)
+
+ def forward(self, input: SparseTensor) -> SparseTensor:
+ nfeats = torch.zeros_like(input.feats)
+ for k in range(input.shape[0]):
+ if DEBUG:
+ assert (input.coords[input.layout[k], 0] == k).all(), f"SparseGroupNorm: batch index mismatch"
+ bfeats = input.feats[input.layout[k]]
+ bfeats = bfeats.permute(1, 0).reshape(1, input.shape[1], -1)
+ bfeats = super().forward(bfeats)
+ bfeats = bfeats.reshape(input.shape[1], -1).permute(1, 0)
+ nfeats[input.layout[k]] = bfeats
+ return input.replace(nfeats)
+
+
+class SparseLayerNorm(nn.LayerNorm):
+ def __init__(self, normalized_shape, eps=1e-5, elementwise_affine=True):
+ super(SparseLayerNorm, self).__init__(normalized_shape, eps, elementwise_affine)
+
+ def forward(self, input: SparseTensor) -> SparseTensor:
+ nfeats = torch.zeros_like(input.feats)
+ for k in range(input.shape[0]):
+ bfeats = input.feats[input.layout[k]]
+ bfeats = bfeats.permute(1, 0).reshape(1, input.shape[1], -1)
+ bfeats = super().forward(bfeats)
+ bfeats = bfeats.reshape(input.shape[1], -1).permute(1, 0)
+ nfeats[input.layout[k]] = bfeats
+ return input.replace(nfeats)
+
+
+class SparseGroupNorm32(SparseGroupNorm):
+ """
+ A GroupNorm layer that converts to float32 before the forward pass.
+ """
+ def forward(self, x: SparseTensor) -> SparseTensor:
+ return super().forward(x.float()).type(x.dtype)
+
+class SparseLayerNorm32(SparseLayerNorm):
+ """
+ A LayerNorm layer that converts to float32 before the forward pass.
+ """
+ def forward(self, x: SparseTensor) -> SparseTensor:
+ return super().forward(x.float()).type(x.dtype)
diff --git a/modules/sparse/spatial.py b/modules/sparse/spatial.py
new file mode 100644
index 0000000000000000000000000000000000000000..828b8fe24653484437c2958645da97a4b60b1dfe
--- /dev/null
+++ b/modules/sparse/spatial.py
@@ -0,0 +1,158 @@
+# MIT License
+
+# Copyright (c) Microsoft Corporation.
+# Copyright (c) 2025 VAST-AI-Research and contributors.
+
+# Permission is hereby granted, free of charge, to any person obtaining a copy
+# of this software and associated documentation files (the "Software"), to deal
+# in the Software without restriction, including without limitation the rights
+# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
+# copies of the Software, and to permit persons to whom the Software is
+# furnished to do so, subject to the following conditions:
+
+# The above copyright notice and this permission notice shall be included in all
+# copies or substantial portions of the Software.
+
+# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
+# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
+# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
+# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
+# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
+# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
+# SOFTWARE
+
+from typing import *
+import torch
+import torch.nn as nn
+from . import SparseTensor
+
+__all__ = [
+ "SparseDownsample",
+ "SparseUpsample",
+ "SparseSubdivide",
+]
+
+
+class SparseDownsample(nn.Module):
+ """
+ Downsample a sparse tensor by a factor of `factor`.
+ Implemented as average pooling.
+ """
+
+ def __init__(self, factor: Union[int, Tuple[int, ...], List[int]]):
+ super(SparseDownsample, self).__init__()
+ self.factor = tuple(factor) if isinstance(factor, (list, tuple)) else factor
+
+ def forward(self, input: SparseTensor) -> SparseTensor:
+ DIM = input.coords.shape[-1] - 1
+ factor = self.factor if isinstance(self.factor, tuple) else (self.factor,) * DIM
+ assert DIM == len(
+ factor
+ ), "Input coordinates must have the same dimension as the downsample factor."
+
+ coord = list(input.coords.unbind(dim=-1))
+ for i, f in enumerate(factor):
+ coord[i + 1] = coord[i + 1] // f
+
+ MAX = [coord[i + 1].max().item() + 1 for i in range(DIM)]
+ OFFSET = torch.cumprod(torch.tensor(MAX[::-1]), 0).tolist()[::-1] + [1]
+ code = sum([c * o for c, o in zip(coord, OFFSET)])
+ code, idx = code.unique(return_inverse=True)
+
+ new_feats = torch.scatter_reduce(
+ torch.zeros(
+ code.shape[0],
+ input.feats.shape[1],
+ device=input.feats.device,
+ dtype=input.feats.dtype,
+ ),
+ dim=0,
+ index=idx.unsqueeze(1).expand(-1, input.feats.shape[1]),
+ src=input.feats,
+ # reduce='mean',
+ reduce="amax",
+ )
+ new_coords = torch.stack(
+ [code // OFFSET[0]]
+ + [(code // OFFSET[i + 1]) % MAX[i] for i in range(DIM)],
+ dim=-1,
+ )
+ out = SparseTensor(
+ new_feats,
+ new_coords,
+ input.shape,
+ )
+ out._scale = tuple([s // f for s, f in zip(input._scale, factor)])
+ out._spatial_cache = input._spatial_cache
+
+ out.register_spatial_cache(f"upsample_{factor}_coords", input.coords)
+ out.register_spatial_cache(f"upsample_{factor}_layout", input.layout)
+ out.register_spatial_cache(f"upsample_{factor}_idx", idx)
+
+ return out
+
+
+class SparseUpsample(nn.Module):
+ """
+ Upsample a sparse tensor by a factor of `factor`.
+ Implemented as nearest neighbor interpolation.
+ """
+
+ def __init__(self, factor: Union[int, Tuple[int, int, int], List[int]]):
+ super(SparseUpsample, self).__init__()
+ self.factor = tuple(factor) if isinstance(factor, (list, tuple)) else factor
+
+ def forward(self, input: SparseTensor) -> SparseTensor:
+ DIM = input.coords.shape[-1] - 1
+ factor = self.factor if isinstance(self.factor, tuple) else (self.factor,) * DIM
+ assert DIM == len(
+ factor
+ ), "Input coordinates must have the same dimension as the upsample factor."
+
+ new_coords = input.get_spatial_cache(f"upsample_{factor}_coords")
+ new_layout = input.get_spatial_cache(f"upsample_{factor}_layout")
+ idx = input.get_spatial_cache(f"upsample_{factor}_idx")
+ if any([x is None for x in [new_coords, new_layout, idx]]):
+ raise ValueError(
+ "Upsample cache not found. SparseUpsample must be paired with SparseDownsample."
+ )
+ new_feats = input.feats[idx]
+ out = SparseTensor(new_feats, new_coords, input.shape, new_layout)
+ out._scale = tuple([s * f for s, f in zip(input._scale, factor)])
+ out._spatial_cache = input._spatial_cache
+ return out
+
+
+class SparseSubdivide(nn.Module):
+ """
+ Upsample a sparse tensor by a factor of `factor`.
+ Implemented as nearest neighbor interpolation.
+ """
+
+ def __init__(self):
+ super(SparseSubdivide, self).__init__()
+
+ def forward(self, input: SparseTensor) -> SparseTensor:
+ DIM = input.coords.shape[-1] - 1
+ # upsample scale=2^DIM
+ n_cube = torch.ones([2] * DIM, device=input.device, dtype=torch.int)
+ n_coords = torch.nonzero(n_cube)
+ n_coords = torch.cat([torch.zeros_like(n_coords[:, :1]), n_coords], dim=-1)
+ factor = n_coords.shape[0]
+ assert factor == 2**DIM
+ # print(n_coords.shape)
+ new_coords = input.coords.clone()
+ new_coords[:, 1:] *= 2
+ new_coords = new_coords.unsqueeze(1) + n_coords.unsqueeze(0).to(
+ new_coords.dtype
+ )
+
+ new_feats = input.feats.unsqueeze(1).expand(
+ input.feats.shape[0], factor, *input.feats.shape[1:]
+ )
+ out = SparseTensor(
+ new_feats.flatten(0, 1), new_coords.flatten(0, 1), input.shape
+ )
+ out._scale = input._scale * 2
+ out._spatial_cache = input._spatial_cache
+ return out
diff --git a/modules/sparse/transformer/__init__.py b/modules/sparse/transformer/__init__.py
new file mode 100644
index 0000000000000000000000000000000000000000..9cb7b69c531b6317c18872e0a977e04cfcb53933
--- /dev/null
+++ b/modules/sparse/transformer/__init__.py
@@ -0,0 +1,26 @@
+# MIT License
+
+# Copyright (c) Microsoft Corporation.
+# Copyright (c) 2025 VAST-AI-Research and contributors.
+
+# Permission is hereby granted, free of charge, to any person obtaining a copy
+# of this software and associated documentation files (the "Software"), to deal
+# in the Software without restriction, including without limitation the rights
+# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
+# copies of the Software, and to permit persons to whom the Software is
+# furnished to do so, subject to the following conditions:
+
+# The above copyright notice and this permission notice shall be included in all
+# copies or substantial portions of the Software.
+
+# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
+# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
+# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
+# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
+# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
+# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
+# SOFTWARE
+
+from .blocks import *
+from .modulated import *
+from .bases import *
\ No newline at end of file
diff --git a/modules/sparse/transformer/bases.py b/modules/sparse/transformer/bases.py
new file mode 100644
index 0000000000000000000000000000000000000000..806c89edf61dd16e528e92ac57c423ac0d86a865
--- /dev/null
+++ b/modules/sparse/transformer/bases.py
@@ -0,0 +1,234 @@
+# MIT License
+
+# Copyright (c) Microsoft Corporation.
+# Copyright (c) 2025 VAST-AI-Research and contributors.
+
+# Permission is hereby granted, free of charge, to any person obtaining a copy
+# of this software and associated documentation files (the "Software"), to deal
+# in the Software without restriction, including without limitation the rights
+# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
+# copies of the Software, and to permit persons to whom the Software is
+# furnished to do so, subject to the following conditions:
+
+# The above copyright notice and this permission notice shall be included in all
+# copies or substantial portions of the Software.
+
+# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
+# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
+# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
+# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
+# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
+# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
+# SOFTWARE
+
+from typing import *
+import torch
+import torch.nn as nn
+from ...utils import convert_module_to_f16, convert_module_to_f32
+from ...transformer import AbsolutePositionEmbedder
+from modules import sparse as sp
+from .blocks import SparseTransformerBlock, SparseTransformerCrossBlock
+
+
+def block_attn_config(self):
+ """
+ Return the attention configuration of the model.
+ """
+ for i in range(self.num_blocks):
+ if self.attn_mode == "shift_window":
+ yield "serialized", self.window_size, 0, (16 * (i % 2),) * 3, sp.SerializeMode.Z_ORDER
+ elif self.attn_mode == "shift_sequence":
+ yield "serialized", self.window_size, self.window_size // 2 * (i % 2), (0, 0, 0), sp.SerializeMode.Z_ORDER
+ elif self.attn_mode == "shift_order":
+ yield "serialized", self.window_size, 0, (0, 0, 0), sp.SerializeModes[i % 4]
+ elif self.attn_mode == "full":
+ yield "full", None, None, None, None
+ elif self.attn_mode == "swin":
+ yield "windowed", self.window_size, None, self.window_size // 2 * (i % 2), None
+
+
+class SparseTransformerBase(nn.Module):
+ """
+ Sparse Transformer without output layers.
+ Serve as the base class for encoder and decoder.
+ """
+ def __init__(
+ self,
+ in_channels: int,
+ model_channels: int,
+ num_blocks: int,
+ num_heads: Optional[int] = None,
+ num_head_channels: Optional[int] = 64,
+ mlp_ratio: float = 4.0,
+ attn_mode: Literal["full", "shift_window", "shift_sequence", "shift_order", "swin"] = "full",
+ window_size: Optional[int] = None,
+ pe_mode: Literal["ape", "rope"] = "ape",
+ use_fp16: bool = False,
+ use_checkpoint: bool = False,
+ qk_rms_norm: bool = False,
+ ):
+ super().__init__()
+ self.in_channels = in_channels
+ self.model_channels = model_channels
+ self.num_blocks = num_blocks
+ self.window_size = window_size
+ self.num_heads = num_heads or model_channels // num_head_channels
+ self.mlp_ratio = mlp_ratio
+ self.attn_mode = attn_mode
+ self.pe_mode = pe_mode
+ self.use_fp16 = use_fp16
+ self.use_checkpoint = use_checkpoint
+ self.qk_rms_norm = qk_rms_norm
+ self.dtype = torch.float16 if use_fp16 else torch.float32
+
+ if pe_mode == "ape":
+ self.pos_embedder = AbsolutePositionEmbedder(model_channels)
+
+ self.input_layer = sp.SparseLinear(in_channels, model_channels)
+ self.blocks = nn.ModuleList([
+ SparseTransformerBlock(
+ model_channels,
+ num_heads=self.num_heads,
+ mlp_ratio=self.mlp_ratio,
+ attn_mode=attn_mode,
+ window_size=window_size,
+ shift_sequence=shift_sequence,
+ shift_window=shift_window,
+ serialize_mode=serialize_mode,
+ use_checkpoint=self.use_checkpoint,
+ use_rope=(pe_mode == "rope"),
+ qk_rms_norm=self.qk_rms_norm,
+ )
+ for attn_mode, window_size, shift_sequence, shift_window, serialize_mode in block_attn_config(self)
+ ])
+
+ @property
+ def device(self) -> torch.device:
+ """
+ Return the device of the model.
+ """
+ return next(self.parameters()).device
+
+ def convert_to_fp16(self) -> None:
+ """
+ Convert the torso of the model to float16.
+ """
+ self.blocks.apply(convert_module_to_f16)
+
+ def convert_to_fp32(self) -> None:
+ """
+ Convert the torso of the model to float32.
+ """
+ self.blocks.apply(convert_module_to_f32)
+
+ def initialize_weights(self) -> None:
+ # Initialize transformer layers:
+ def _basic_init(module):
+ if isinstance(module, nn.Linear):
+ torch.nn.init.xavier_uniform_(module.weight)
+ if module.bias is not None:
+ nn.init.constant_(module.bias, 0)
+ self.apply(_basic_init)
+
+ def forward(self, x: sp.SparseTensor) -> sp.SparseTensor:
+ h = self.input_layer(x)
+ if self.pe_mode == "ape" and len(self.blocks) != 0:
+ h = h + self.pos_embedder(x.coords[:, 1:])
+ for block in self.blocks:
+ h = block(h)
+ return h
+
+class SparseTransformerCrossBase(nn.Module):
+ """
+ Sparse Transformer without output layers.
+ Serve as the base class for encoder and decoder.
+ """
+ def __init__(
+ self,
+ in_channels: int,
+ model_channels: int,
+ context_channels: int,
+ num_blocks: int,
+ num_heads: Optional[int] = None,
+ num_head_channels: Optional[int] = 64,
+ mlp_ratio: float = 4.0,
+ attn_mode: Literal["full", "shift_window", "shift_sequence", "shift_order", "swin"] = "full",
+ window_size: Optional[int] = None,
+ pe_mode: Literal["ape", "rope"] = "ape",
+ use_fp16: bool = False,
+ use_checkpoint: bool = False,
+ qk_rms_norm: bool = False,
+ ):
+ super().__init__()
+ self.in_channels = in_channels
+ self.model_channels = model_channels
+ self.num_blocks = num_blocks
+ self.window_size = window_size
+ self.num_heads = num_heads or model_channels // num_head_channels
+ self.mlp_ratio = mlp_ratio
+ self.attn_mode = attn_mode
+ self.pe_mode = pe_mode
+ self.use_fp16 = use_fp16
+ self.use_checkpoint = use_checkpoint
+ self.qk_rms_norm = qk_rms_norm
+ self.dtype = torch.float16 if use_fp16 else torch.float32
+
+ if pe_mode == "ape":
+ self.pos_embedder_x = AbsolutePositionEmbedder(model_channels)
+ self.pos_embedder_ctx = AbsolutePositionEmbedder(context_channels)
+
+ self.input_layer = sp.SparseLinear(in_channels, model_channels)
+ self.blocks = nn.ModuleList([
+ SparseTransformerCrossBlock(
+ model_channels,
+ num_heads=self.num_heads,
+ ctx_channels=context_channels,
+ mlp_ratio=self.mlp_ratio,
+ attn_mode=attn_mode,
+ window_size=window_size,
+ shift_sequence=shift_sequence,
+ shift_window=shift_window,
+ serialize_mode=serialize_mode,
+ use_checkpoint=self.use_checkpoint,
+ use_rope=(pe_mode == "rope"),
+ qk_rms_norm=self.qk_rms_norm,
+ )
+ for attn_mode, window_size, shift_sequence, shift_window, serialize_mode in block_attn_config(self)
+ ])
+
+ @property
+ def device(self) -> torch.device:
+ """
+ Return the device of the model.
+ """
+ return next(self.parameters()).device
+
+ def convert_to_fp16(self) -> None:
+ """
+ Convert the torso of the model to float16.
+ """
+ self.blocks.apply(convert_module_to_f16)
+
+ def convert_to_fp32(self) -> None:
+ """
+ Convert the torso of the model to float32.
+ """
+ self.blocks.apply(convert_module_to_f32)
+
+ def initialize_weights(self) -> None:
+ # Initialize transformer layers:
+ def _basic_init(module):
+ if isinstance(module, nn.Linear):
+ torch.nn.init.xavier_uniform_(module.weight)
+ if module.bias is not None:
+ nn.init.constant_(module.bias, 0)
+ self.apply(_basic_init)
+
+ def forward(self, x: sp.SparseTensor, context: sp.SparseTensor) -> sp.SparseTensor:
+ h = self.input_layer(x)
+ if self.pe_mode == "ape" and len(self.blocks) != 0:
+ h = h + self.pos_embedder_x(x.coords[:, 1:])
+ context = context + self.pos_embedder_ctx(context.coords[:, 1:])
+ for block in self.blocks:
+ h = block(h, context)
+ return h
diff --git a/modules/sparse/transformer/blocks.py b/modules/sparse/transformer/blocks.py
new file mode 100644
index 0000000000000000000000000000000000000000..5bc3bb25226a07ab4fdda5ea2d084d7dd34c39cc
--- /dev/null
+++ b/modules/sparse/transformer/blocks.py
@@ -0,0 +1,165 @@
+from typing import *
+import torch
+import torch.nn as nn
+from ..basic import SparseTensor
+from ..linear import SparseLinear
+from ..nonlinearity import SparseGELU
+from ..attention import SparseMultiHeadAttention, SerializeMode
+from ...norm import LayerNorm32
+
+
+class SparseFeedForwardNet(nn.Module):
+ def __init__(self, channels: int, mlp_ratio: float = 4.0):
+ super().__init__()
+ self.mlp = nn.Sequential(
+ SparseLinear(channels, int(channels * mlp_ratio)),
+ SparseGELU(approximate="tanh"),
+ SparseLinear(int(channels * mlp_ratio), channels),
+ )
+
+ def forward(self, x: SparseTensor) -> SparseTensor:
+ return self.mlp(x)
+
+
+class SparseTransformerBlock(nn.Module):
+ """
+ Sparse Transformer block (MSA + FFN).
+ """
+
+ def __init__(
+ self,
+ channels: int,
+ num_heads: int,
+ mlp_ratio: float = 4.0,
+ attn_mode: Literal[
+ "full", "shift_window", "shift_sequence", "shift_order", "swin"
+ ] = "full",
+ window_size: Optional[int] = None,
+ shift_sequence: Optional[int] = None,
+ shift_window: Optional[Tuple[int, int, int]] = None,
+ serialize_mode: Optional[SerializeMode] = None,
+ use_checkpoint: bool = False,
+ use_rope: bool = False,
+ qk_rms_norm: bool = False,
+ qkv_bias: bool = True,
+ ln_affine: bool = False,
+ ):
+ super().__init__()
+ self.use_checkpoint = use_checkpoint
+ self.norm1 = LayerNorm32(channels, elementwise_affine=ln_affine, eps=1e-6)
+ self.norm2 = LayerNorm32(channels, elementwise_affine=ln_affine, eps=1e-6)
+ self.attn = SparseMultiHeadAttention(
+ channels,
+ num_heads=num_heads,
+ attn_mode=attn_mode,
+ window_size=window_size,
+ shift_sequence=shift_sequence,
+ shift_window=shift_window,
+ serialize_mode=serialize_mode,
+ qkv_bias=qkv_bias,
+ use_rope=use_rope,
+ qk_rms_norm=qk_rms_norm,
+ )
+ self.mlp = SparseFeedForwardNet(
+ channels,
+ mlp_ratio=mlp_ratio,
+ )
+
+ def _forward(self, x: SparseTensor) -> SparseTensor:
+ h = x.replace(self.norm1(x.feats))
+ h = self.attn(h)
+ x = x + h
+ h = x.replace(self.norm2(x.feats))
+ h = self.mlp(h)
+ x = x + h
+ return x
+
+ def forward(self, x: SparseTensor) -> SparseTensor:
+ if self.use_checkpoint:
+ return torch.utils.checkpoint.checkpoint(
+ self._forward, x, use_reentrant=False
+ )
+ else:
+ return self._forward(x)
+
+
+class SparseTransformerCrossBlock(nn.Module):
+ """
+ Sparse Transformer cross-attention block (MSA + MCA + FFN).
+ """
+
+ def __init__(
+ self,
+ channels: int,
+ ctx_channels: int,
+ num_heads: int,
+ mlp_ratio: float = 4.0,
+ attn_mode: Literal[
+ "full", "shift_window", "shift_sequence", "shift_order", "swin"
+ ] = "full",
+ window_size: Optional[int] = None,
+ shift_sequence: Optional[int] = None,
+ shift_window: Optional[Tuple[int, int, int]] = None,
+ serialize_mode: Optional[SerializeMode] = None,
+ use_checkpoint: bool = False,
+ use_rope: bool = False,
+ qk_rms_norm: bool = False,
+ qk_rms_norm_cross: bool = False,
+ qkv_bias: bool = True,
+ ln_affine: bool = False,
+ ):
+ super().__init__()
+ self.use_checkpoint = use_checkpoint
+ self.norm1 = LayerNorm32(channels, elementwise_affine=ln_affine, eps=1e-6)
+ self.norm2 = LayerNorm32(channels, elementwise_affine=ln_affine, eps=1e-6)
+ self.norm3 = LayerNorm32(channels, elementwise_affine=ln_affine, eps=1e-6)
+ self.context_norm = LayerNorm32(
+ ctx_channels, elementwise_affine=ln_affine, eps=1e-6
+ )
+ self.self_attn = SparseMultiHeadAttention(
+ channels,
+ num_heads=num_heads,
+ type="self",
+ attn_mode=attn_mode,
+ window_size=window_size,
+ shift_sequence=shift_sequence,
+ shift_window=shift_window,
+ serialize_mode=serialize_mode,
+ qkv_bias=qkv_bias,
+ use_rope=use_rope,
+ qk_rms_norm=qk_rms_norm,
+ )
+ self.cross_attn = SparseMultiHeadAttention(
+ channels,
+ ctx_channels=ctx_channels,
+ num_heads=num_heads,
+ type="cross",
+ attn_mode="full",
+ qkv_bias=qkv_bias,
+ qk_rms_norm=qk_rms_norm_cross,
+ )
+ self.mlp = SparseFeedForwardNet(
+ channels,
+ mlp_ratio=mlp_ratio,
+ )
+
+ def _forward(self, x: SparseTensor, context: torch.Tensor):
+ h = x.replace(self.norm1(x.feats))
+ h = self.self_attn(h)
+ x = x + h
+ h = x.replace(self.norm2(x.feats))
+
+ h = self.cross_attn(h, context)
+ x = x + h
+ h = x.replace(self.norm3(x.feats))
+ h = self.mlp(h)
+ x = x + h
+ return x
+
+ def forward(self, x: SparseTensor, context: torch.Tensor):
+ if self.use_checkpoint:
+ return torch.utils.checkpoint.checkpoint(
+ self._forward, x, context, use_reentrant=False
+ )
+ else:
+ return self._forward(x, context)
\ No newline at end of file
diff --git a/modules/sparse/transformer/modulated.py b/modules/sparse/transformer/modulated.py
new file mode 100644
index 0000000000000000000000000000000000000000..e991497bf40e07aba42232c2f2727363ad463a1a
--- /dev/null
+++ b/modules/sparse/transformer/modulated.py
@@ -0,0 +1,119 @@
+# MIT License
+
+# Copyright (c) Microsoft Corporation.
+# Copyright (c) 2025 VAST-AI-Research and contributors.
+
+# Permission is hereby granted, free of charge, to any person obtaining a copy
+# of this software and associated documentation files (the "Software"), to deal
+# in the Software without restriction, including without limitation the rights
+# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
+# copies of the Software, and to permit persons to whom the Software is
+# furnished to do so, subject to the following conditions:
+
+# The above copyright notice and this permission notice shall be included in all
+# copies or substantial portions of the Software.
+
+# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
+# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
+# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
+# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
+# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
+# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
+# SOFTWARE
+
+from typing import *
+import torch
+import torch.nn as nn
+import torch.utils.checkpoint
+from ..basic import SparseTensor
+from ..attention import SparseMultiHeadAttention, SerializeMode
+from ...norm import LayerNorm32
+from .blocks import SparseFeedForwardNet
+
+
+class ModulatedSparseTransformerCrossBlock(nn.Module):
+ """
+ Sparse Transformer cross-attention block (MSA + MCA + FFN) with adaptive layer norm conditioning.
+ """
+ def __init__(
+ self,
+ channels: int,
+ ctx_channels: int,
+ num_heads: int,
+ mlp_ratio: float = 4.0,
+ attn_mode: Literal["full", "shift_window", "shift_sequence", "shift_order", "swin"] = "full",
+ window_size: Optional[int] = None,
+ shift_sequence: Optional[int] = None,
+ shift_window: Optional[Tuple[int, int, int]] = None,
+ serialize_mode: Optional[SerializeMode] = None,
+ use_checkpoint: bool = False,
+ use_rope: bool = False,
+ qk_rms_norm: bool = False,
+ qk_rms_norm_cross: bool = False,
+ qkv_bias: bool = True,
+ share_mod: bool = False,
+
+ ):
+ super().__init__()
+ self.use_checkpoint = use_checkpoint
+ self.share_mod = share_mod
+ self.norm1 = LayerNorm32(channels, elementwise_affine=False, eps=1e-6)
+ self.norm2 = LayerNorm32(channels, elementwise_affine=True, eps=1e-6)
+ self.norm3 = LayerNorm32(channels, elementwise_affine=False, eps=1e-6)
+ self.self_attn = SparseMultiHeadAttention(
+ channels,
+ num_heads=num_heads,
+ type="self",
+ attn_mode=attn_mode,
+ window_size=window_size,
+ shift_sequence=shift_sequence,
+ shift_window=shift_window,
+ serialize_mode=serialize_mode,
+ qkv_bias=qkv_bias,
+ use_rope=use_rope,
+ qk_rms_norm=qk_rms_norm,
+ )
+ self.cross_attn = SparseMultiHeadAttention(
+ channels,
+ ctx_channels=ctx_channels,
+ num_heads=num_heads,
+ type="cross",
+ attn_mode="full",
+ qkv_bias=qkv_bias,
+ qk_rms_norm=qk_rms_norm_cross,
+ )
+ self.mlp = SparseFeedForwardNet(
+ channels,
+ mlp_ratio=mlp_ratio,
+ )
+ if not share_mod:
+ self.adaLN_modulation = nn.Sequential(
+ nn.SiLU(),
+ nn.Linear(channels, 6 * channels, bias=True)
+ )
+
+ def _forward(self, x: SparseTensor, mod: torch.Tensor, context: torch.Tensor) -> SparseTensor:
+ if self.share_mod:
+ shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = mod.chunk(6, dim=1)
+ else:
+ shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = self.adaLN_modulation(mod).chunk(6, dim=1)
+ h = x.replace(self.norm1(x.feats))
+ h = h * (1 + scale_msa) + shift_msa
+ h = self.self_attn(h)
+ h = h * gate_msa
+ x = x + h
+ h = x.replace(self.norm2(x.feats))
+ h = self.cross_attn(h, context)
+ x = x + h
+ h = x.replace(self.norm3(x.feats))
+ h = h * (1 + scale_mlp) + shift_mlp
+ h = self.mlp(h)
+ h = h * gate_mlp
+ x = x + h
+ return x
+
+ def forward(self, x: SparseTensor, mod: torch.Tensor, context: torch.Tensor) -> SparseTensor:
+ if self.use_checkpoint:
+ return torch.utils.checkpoint.checkpoint(self._forward, x, mod, context, use_reentrant=False)
+ else:
+ return self._forward(x, mod, context)
diff --git a/modules/transformer/__init__.py b/modules/transformer/__init__.py
new file mode 100644
index 0000000000000000000000000000000000000000..753e5c32da0c87a33752729207b0651ae007b6f7
--- /dev/null
+++ b/modules/transformer/__init__.py
@@ -0,0 +1,24 @@
+# MIT License
+
+# Copyright (c) Microsoft Corporation.
+# Copyright (c) 2025 VAST-AI-Research and contributors.
+
+# Permission is hereby granted, free of charge, to any person obtaining a copy
+# of this software and associated documentation files (the "Software"), to deal
+# in the Software without restriction, including without limitation the rights
+# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
+# copies of the Software, and to permit persons to whom the Software is
+# furnished to do so, subject to the following conditions:
+
+# The above copyright notice and this permission notice shall be included in all
+# copies or substantial portions of the Software.
+
+# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
+# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
+# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
+# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
+# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
+# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
+# SOFTWARE
+
+from .blocks import *
\ No newline at end of file
diff --git a/modules/transformer/blocks.py b/modules/transformer/blocks.py
new file mode 100644
index 0000000000000000000000000000000000000000..b8f36d1063cb94d6ca43cdea9b6417d2a3e2ace2
--- /dev/null
+++ b/modules/transformer/blocks.py
@@ -0,0 +1,276 @@
+# MIT License
+
+# Copyright (c) Microsoft Corporation.
+# Copyright (c) 2025 VAST-AI-Research and contributors.
+
+# Permission is hereby granted, free of charge, to any person obtaining a copy
+# of this software and associated documentation files (the "Software"), to deal
+# in the Software without restriction, including without limitation the rights
+# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
+# copies of the Software, and to permit persons to whom the Software is
+# furnished to do so, subject to the following conditions:
+
+# The above copyright notice and this permission notice shall be included in all
+# copies or substantial portions of the Software.
+
+# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
+# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
+# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
+# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
+# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
+# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
+# SOFTWARE
+
+from typing import *
+import numpy as np
+import torch
+import torch.nn as nn
+import torch.nn.functional as F
+
+
+class AbsolutePositionEmbedder(nn.Module):
+ """
+ Embeds spatial positions into vector representations.
+ """
+
+ def __init__(self, channels: int, in_channels: int = 3):
+ super().__init__()
+ self.channels = channels
+ self.in_channels = in_channels
+ self.freq_dim = channels // in_channels // 2
+ self.freqs = torch.arange(self.freq_dim, dtype=torch.float32) / self.freq_dim
+ self.freqs = 1.0 / (10000**self.freqs)
+
+ def _sin_cos_embedding(self, x: torch.Tensor) -> torch.Tensor:
+ """
+ Create sinusoidal position embeddings.
+
+ Args:
+ x: a 1-D Tensor of N indices
+
+ Returns:
+ an (N, D) Tensor of positional embeddings.
+ """
+ self.freqs = self.freqs.to(x.device)
+ out = torch.outer(x, self.freqs)
+ out = torch.cat([torch.sin(out), torch.cos(out)], dim=-1)
+ return out
+
+ def forward(self, x: torch.Tensor) -> torch.Tensor:
+ """
+ Args:
+ x (torch.Tensor): (N, D) tensor of spatial positions
+ """
+ N, D = x.shape
+ assert (
+ D == self.in_channels
+ ), "Input dimension must match number of input channels"
+ embed = self._sin_cos_embedding(x.reshape(-1))
+ embed = embed.reshape(N, -1)
+ if embed.shape[1] < self.channels:
+ embed = torch.cat(
+ [
+ embed,
+ torch.zeros(N, self.channels - embed.shape[1], device=embed.device),
+ ],
+ dim=-1,
+ )
+ return embed
+
+
+class RotaryPositionPhasesEmbedder(nn.Module):
+ def __init__(
+ self,
+ head_dim: int,
+ dim: int = 3,
+ rope_freq: Tuple[float, float] = (1.0, 10000.0),
+ ):
+ super().__init__()
+ assert head_dim % 2 == 0, "Head dim must be divisible by 2"
+ self.head_dim = head_dim
+ self.dim = dim
+ self.rope_freq = rope_freq
+ self.freq_dim = head_dim // 2 // dim
+ self.freqs = torch.arange(self.freq_dim, dtype=torch.float32) / self.freq_dim
+ self.freqs = rope_freq[0] / (rope_freq[1] ** (self.freqs))
+
+ def _get_phases(self, indices: torch.Tensor) -> torch.Tensor:
+ self.freqs = self.freqs.to(indices.device)
+ phases = torch.outer(indices, self.freqs)
+ phases = torch.polar(torch.ones_like(phases), phases)
+ return phases
+
+ @staticmethod
+ def apply_rotary_embedding(x: torch.Tensor, phases: torch.Tensor) -> torch.Tensor:
+ x_complex = torch.view_as_complex(x.float().reshape(*x.shape[:-1], -1, 2))
+ if phases.ndim == 3:
+ phases = phases.unsqueeze(1)
+ x_rotated = x_complex * phases
+ x_embed = (
+ torch.view_as_real(x_rotated).reshape(*x_rotated.shape[:-1], -1).to(x.dtype)
+ )
+ return x_embed
+
+ def forward(self, indices: torch.Tensor) -> torch.Tensor:
+ assert indices.shape[-1] == self.dim, f"Last dim of indices must be {self.dim}"
+ phases = self._get_phases(indices.reshape(-1)).reshape(*indices.shape[:-1], -1)
+ if phases.shape[-1] < self.head_dim // 2:
+ padn = self.head_dim // 2 - phases.shape[-1]
+ phases = torch.cat(
+ [
+ phases,
+ torch.polar(
+ torch.ones(*phases.shape[:-1], padn, device=phases.device),
+ torch.zeros(*phases.shape[:-1], padn, device=phases.device),
+ ),
+ ],
+ dim=-1,
+ )
+ return phases
+
+
+class TimestepEmbedder(nn.Module):
+ """
+ Embeds scalar timesteps into vector representations.
+ """
+
+ def __init__(self, hidden_size, frequency_embedding_size=256):
+ super().__init__()
+ self.mlp = nn.Sequential(
+ nn.Linear(frequency_embedding_size, hidden_size, bias=True),
+ nn.SiLU(),
+ nn.Linear(hidden_size, hidden_size, bias=True),
+ )
+ self.frequency_embedding_size = frequency_embedding_size
+
+ @staticmethod
+ def timestep_embedding(t, dim, max_period=10000):
+ """
+ Create sinusoidal timestep embeddings.
+
+ Args:
+ t: a 1-D Tensor of N indices, one per batch element.
+ These may be fractional.
+ dim: the dimension of the output.
+ max_period: controls the minimum frequency of the embeddings.
+
+ Returns:
+ an (N, D) Tensor of positional embeddings.
+ """
+ # https://github.com/openai/glide-text2im/blob/main/glide_text2im/nn.py
+ half = dim // 2
+ freqs = torch.exp(
+ -np.log(max_period)
+ * torch.arange(start=0, end=half, dtype=torch.float32)
+ / half
+ ).to(device=t.device)
+ args = t[:, None].float() * freqs[None]
+ embedding = torch.cat([torch.cos(args), torch.sin(args)], dim=-1)
+ if dim % 2:
+ embedding = torch.cat(
+ [embedding, torch.zeros_like(embedding[:, :1])], dim=-1
+ )
+ return embedding
+
+ def forward(self, t):
+ t_freq = self.timestep_embedding(t, self.frequency_embedding_size)
+ t_emb = self.mlp(t_freq)
+ return t_emb
+
+
+class PointEmbed(nn.Module):
+ def __init__(self, hidden_dim=48, dim=128):
+ super().__init__()
+
+ assert hidden_dim % 6 == 0
+
+ self.embedding_dim = hidden_dim
+ e = torch.pow(2, torch.arange(self.embedding_dim // 6)).float() * np.pi
+ e = torch.stack(
+ [
+ torch.cat(
+ [
+ e,
+ torch.zeros(self.embedding_dim // 6),
+ torch.zeros(self.embedding_dim // 6),
+ ]
+ ),
+ torch.cat(
+ [
+ torch.zeros(self.embedding_dim // 6),
+ e,
+ torch.zeros(self.embedding_dim // 6),
+ ]
+ ),
+ torch.cat(
+ [
+ torch.zeros(self.embedding_dim // 6),
+ torch.zeros(self.embedding_dim // 6),
+ e,
+ ]
+ ),
+ ]
+ )
+ self.register_buffer("basis", e) # 3 x 16
+
+ self.mlp = nn.Linear(self.embedding_dim + 3, dim)
+
+ @staticmethod
+ def embed(input, basis):
+ projections = torch.einsum("bnd,de->bne", input, basis)
+ embeddings = torch.cat([projections.sin(), projections.cos()], dim=2)
+ return embeddings
+
+ def forward(self, input):
+ dt = self.mlp.weight.dtype
+ if input.dtype != dt:
+ input = input.to(dtype=dt)
+ basis = self.basis.to(dtype=dt)
+ embed = self.mlp(torch.cat([self.embed(input, basis), input], dim=2))
+ return embed
+
+
+class MaskedTransformerCrossAttnBlock(nn.Module):
+ def __init__(self, hidden_size: int, num_heads: int, cond_dim: int):
+ super().__init__()
+ self.hidden_size = hidden_size
+ self.num_heads = num_heads
+ self.head_dim = hidden_size // num_heads
+
+ self.norm = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6)
+ self.q_cross = nn.Linear(hidden_size, hidden_size, bias=True)
+ self.kv_cross = nn.Linear(cond_dim, hidden_size * 2, bias=True)
+ self.proj_out_cross = nn.Linear(hidden_size, hidden_size, bias=True)
+ self.scale_cross = nn.Parameter(torch.zeros(hidden_size))
+
+ def forward(
+ self,
+ x: torch.Tensor,
+ c_tokens: torch.Tensor,
+ x_mask: Optional[torch.Tensor] = None,
+ c_mask: Optional[torch.Tensor] = None,
+ ) -> torch.Tensor:
+ b, n, d = x.shape
+ q_c = (
+ self.q_cross(self.norm(x))
+ .view(b, n, self.num_heads, self.head_dim)
+ .transpose(1, 2)
+ )
+ kv_c = (
+ self.kv_cross(c_tokens)
+ .view(b, c_tokens.shape[1], 2, self.num_heads, self.head_dim)
+ .permute(2, 0, 3, 1, 4)
+ )
+ k_c, v_c = kv_c[0], kv_c[1]
+ cross_attn_mask = c_mask.view(b, 1, 1, -1) if c_mask is not None else None
+ cross_out = F.scaled_dot_product_attention(
+ q_c,
+ k_c,
+ v_c,
+ attn_mask=cross_attn_mask,
+ )
+ cross_out = cross_out.transpose(1, 2).reshape(b, n, d)
+ x = x + self.scale_cross * self.proj_out_cross(cross_out)
+ if x_mask is not None:
+ x = torch.where(x_mask.unsqueeze(-1), x, torch.zeros_like(x))
+ return x
diff --git a/modules/transformer/hybrid.py b/modules/transformer/hybrid.py
new file mode 100644
index 0000000000000000000000000000000000000000..90b25add600f1ff727cf407cf50e0dcc00af56cd
--- /dev/null
+++ b/modules/transformer/hybrid.py
@@ -0,0 +1,236 @@
+from __future__ import annotations
+
+from typing import Optional
+import torch
+import torch.nn as nn
+import torch.nn.functional as F
+from torch.utils.checkpoint import checkpoint
+
+from ..attention import (
+ can_flash_varlen,
+ flash_varlen_self_attention,
+ graph_adj_varlen_attention,
+ sdpa_padding_mask,
+)
+from .blocks import RotaryPositionPhasesEmbedder
+
+
+class GraphAttnVarlenBlock(nn.Module):
+ def __init__(
+ self, hidden_size: int, num_heads: int, gradient_checkpointing: bool = False
+ ):
+ super().__init__()
+ self.hidden_size = hidden_size
+ self.num_heads = num_heads
+ self.head_dim = hidden_size // num_heads
+
+ self.norm1 = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6)
+ self.norm2 = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6)
+ self.qkv = nn.Linear(hidden_size, hidden_size * 3, bias=True)
+ self.proj_out = nn.Linear(hidden_size, hidden_size, bias=True)
+ self.ffn = nn.Sequential(
+ nn.Linear(hidden_size, hidden_size * 4, bias=True),
+ nn.GELU(approximate="tanh"),
+ nn.Linear(hidden_size * 4, hidden_size, bias=True),
+ )
+ self.scale_msa = nn.Parameter(torch.zeros(hidden_size))
+ self.scale_mlp = nn.Parameter(torch.zeros(hidden_size))
+ self.gradient_checkpointing = bool(gradient_checkpointing)
+
+ def _forward_once(
+ self,
+ x: torch.Tensor,
+ x_mask: Optional[torch.Tensor],
+ adj_matrix: Optional[torch.Tensor],
+ rope_phases: Optional[torch.Tensor],
+ ) -> torch.Tensor:
+ B, N, D = x.shape
+ qkv = (
+ self.qkv(self.norm1(x))
+ .view(B, N, 3, self.num_heads, self.head_dim)
+ .permute(2, 0, 3, 1, 4)
+ )
+ q, k, v = qkv[0], qkv[1], qkv[2]
+
+ if rope_phases is not None:
+ q = RotaryPositionPhasesEmbedder.apply_rotary_embedding(q, rope_phases)
+ k = RotaryPositionPhasesEmbedder.apply_rotary_embedding(k, rope_phases)
+
+ if x_mask is None:
+ attn_mask = None
+ if adj_matrix is not None:
+ adj_mask = adj_matrix.bool()
+ eye = torch.eye(N, dtype=torch.bool, device=x.device).unsqueeze(0)
+ attn_mask = (adj_mask | eye).unsqueeze(1)
+ attn_out = F.scaled_dot_product_attention(q, k, v, attn_mask=attn_mask)
+ else:
+ attn_out = graph_adj_varlen_attention(q, k, v, x_mask, adj_matrix)
+
+ attn_out = attn_out.transpose(1, 2).reshape(B, N, D)
+ x = x + self.scale_msa * self.proj_out(attn_out)
+ x = x + self.scale_mlp * self.ffn(self.norm2(x))
+ if x_mask is not None:
+ x = torch.where(x_mask.unsqueeze(-1), x, torch.zeros_like(x))
+ return x
+
+ def forward(
+ self,
+ x: torch.Tensor,
+ x_mask: Optional[torch.Tensor],
+ adj_matrix: Optional[torch.Tensor],
+ rope_phases: Optional[torch.Tensor] = None,
+ ) -> torch.Tensor:
+ if self.training and self.gradient_checkpointing:
+ return checkpoint(
+ self._forward_once,
+ x,
+ x_mask,
+ adj_matrix,
+ rope_phases,
+ use_reentrant=False,
+ )
+ return self._forward_once(x, x_mask, adj_matrix, rope_phases)
+
+
+class FlashVarlenTransformerBlock(nn.Module):
+ def __init__(
+ self, hidden_size: int, num_heads: int, gradient_checkpointing: bool = False
+ ):
+ super().__init__()
+ self.hidden_size = hidden_size
+ self.num_heads = num_heads
+ self.head_dim = hidden_size // num_heads
+
+ self.norm1 = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6)
+ self.norm2 = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6)
+ self.qkv = nn.Linear(hidden_size, hidden_size * 3, bias=True)
+ self.proj_out = nn.Linear(hidden_size, hidden_size, bias=True)
+ self.ffn = nn.Sequential(
+ nn.Linear(hidden_size, hidden_size * 4, bias=True),
+ nn.GELU(approximate="tanh"),
+ nn.Linear(hidden_size * 4, hidden_size, bias=True),
+ )
+ self.scale_msa = nn.Parameter(torch.zeros(hidden_size))
+ self.scale_mlp = nn.Parameter(torch.zeros(hidden_size))
+ self.gradient_checkpointing = bool(gradient_checkpointing)
+
+ def _forward_once(
+ self,
+ x: torch.Tensor,
+ x_mask: Optional[torch.Tensor],
+ rope_phases: Optional[torch.Tensor],
+ ) -> torch.Tensor:
+ B, N, D = x.shape
+ qkv = (
+ self.qkv(self.norm1(x))
+ .view(B, N, 3, self.num_heads, self.head_dim)
+ .permute(2, 0, 3, 1, 4)
+ )
+ q, k, v = qkv[0], qkv[1], qkv[2]
+
+ if rope_phases is not None:
+ q = RotaryPositionPhasesEmbedder.apply_rotary_embedding(q, rope_phases)
+ k = RotaryPositionPhasesEmbedder.apply_rotary_embedding(k, rope_phases)
+
+ if can_flash_varlen(q, x_mask):
+ attn_out = flash_varlen_self_attention(q, k, v, x_mask)
+ elif x_mask is not None:
+ pad_mask = sdpa_padding_mask(x_mask)
+ attn_out = F.scaled_dot_product_attention(q, k, v, attn_mask=pad_mask)
+ else:
+ attn_out = F.scaled_dot_product_attention(q, k, v, attn_mask=None)
+
+ attn_out = attn_out.transpose(1, 2).reshape(B, N, D)
+ x = x + self.scale_msa * self.proj_out(attn_out)
+ x = x + self.scale_mlp * self.ffn(self.norm2(x))
+ if x_mask is not None:
+ x = torch.where(x_mask.unsqueeze(-1), x, torch.zeros_like(x))
+ return x
+
+ def forward(
+ self,
+ x: torch.Tensor,
+ x_mask: Optional[torch.Tensor],
+ rope_phases: Optional[torch.Tensor] = None,
+ ) -> torch.Tensor:
+ if self.training and self.gradient_checkpointing:
+ return checkpoint(
+ self._forward_once,
+ x,
+ x_mask,
+ rope_phases,
+ use_reentrant=False,
+ )
+ return self._forward_once(x, x_mask, rope_phases)
+
+
+class HybridGraphFlashStage(nn.Module):
+ def __init__(
+ self,
+ hidden_size: int,
+ num_heads: int,
+ num_flash: int,
+ gradient_checkpointing: bool = False,
+ ):
+ super().__init__()
+ self.graph_block = GraphAttnVarlenBlock(
+ hidden_size, num_heads, gradient_checkpointing=gradient_checkpointing
+ )
+ self.flash_blocks = nn.ModuleList(
+ [
+ FlashVarlenTransformerBlock(
+ hidden_size,
+ num_heads,
+ gradient_checkpointing=gradient_checkpointing,
+ )
+ for _ in range(num_flash)
+ ]
+ )
+
+ def forward(
+ self,
+ x: torch.Tensor,
+ x_mask: Optional[torch.Tensor],
+ adj_matrix: Optional[torch.Tensor],
+ rope_phases: Optional[torch.Tensor],
+ ) -> torch.Tensor:
+ x = self.graph_block(
+ x, x_mask=x_mask, adj_matrix=adj_matrix, rope_phases=rope_phases
+ )
+ for fb in self.flash_blocks:
+ x = fb(x, x_mask=x_mask, rope_phases=rope_phases)
+ return x
+
+
+class HybridGraphFlashStack(nn.Module):
+ def __init__(
+ self,
+ hidden_size: int,
+ num_heads: int,
+ num_stages: int,
+ num_flash_per_stage: int,
+ gradient_checkpointing: bool = False,
+ ):
+ super().__init__()
+ self.stages = nn.ModuleList(
+ [
+ HybridGraphFlashStage(
+ hidden_size,
+ num_heads,
+ num_flash_per_stage,
+ gradient_checkpointing=gradient_checkpointing,
+ )
+ for _ in range(num_stages)
+ ]
+ )
+
+ def forward(
+ self,
+ x: torch.Tensor,
+ x_mask: Optional[torch.Tensor],
+ adj_matrix: Optional[torch.Tensor],
+ rope_phases: Optional[torch.Tensor],
+ ) -> torch.Tensor:
+ for stage in self.stages:
+ x = stage(x, x_mask=x_mask, adj_matrix=adj_matrix, rope_phases=rope_phases)
+ return x
diff --git a/modules/utils.py b/modules/utils.py
new file mode 100644
index 0000000000000000000000000000000000000000..1baaab2318b4afbabb0b07e9fbe8f19b137cba3b
--- /dev/null
+++ b/modules/utils.py
@@ -0,0 +1,145 @@
+import torch
+import torch.nn as nn
+from typing import *
+import numpy as np
+from modules import sparse as sp
+
+FP16_MODULES = (
+ nn.Conv1d,
+ nn.Conv2d,
+ nn.Conv3d,
+ nn.ConvTranspose1d,
+ nn.ConvTranspose2d,
+ nn.ConvTranspose3d,
+ nn.Linear,
+ sp.SparseConv3d,
+ sp.SparseInverseConv3d,
+ sp.SparseLinear,
+)
+
+
+def convert_module_to_f16(l):
+ """
+ Convert primitive modules to float16.
+ """
+ if isinstance(l, FP16_MODULES):
+ for p in l.parameters():
+ p.data = p.data.half()
+
+
+def convert_module_to_f32(l):
+ """
+ Convert primitive modules to float32, undoing convert_module_to_f16().
+ """
+ if isinstance(l, FP16_MODULES):
+ for p in l.parameters():
+ p.data = p.data.float()
+
+
+def zero_module(module):
+ """
+ Zero out the parameters of a module and return it.
+ """
+ for p in module.parameters():
+ p.detach().zero_()
+ return module
+
+
+def modulate(x, shift, scale):
+ return x * (1 + scale.unsqueeze(1)) + shift.unsqueeze(1)
+
+
+class DiagonalGaussianDistribution(object):
+ def __init__(
+ self,
+ parameters: Union[torch.Tensor, List[torch.Tensor]],
+ deterministic=False,
+ feat_dim=1,
+ ):
+ self.feat_dim = feat_dim
+ self.parameters = parameters
+
+ if isinstance(parameters, list):
+ self.mean = parameters[0]
+ self.logvar = parameters[1]
+ else:
+ self.mean, self.logvar = torch.chunk(parameters, 2, dim=feat_dim)
+ self.logvar = torch.clamp(self.logvar, -30.0, 20.0)
+ self.deterministic = deterministic
+ self.std = torch.exp(0.5 * self.logvar)
+ self.var = torch.exp(self.logvar)
+ if self.deterministic:
+ self.var = self.std = torch.zeros_like(self.mean)
+
+ def sample(self):
+ x = self.mean + self.std * torch.randn_like(self.mean)
+ return x
+
+ def kl(self, other=None, dims=(1, 2, 3)):
+ if self.deterministic:
+ return torch.Tensor([0.0])
+ else:
+ if other is None:
+ return 0.5 * torch.mean(
+ torch.pow(self.mean, 2) + self.var - 1.0 - self.logvar, dim=dims
+ )
+ else:
+ return 0.5 * torch.mean(
+ torch.pow(self.mean - other.mean, 2) / other.var
+ + self.var / other.var
+ - 1.0
+ - self.logvar
+ + other.logvar,
+ dim=dims,
+ )
+
+ def nll(self, sample, dims=(1, 2, 3)):
+ if self.deterministic:
+ return torch.Tensor([0.0])
+ logtwopi = np.log(2.0 * np.pi)
+ return 0.5 * torch.sum(
+ logtwopi + self.logvar + torch.pow(sample - self.mean, 2) / self.var,
+ dim=dims,
+ )
+
+ def mode(self):
+ return self.mean
+
+
+def per_batch_counts(batch_indices: torch.Tensor, num_batches: int) -> List[int]:
+ """Count elements per batch, returned as a list of length num_batches."""
+ return torch.bincount(batch_indices.long(), minlength=num_batches).tolist()
+
+
+def flatten_coords(coords_4d: torch.Tensor):
+ coords_4d_long = coords_4d.long()
+
+ base_x = 1024
+ base_y = 1024 * 1024
+ base_z = 1024 * 1024 * 1024
+
+ flat_coords = (
+ coords_4d_long[:, 0] * base_z
+ + coords_4d_long[:, 1] * base_y
+ + coords_4d_long[:, 2] * base_x
+ + coords_4d_long[:, 3]
+ )
+ return flat_coords
+
+def manual_cast(tensor, dtype):
+ if not torch.is_autocast_enabled():
+ return tensor.type(dtype)
+ return tensor
+
+
+def str_to_dtype(dtype_str: str):
+ return {
+ "f16": torch.float16,
+ "fp16": torch.float16,
+ "float16": torch.float16,
+ "bf16": torch.bfloat16,
+ "bfloat16": torch.bfloat16,
+ "f32": torch.float32,
+ "fp32": torch.float32,
+ "float32": torch.float32,
+ }[dtype_str]
\ No newline at end of file
diff --git a/requirements.txt b/requirements.txt
new file mode 100644
index 0000000000000000000000000000000000000000..4e17eace2776aa0a4ec51f08b1f561ee9bd2bf59
--- /dev/null
+++ b/requirements.txt
@@ -0,0 +1,23 @@
+# ──────────────────────────────────────────────────────────────────────
+# LATO.2 Gradio App — requirements.txt
+# ──────────────────────────────────────────────────────────────────────
+# Install AFTER running `setup.sh --all` (which sets up the base
+# conda env with PyTorch, spconv, flash-attn, o_voxel, etc.)
+#
+# pip install -r requirements.txt
+# ──────────────────────────────────────────────────────────────────────
+
+# Gradio app framework
+gradio>=4.44.0
+gradio_rerun>=0.0.4
+
+# Rerun 3D viewer SDK
+rerun-sdk>=0.22.0
+
+# ── Already in setup.sh but listed for completeness ──────────────────
+numpy
+trimesh
+tqdm
+pillow
+huggingface_hub
+open3d==0.19.0
diff --git a/scripts/ckpt_download.py b/scripts/ckpt_download.py
new file mode 100644
index 0000000000000000000000000000000000000000..474a84fe41b24888e68df9a2ae3fb009b9f082f9
--- /dev/null
+++ b/scripts/ckpt_download.py
@@ -0,0 +1,90 @@
+"""
+Usage:
+ python scripts/ckpt_download.py \
+ [--out_dir ]
+"""
+
+import argparse
+import os
+import sys
+
+ROOT = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
+sys.path.insert(0, ROOT)
+
+from utils import logging
+
+DEFAULT_REPO_ID = "0x4c48/LATO.2"
+DEFAULT_OUT_DIR = os.path.join(ROOT, "ckpt")
+
+
+def parse_args():
+ p = argparse.ArgumentParser(
+ description="Download LATO.2 checkpoints from the Hugging Face Hub.",
+ formatter_class=argparse.ArgumentDefaultsHelpFormatter,
+ )
+ p.add_argument("--repo_id", default=DEFAULT_REPO_ID, help="HF model repo id")
+ p.add_argument(
+ "--out_dir",
+ default=DEFAULT_OUT_DIR,
+ help="Directory to download checkpoints into",
+ )
+ p.add_argument("--revision", default=None, help="Git revision / branch / tag")
+ p.add_argument(
+ "--token",
+ default=os.environ.get("HF_TOKEN"),
+ help="HF access token for gated/private repos (or set HF_TOKEN)",
+ )
+ p.add_argument(
+ "--all",
+ action="store_true",
+ help="Download every file in the repo (not just *.pt)",
+ )
+ p.add_argument(
+ "--include-readme",
+ action="store_true",
+ help="Also download README.md alongside the *.pt weights",
+ )
+ return p.parse_args()
+
+
+def main():
+ args = parse_args()
+
+ try:
+ from huggingface_hub import snapshot_download
+ except ImportError:
+ sys.exit(
+ "huggingface_hub is not installed. Activate the `trellis2` conda env "
+ "or run: pip install -U huggingface_hub"
+ )
+
+ if args.all:
+ allow_patterns = None
+ else:
+ allow_patterns = ["*.pt"]
+ if args.include_readme:
+ allow_patterns.append("README.md")
+
+ os.makedirs(args.out_dir, exist_ok=True)
+ logging.info(f"Downloading {args.repo_id} -> {args.out_dir}")
+ if allow_patterns:
+ logging.info(f" patterns: {allow_patterns}")
+
+ path = snapshot_download(
+ repo_id=args.repo_id,
+ repo_type="model",
+ revision=args.revision,
+ local_dir=args.out_dir,
+ allow_patterns=allow_patterns,
+ token=args.token,
+ )
+
+ logging.info(f"\nDone. Checkpoints available in: {path}")
+ files = sorted(f for f in os.listdir(path) if os.path.isfile(os.path.join(path, f)))
+ for f in files:
+ size = os.path.getsize(os.path.join(path, f)) / (1024 * 1024)
+ logging.info(f" {f:24s} {size:8.1f} MB")
+
+
+if __name__ == "__main__":
+ main()
diff --git a/scripts/e2e_inference.py b/scripts/e2e_inference.py
new file mode 100644
index 0000000000000000000000000000000000000000..a0e3c23a82d9068cf7d690c3dee6cee8c65e5d47
--- /dev/null
+++ b/scripts/e2e_inference.py
@@ -0,0 +1,380 @@
+"""
+Outputs into --out_dir:
+ _pred.ply generated vertices (+offset head) in [-0.5, 0.5]
+ _pred_coords.ply generated vertex voxel coords in [0, 1024)
+ _pred.obj generated mesh (offset vertices, faces) in [-0.5, 0.5]
+ _pred_coords.obj generated mesh on integer voxel coords in [0, 1024)
+ _render.png the conditioning view fed to DINO-v2
+
+Usage:
+ python scripts/e2e_inference.py --mesh_dir --out_dir outputs/e2e_run/ \
+ [--vert_num 2000] [--cfg_strength 3.0] [--vflow_steps 24] [--tflow_steps 50] \
+ [--render_azimuth 45 --render_elevation 30] [--no-fill_quad_rings]
+"""
+
+import argparse
+import os
+import sys
+import time
+from collections import Counter
+from functools import partial
+
+import tqdm
+
+ROOT = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
+sys.path.insert(0, ROOT)
+os.environ.setdefault("PYTORCH_CUDA_ALLOC_CONF", "expandable_segments:True")
+os.environ.setdefault("XDG_RUNTIME_DIR", "/tmp/runtime-root")
+os.makedirs(os.environ["XDG_RUNTIME_DIR"], exist_ok=True)
+# Open3D headless rendering: without this the default EGL platform can hang
+# (e.g. when every GPU is busy or no display device is exposed).
+os.environ.setdefault("EGL_PLATFORM", "surfaceless")
+
+import numpy as np
+import torch
+import trimesh
+from PIL import Image
+from torch.utils.data import DataLoader
+
+from dataset.voxel_dataset import VoxelVertexDataset, collate_fn
+from models import (
+ DinoV2Encoder,
+ OffsetHead,
+ TopoFlowEulerSampler,
+ TopologySiTFlow,
+ TopologyVAE,
+ VertexSLatFlowModel,
+ VertFlowEulerCfgSampler,
+ VertexVAE,
+ VoxelFieldConditioner,
+)
+from modules.sparse import SparseTensor
+import utils.logging as logging
+from utils.export import export_vertex
+from utils.inference import (
+ build_voxel_fields,
+ compute_density,
+ decode_vertices,
+ edges_to_faces,
+ pad_verts,
+ worker_init,
+)
+from utils.load import load_latov2_model
+
+
+def parse_args():
+ p = argparse.ArgumentParser(
+ description="end-to-end vertex + topology generation inference"
+ )
+ p.add_argument("--mesh_dir", required=True, help="directory of input meshes")
+ p.add_argument(
+ "--out_dir", required=True, help="output directory for the PLYs / OBJs"
+ )
+ p.add_argument("--vflow_ckpt", default=os.path.join(ROOT, "ckpt", "vflow.pt"))
+ p.add_argument("--vvae_ckpt", default=os.path.join(ROOT, "ckpt", "vvae.pt"))
+ p.add_argument(
+ "--offset_head_ckpt", default=os.path.join(ROOT, "ckpt", "offset_head.pt")
+ )
+ p.add_argument("--tflow_ckpt", default=os.path.join(ROOT, "ckpt", "tflow.pt"))
+ p.add_argument("--tvae_ckpt", default=os.path.join(ROOT, "ckpt", "tvae.pt"))
+ p.add_argument(
+ "--voxel_encoder_ckpt",
+ default=os.path.join(ROOT, "ckpt", "voxel_encoder.pt"),
+ )
+ p.add_argument("--batch_size", type=int, default=1)
+ p.add_argument(
+ "--num_samples", type=int, default=None, help="only run the first N meshes"
+ )
+ p.add_argument("--num_workers", type=int, default=4)
+ p.add_argument("--inference_threshold", type=float, default=0.5)
+ p.add_argument("--seed", type=int, default=42)
+ # vertex flow sampling
+ p.add_argument("--vflow_steps", type=int, default=24, help="V-Flow Euler steps")
+ p.add_argument("--cfg_strength", type=float, default=3.0)
+ p.add_argument("--rescale_t", type=float, default=1.0)
+ # vertex-count density conditioning
+ p.add_argument("--vert_num", type=int, default=2000, help="target vertex count")
+ p.add_argument(
+ "--use_gt_vert_count",
+ action=argparse.BooleanOptionalAction,
+ default=False,
+ help="condition on the GT quantized vertex count instead of --vert_num",
+ )
+ p.add_argument(
+ "--scaler",
+ type=float,
+ default=1.0,
+ help="multiplier on the GT count when --use_gt_vert_count",
+ )
+ p.add_argument("--min_verts", type=float, default=200.0)
+ p.add_argument("--max_verts", type=float, default=5000.0)
+ # topology flow sampling / decoding
+ p.add_argument("--tflow_steps", type=int, default=50, help="T-Flow Euler steps")
+ p.add_argument("--edge_threshold", type=float, default=0.0)
+ p.add_argument("--chunk_size", type=int, default=20000)
+ p.add_argument(
+ "--fill_quad_rings",
+ action=argparse.BooleanOptionalAction,
+ default=True,
+ help=(
+ "post-process: split chordless 4-vertex rings into two triangles "
+ "(pure topology, not the voxel support filter)"
+ ),
+ )
+ # conditioning render
+ p.add_argument("--render_azimuth", type=float, default=45.0)
+ p.add_argument("--render_elevation", type=float, default=30.0)
+ p.add_argument("--img_res", type=int, default=518)
+ p.add_argument(
+ "--dino_hub_dir",
+ default=os.path.join(ROOT, "ckpt", "dinov2"),
+ help="torch.hub cache for DINO-v2; reused when present, downloaded otherwise",
+ )
+ args = p.parse_args()
+ if args.num_samples is not None and args.num_samples <= 0:
+ args.num_samples = None # <= 0 means "all", not python slice semantics
+ return args
+
+
+def export_mesh(out_dir, base_name, vert_int, vert_offsets, faces, resolution):
+ """OBJ pair matching export_vertex's PLY conventions (offset verts / int coords)."""
+ res = float(resolution)
+ vert_with_offset = (
+ vert_int.astype(np.float64) / res
+ - 0.5
+ + vert_offsets.astype(np.float64) / (res * 2.0)
+ )
+ trimesh.Trimesh(vertices=vert_with_offset, faces=faces).export(
+ os.path.join(out_dir, f"{base_name}_pred.obj")
+ )
+ trimesh.Trimesh(vertices=vert_int.astype(np.float64), faces=faces).export(
+ os.path.join(out_dir, f"{base_name}_pred_coords.obj")
+ )
+
+
+def main():
+ logging.info("End-to-end inference starting...")
+
+ args = parse_args()
+ device = torch.device("cuda")
+ torch.manual_seed(args.seed)
+ np.random.seed(args.seed)
+ os.makedirs(args.out_dir, exist_ok=True)
+
+ # stage 1: vertex generation
+ vflow, vflow_cfg = load_latov2_model(VertexSLatFlowModel, args.vflow_ckpt, device)
+ vvae, vvae_cfg = load_latov2_model(VertexVAE, args.vvae_ckpt, device)
+ offset_head, _ = load_latov2_model(OffsetHead, args.offset_head_ckpt, device)
+ # stage 2: topology generation
+ tflow, tflow_cfg = load_latov2_model(TopologySiTFlow, args.tflow_ckpt, device)
+ tvae, _ = load_latov2_model(TopologyVAE, args.tvae_ckpt, device)
+ voxel_encoder, venc_cfg = load_latov2_model(
+ VoxelFieldConditioner, args.voxel_encoder_ckpt, device
+ )
+
+ res = vvae_cfg["resolution"]
+ min_res = vvae_cfg["min_resolution"]
+ latent_dim = vflow_cfg["latent_dim"]
+ density_max = vflow_cfg["max_vertex_num"]
+ z_dim = int(tflow_cfg["args"]["z_dim"])
+ num_discrete = int(tflow_cfg["args"]["num_discrete"])
+ max_vertices = int(tflow_cfg["args"]["max_vertices"])
+ latent_scale = float(tflow_cfg["latent_scale"])
+ voxel_res = int(venc_cfg["resolution"])
+ if num_discrete != res:
+ raise ValueError(
+ f"T-Flow num_discrete={num_discrete} != V-VAE resolution={res}; "
+ "the generated vertex voxels would be in the wrong coordinate space."
+ )
+ if voxel_res != min_res:
+ raise ValueError(
+ f"voxel encoder resolution={voxel_res} != V-VAE min_resolution={min_res}; "
+ "both stages must share the same active-voxel conditioning grid."
+ )
+
+ dino = (
+ DinoV2Encoder(
+ model_name=vflow_cfg["dino_version"],
+ hub_dir=args.dino_hub_dir,
+ img_res=vflow_cfg["image_resolution"],
+ )
+ .to(device)
+ .eval()
+ )
+ logging.info(f"loaded {vflow_cfg['dino_version']} from {args.dino_hub_dir}")
+ vertex_sampler = VertFlowEulerCfgSampler()
+ topo_sampler = TopoFlowEulerSampler()
+
+ dataset = VoxelVertexDataset(
+ root_dir=args.mesh_dir,
+ resolution=res,
+ min_resolution=min_res,
+ need_encoder_inputs=False,
+ num_samples=args.num_samples,
+ render=True,
+ img_res=args.img_res,
+ render_azimuth=args.render_azimuth,
+ render_elevation=args.render_elevation,
+ )
+ loader = DataLoader(
+ dataset,
+ batch_size=args.batch_size,
+ shuffle=False,
+ collate_fn=partial(collate_fn, resolution=res, min_resolution=min_res),
+ num_workers=args.num_workers,
+ pin_memory=True,
+ # EGL rendering hangs inside fork-ed children of a CUDA-initialized
+ # parent; spawn gives each worker a clean process for its EGL context.
+ multiprocessing_context="spawn" if args.num_workers > 0 else None,
+ worker_init_fn=worker_init if args.num_workers > 0 else None,
+ )
+ dupes = sorted(
+ s
+ for s, c in Counter(os.path.splitext(f)[0] for f in dataset.files).items()
+ if c > 1
+ )
+ if dupes:
+ logging.warning(
+ f"WARNING: {len(dupes)} duplicate mesh basename(s) — later samples will overwrite earlier outputs."
+ )
+ logging.info(
+ f"{len(dataset)} meshes from {args.mesh_dir} "
+ f"(vflow_steps={args.vflow_steps}, cfg={args.cfg_strength}, "
+ f"vert_num={args.vert_num}, use_gt_vert_count={args.use_gt_vert_count}, "
+ f"scaler={args.scaler}, density_max={density_max}, tflow_steps={args.tflow_steps}, "
+ f"view=az{args.render_azimuth}/el{args.render_elevation}, seed={args.seed}) "
+ f"-> {args.out_dir}"
+ )
+
+ n_ok = n_no_topo = n_fail = 0
+ t_start = time.time()
+ qbar = tqdm.tqdm(loader, desc="inference", unit="batch", dynamic_ncols=True)
+ for batch in qbar:
+ for err in batch["errors"]:
+ n_fail += 1
+ logging.error(
+ f"{err['name']}: FAILED during preprocessing: {err['error'].splitlines()[0]}"
+ )
+ if "name" not in batch:
+ continue
+
+ density = compute_density(batch, args, density_max, device)
+ with torch.no_grad():
+ cond = dino(np.stack(batch["image"])).float()
+ neg_cond = torch.zeros_like(cond)
+
+ # ---- stage 1: V-Flow on the 64^3 active voxels -> V-VAE vertex decode ----
+ min_active = batch[f"active_voxels_{min_res}"]
+ with torch.no_grad(), torch.autocast("cuda", dtype=torch.bfloat16):
+ min_active_coords = min_active.to(device)
+ noise = SparseTensor(
+ coords=min_active_coords.int(),
+ feats=torch.randn(
+ min_active_coords.shape[0], latent_dim, device=device
+ ),
+ )
+ z_pred = vertex_sampler.sample(
+ model=vflow,
+ noise=noise,
+ cond=cond,
+ neg_cond=neg_cond,
+ steps=args.vflow_steps,
+ cfg_strength=args.cfg_strength,
+ rescale_t=args.rescale_t,
+ density=density,
+ )
+ pred_coords, pred_offsets = decode_vertices(
+ vvae, offset_head, z_pred, args.inference_threshold
+ )
+
+ keep_idx, verts_list, offsets_list = [], [], []
+ for b, name in enumerate(batch["name"]):
+ pred_sel = pred_coords[:, 0] == b
+ vert_int = pred_coords[pred_sel, 1:].long()
+ vert_off = pred_offsets[pred_sel]
+ export_vertex(
+ args.out_dir,
+ name,
+ type_name="pred",
+ vert_int=vert_int.numpy(),
+ vert_offsets=vert_off.numpy(),
+ resolution=res,
+ )
+ Image.fromarray(batch["image"][b]).save(
+ os.path.join(args.out_dir, f"{name}_render.png")
+ )
+ num_pred = int(vert_int.shape[0])
+ if num_pred < 3:
+ n_no_topo += 1
+ logging.warning(
+ f"{name}: only {num_pred} generated vertices; skipping topology."
+ )
+ elif num_pred > max_vertices:
+ n_no_topo += 1
+ logging.warning(
+ f"{name}: {num_pred} generated vertices exceed T-Flow "
+ f"max_vertices={max_vertices}; skipping topology."
+ )
+ else:
+ keep_idx.append(b)
+ verts_list.append(vert_int)
+ offsets_list.append(vert_off)
+ if not keep_idx:
+ continue
+
+ # ---- stage 2: T-Flow on the generated vertices -> T-VAE edge decode ----
+ with torch.no_grad():
+ verts, mask, lengths = pad_verts(verts_list, device)
+ voxel_list = [
+ min_active[min_active[:, 0] == b, 1:].long() for b in keep_idx
+ ]
+ field = build_voxel_fields(voxel_list, voxel_res, device) # (B', R, R, R)
+ cond_vox = voxel_encoder(field) # (B', R'^3, cond_in_dim)
+
+ z0 = torch.randn(verts.shape[0], verts.shape[1], z_dim, device=device)
+ z_flow = topo_sampler.sample(
+ model=tflow,
+ noise=z0,
+ verts=verts,
+ mask=mask,
+ cond=cond_vox,
+ steps=args.tflow_steps,
+ )
+ z = z_flow.float() / latent_scale
+
+ with torch.autocast("cuda", dtype=torch.bfloat16):
+ edges_list = tvae.decode(
+ z,
+ verts=verts,
+ verts_mask=mask,
+ chunk_size=args.chunk_size,
+ threshold=args.edge_threshold,
+ )
+
+ for k, b in enumerate(keep_idx):
+ name = batch["name"][b]
+ faces = edges_to_faces(edges_list[k], lengths[k], args.fill_quad_rings)
+ if faces.shape[0] == 0:
+ n_no_topo += 1
+ logging.warning(
+ f"{name}: no faces decoded; the _pred PLYs are the only outputs."
+ )
+ continue
+ export_mesh(
+ args.out_dir,
+ name,
+ vert_int=verts_list[k].numpy(),
+ vert_offsets=offsets_list[k].numpy(),
+ faces=faces,
+ resolution=res,
+ )
+ n_ok += 1
+
+ logging.info(
+ f"done: {n_ok} ok, {n_no_topo} without topology, {n_fail} failed "
+ f"in {time.time() - t_start:.0f}s -> {args.out_dir}"
+ )
+
+
+if __name__ == "__main__":
+ main()
diff --git a/scripts/tflow_inference.py b/scripts/tflow_inference.py
new file mode 100644
index 0000000000000000000000000000000000000000..16a89eb834f90c881cee9f0d2fa0616614b411ac
--- /dev/null
+++ b/scripts/tflow_inference.py
@@ -0,0 +1,221 @@
+"""
+Outputs into --out_dir:
+ _pred.obj generated mesh (known verts, faces) in [-0.5, 0.5]
+ _pred.ply fallback point cloud when no faces were generated
+ _known.ply the known (dequantized) vertices fed to the flow
+ _voxel_field.ply the active-voxel conditioning field (debug, --save_voxel_field)
+
+Usage:
+ python scripts/tflow_inference.py --mesh_dir --out_dir outputs/tflow_run/ \
+ [--steps 50] [--no-use_cond] [--no-fill_quad_rings]
+"""
+
+import argparse
+import os
+import sys
+import time
+
+import numpy as np
+import torch
+import trimesh
+import tqdm
+
+ROOT = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
+sys.path.insert(0, ROOT)
+os.environ.setdefault("PYTORCH_CUDA_ALLOC_CONF", "expandable_segments:True")
+
+from torch.utils.data import DataLoader
+
+from dataset.topo_dataset import TopoVoxelDataset, collate_fn
+from models import (
+ TopologyVAE,
+ TopologySiTFlow,
+ TopoFlowEulerSampler,
+ VoxelFieldConditioner,
+)
+import utils.logging as logging
+from utils.inference import build_voxel_fields, edges_to_faces, pad_verts
+from utils.load import load_latov2_model
+
+
+def parse_args():
+ p = argparse.ArgumentParser(description="T-Flow topology generation inference")
+ p.add_argument("--mesh_dir", required=True, help="directory of input meshes")
+ p.add_argument("--out_dir", required=True, help="output directory for the meshes")
+ p.add_argument("--tflow_ckpt", default=os.path.join(ROOT, "ckpt", "tflow.pt"))
+ p.add_argument("--tvae_ckpt", default=os.path.join(ROOT, "ckpt", "tvae.pt"))
+ p.add_argument(
+ "--voxel_encoder_ckpt",
+ default=os.path.join(ROOT, "ckpt", "voxel_encoder.pt"),
+ )
+ p.add_argument("--batch_size", type=int, default=1)
+ p.add_argument(
+ "--num_samples", type=int, default=None, help="only run the first N meshes"
+ )
+ p.add_argument("--num_workers", type=int, default=4)
+ p.add_argument("--seed", type=int, default=42)
+ # flow sampling
+ p.add_argument("--steps", type=int, default=50, help="Euler steps")
+ p.add_argument(
+ "--use_cond",
+ action=argparse.BooleanOptionalAction,
+ default=True,
+ help=(
+ "condition on the active-voxel field. --no-use_cond runs the flow "
+ "unconditionally (the model's learned null token)."
+ ),
+ )
+ # topology decoding
+ p.add_argument("--edge_threshold", type=float, default=0.0)
+ p.add_argument("--chunk_size", type=int, default=20000)
+ p.add_argument(
+ "--fill_quad_rings",
+ action=argparse.BooleanOptionalAction,
+ default=True,
+ help=(
+ "post-process: split chordless 4-vertex rings into two triangles "
+ "(pure topology, not the voxel support filter)"
+ ),
+ )
+ p.add_argument(
+ "--save_voxel_field",
+ action=argparse.BooleanOptionalAction,
+ default=True,
+ help="also dump the active-voxel conditioning field as a point cloud",
+ )
+ args = p.parse_args()
+ if args.num_samples is not None and args.num_samples <= 0:
+ args.num_samples = None # <= 0 means "all", not python slice semantics
+ return args
+
+
+def main():
+ logging.info("T-Flow inference starting...")
+
+ args = parse_args()
+ device = torch.device("cuda")
+ torch.manual_seed(args.seed)
+ np.random.seed(args.seed)
+ os.makedirs(args.out_dir, exist_ok=True)
+
+ tflow, tflow_cfg = load_latov2_model(TopologySiTFlow, args.tflow_ckpt, device)
+ tvae, _ = load_latov2_model(TopologyVAE, args.tvae_ckpt, device)
+ voxel_encoder, venc_cfg = load_latov2_model(
+ VoxelFieldConditioner, args.voxel_encoder_ckpt, device
+ )
+ z_dim = int(tflow_cfg["args"]["z_dim"])
+ num_discrete = int(tflow_cfg["args"]["num_discrete"])
+ max_vertices = int(tflow_cfg["args"]["max_vertices"])
+ latent_scale = float(tflow_cfg["latent_scale"])
+ voxel_res = int(venc_cfg["resolution"])
+ sampler = TopoFlowEulerSampler()
+
+ dataset = TopoVoxelDataset(
+ root_dir=args.mesh_dir,
+ num_discrete=num_discrete,
+ voxel_res=voxel_res,
+ max_vertices=max_vertices,
+ num_samples=args.num_samples,
+ )
+ loader = DataLoader(
+ dataset,
+ batch_size=args.batch_size,
+ shuffle=False,
+ collate_fn=collate_fn,
+ num_workers=args.num_workers,
+ pin_memory=True,
+ )
+ logging.info(
+ f"{len(dataset)} meshes from {args.mesh_dir} "
+ f"(steps={args.steps}, use_cond={args.use_cond}, num_discrete={num_discrete}, "
+ f"voxel_res={voxel_res}, latent_scale={latent_scale}, seed={args.seed}) "
+ f"-> {args.out_dir}"
+ )
+
+ n_ok = n_fail = 0
+ t_start = time.time()
+ qbar = tqdm.tqdm(loader, desc="inference", unit="batch", dynamic_ncols=True)
+ for batch in qbar:
+ for err in batch["errors"]:
+ n_fail += 1
+ logging.error(
+ f"{err['name']}: FAILED during preprocessing: {err['error'].splitlines()[0]}"
+ )
+ if "name" not in batch:
+ continue
+
+ names = batch["name"]
+ verts_list = batch["vertices"] # list of (N_i, 3) long in [0, num_discrete)
+ voxel_list = batch["voxel_coords"] # list of (M_i, 3) long in [0, voxel_res)
+
+ with torch.no_grad():
+ verts, mask, lengths = pad_verts(
+ verts_list, device
+ ) # (B, N_max, 3), (B, N_max)
+ if args.use_cond:
+ field = build_voxel_fields(
+ voxel_list, voxel_res, device
+ ) # (B, R, R, R)
+ cond = voxel_encoder(field) # (B, R'^3, cond_in_dim)
+ else:
+ cond = None # unconditional
+
+ z0 = torch.randn(verts.shape[0], verts.shape[1], z_dim, device=device)
+ z_flow = sampler.sample(
+ model=tflow,
+ noise=z0,
+ verts=verts,
+ mask=mask,
+ cond=cond,
+ steps=args.steps,
+ )
+ z = z_flow.float() / latent_scale
+
+ with torch.autocast("cuda", dtype=torch.bfloat16):
+ edges_list = tvae.decode(
+ z,
+ verts=verts,
+ verts_mask=mask,
+ chunk_size=args.chunk_size,
+ threshold=args.edge_threshold,
+ )
+
+ for b, name in enumerate(names):
+ num_vertices = lengths[b]
+ if num_vertices == 0:
+ logging.warning(f"{name}: no known vertices; skipping.")
+ continue
+ edges = edges_list[b]
+ faces = edges_to_faces(edges, num_vertices, args.fill_quad_rings)
+
+ verts_int = verts_list[b]
+ disp = (verts_int.numpy().astype(np.float64) + 0.5) / num_discrete - 0.5
+
+ if faces.shape[0] > 0:
+ trimesh.Trimesh(vertices=disp, faces=faces).export(
+ os.path.join(args.out_dir, f"{name}_pred.obj")
+ )
+ else:
+ trimesh.PointCloud(disp).export(
+ os.path.join(args.out_dir, f"{name}_pred.ply")
+ )
+ trimesh.PointCloud(disp).export(
+ os.path.join(args.out_dir, f"{name}_known.ply")
+ )
+ if args.use_cond and args.save_voxel_field and voxel_list[b].shape[0] > 0:
+ vox_pts = (
+ voxel_list[b].numpy().astype(np.float64) + 0.5
+ ) / voxel_res - 0.5
+ trimesh.PointCloud(vox_pts).export(
+ os.path.join(args.out_dir, f"{name}_voxel_field.ply")
+ )
+
+ n_ok += 1
+
+ logging.info(
+ f"done: {n_ok} ok, {n_fail} failed in {time.time() - t_start:.0f}s -> {args.out_dir}"
+ )
+
+
+if __name__ == "__main__":
+ main()
diff --git a/scripts/vflow_inference.py b/scripts/vflow_inference.py
new file mode 100644
index 0000000000000000000000000000000000000000..37d435271556644fcecd72fc3cdf120debd92f82
--- /dev/null
+++ b/scripts/vflow_inference.py
@@ -0,0 +1,298 @@
+"""
+Outputs PLYs and renders into --out_dir:
+ _gt_coords.ply GT vertex voxel coords in [0, 1024)
+ _gt.ply GT vertices (+GT offsets) in [-0.5, 0.5],
+ _recon_coords.ply reconstructed vertex voxel coords in [0, 1024) (if --reconstruct)
+ _recon.ply reconstructed vertices (+offset head) in [-0.5, 0.5] (if --reconstruct)
+ _pred_coords.ply flow-generated vertex voxel coords in [0, 1024)
+ _pred.ply flow-generated vertices (+offset head) in [-0.5, 0.5]
+ _render.png the conditioning view fed to DINO-v2
+
+Usage:
+ python scripts/vflow_inference.py --mesh_dir --out_dir outputs/vflow_run/ \
+ [--vert_num 2000] [--cfg_strength 3.0] [--steps 24] \
+ [--render_azimuth 45 --render_elevation 30] [--reconstruct]
+"""
+
+import argparse
+import os
+import sys
+import time
+from collections import Counter
+from functools import partial
+
+import tqdm
+
+ROOT = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
+sys.path.insert(0, ROOT)
+os.environ.setdefault("PYTORCH_CUDA_ALLOC_CONF", "expandable_segments:True")
+os.environ.setdefault("XDG_RUNTIME_DIR", "/tmp/runtime-root")
+os.makedirs(os.environ["XDG_RUNTIME_DIR"], exist_ok=True)
+# Open3D headless rendering: without this the default EGL platform can hang
+# (e.g. when every GPU is busy or no display device is exposed).
+os.environ.setdefault("EGL_PLATFORM", "surfaceless")
+
+import numpy as np
+import torch
+from PIL import Image
+from torch.utils.data import DataLoader
+
+from dataset.voxel_dataset import VoxelVertexDataset, collate_fn
+from models import (
+ DinoV2Encoder,
+ OffsetHead,
+ VDFEncoder,
+ VertexVAE,
+ VertexSLatFlowModel,
+ VertFlowEulerCfgSampler,
+)
+from modules.sparse import SparseTensor
+import utils.logging as logging
+from utils.export import export_vertex
+from utils.inference import compute_density, decode_vertices, worker_init
+from utils.load import load_latov2_model
+
+
+def parse_args():
+ p = argparse.ArgumentParser(description="V-Flow vertex generation inference")
+ p.add_argument("--mesh_dir", required=True, help="directory of input meshes")
+ p.add_argument(
+ "--out_dir", required=True, help="output directory for the PLYs / renders"
+ )
+ p.add_argument("--vflow_ckpt", default=os.path.join(ROOT, "ckpt", "vflow.pt"))
+ p.add_argument("--vvae_ckpt", default=os.path.join(ROOT, "ckpt", "vvae.pt"))
+ p.add_argument(
+ "--vdf_encoder_ckpt", default=os.path.join(ROOT, "ckpt", "vdf_encoder.pt")
+ )
+ p.add_argument(
+ "--offset_head_ckpt", default=os.path.join(ROOT, "ckpt", "offset_head.pt")
+ )
+ p.add_argument("--batch_size", type=int, default=1)
+ p.add_argument(
+ "--num_samples", type=int, default=None, help="only run the first N meshes"
+ )
+ p.add_argument("--num_workers", type=int, default=4)
+ p.add_argument("--pc_sample_number", type=int, default=819200)
+ p.add_argument("--sample_type", choices=["dora", "uniform"], default="dora")
+ p.add_argument("--inference_threshold", type=float, default=0.5)
+ p.add_argument(
+ "--reconstruct",
+ action=argparse.BooleanOptionalAction,
+ default=False,
+ help="also reconstruct the vertex using V-VAE",
+ )
+ p.add_argument(
+ "--sample_posterior",
+ action=argparse.BooleanOptionalAction,
+ default=True,
+ help="sample the GT posterior (training-style) instead of taking its mode",
+ )
+ p.add_argument("--seed", type=int, default=42)
+ # flow sampling
+ p.add_argument("--steps", type=int, default=24, help="Euler steps")
+ p.add_argument("--cfg_strength", type=float, default=3.0)
+ p.add_argument("--rescale_t", type=float, default=1.0)
+ # vertex-count density conditioning
+ p.add_argument("--vert_num", type=int, default=2000, help="target vertex count")
+ p.add_argument(
+ "--use_gt_vert_count",
+ action=argparse.BooleanOptionalAction,
+ default=False,
+ help="condition on the GT quantized vertex count instead of --vert_num",
+ )
+ p.add_argument(
+ "--scaler",
+ type=float,
+ default=1.0,
+ help="multiplier on the GT count when --use_gt_vert_count",
+ )
+ p.add_argument("--min_verts", type=float, default=200.0)
+ p.add_argument("--max_verts", type=float, default=5000.0)
+ # conditioning render
+ p.add_argument("--render_azimuth", type=float, default=45.0)
+ p.add_argument("--render_elevation", type=float, default=30.0)
+ p.add_argument("--img_res", type=int, default=518)
+ p.add_argument(
+ "--dino_hub_dir",
+ default=os.path.join(ROOT, "ckpt", "dinov2"),
+ help="torch.hub cache for DINO-v2; reused when present, downloaded otherwise",
+ )
+ args = p.parse_args()
+ if args.num_samples is not None and args.num_samples <= 0:
+ args.num_samples = None # <= 0 means "all", not python slice semantics
+ return args
+
+
+def main():
+ logging.info("V-Flow inference starting...")
+
+ args = parse_args()
+ device = torch.device("cuda")
+ torch.manual_seed(args.seed)
+ np.random.seed(args.seed)
+ os.makedirs(args.out_dir, exist_ok=True)
+
+ vflow, vflow_cfg = load_latov2_model(VertexSLatFlowModel, args.vflow_ckpt, device)
+ vvae, vvae_cfg = load_latov2_model(VertexVAE, args.vvae_ckpt, device)
+ vdf_encoder, _ = load_latov2_model(VDFEncoder, args.vdf_encoder_ckpt, device)
+ offset_head, _ = load_latov2_model(OffsetHead, args.offset_head_ckpt, device)
+ res = vvae_cfg["resolution"]
+ min_res = vvae_cfg["min_resolution"]
+ latent_dim = vflow_cfg["latent_dim"]
+ density_max = vflow_cfg["max_vertex_num"]
+
+ dino = (
+ DinoV2Encoder(
+ model_name=vflow_cfg["dino_version"],
+ hub_dir=args.dino_hub_dir,
+ img_res=vflow_cfg["image_resolution"],
+ )
+ .to(device)
+ .eval()
+ )
+ logging.info(
+ f"loaded {vflow_cfg['dino_version']} from {args.dino_hub_dir}"
+ )
+ sampler = VertFlowEulerCfgSampler()
+
+ dataset = VoxelVertexDataset(
+ root_dir=args.mesh_dir,
+ resolution=res,
+ min_resolution=min_res,
+ pc_sample_number=args.pc_sample_number,
+ sample_type=args.sample_type,
+ num_samples=args.num_samples,
+ render=True,
+ img_res=args.img_res,
+ render_azimuth=args.render_azimuth,
+ render_elevation=args.render_elevation,
+ )
+ loader = DataLoader(
+ dataset,
+ batch_size=args.batch_size,
+ shuffle=False,
+ collate_fn=partial(collate_fn, resolution=res, min_resolution=min_res),
+ num_workers=args.num_workers,
+ pin_memory=True,
+ # EGL rendering hangs inside fork-ed children of a CUDA-initialized
+ # parent; spawn gives each worker a clean process for its EGL context.
+ multiprocessing_context="spawn" if args.num_workers > 0 else None,
+ worker_init_fn=worker_init if args.num_workers > 0 else None,
+ )
+ dupes = sorted(
+ s
+ for s, c in Counter(os.path.splitext(f)[0] for f in dataset.files).items()
+ if c > 1
+ )
+ if dupes:
+ logging.warning(
+ f"WARNING: {len(dupes)} duplicate mesh basename(s) — later samples will overwrite earlier outputs."
+ )
+ logging.info(
+ f"{len(dataset)} meshes from {args.mesh_dir} "
+ f"(steps={args.steps}, cfg={args.cfg_strength}, vert_num={args.vert_num}, "
+ f"use_gt_vert_count={args.use_gt_vert_count}, scaler={args.scaler}, density_max={density_max}, "
+ f"view=az{args.render_azimuth}/el{args.render_elevation}, seed={args.seed}) "
+ f"-> {args.out_dir}"
+ )
+
+ n_ok = n_fail = 0
+ t_start = time.time()
+ qbar = tqdm.tqdm(loader, desc="inference", unit="batch", dynamic_ncols=True)
+ for batch in qbar:
+ for err in batch["errors"]:
+ n_fail += 1
+ logging.error(
+ f"{err['name']}: FAILED during preprocessing: {err['error'].splitlines()[0]}"
+ )
+ if "name" not in batch:
+ continue
+
+ density = compute_density(batch, args, density_max, device)
+ with torch.no_grad():
+ cond = dino(np.stack(batch["image"])).float()
+ neg_cond = torch.zeros_like(cond)
+
+ with torch.no_grad(), torch.autocast("cuda", dtype=torch.bfloat16):
+ if args.reconstruct:
+ point_cloud = batch[f"point_cloud_{res}"].to(device)
+ vertex_added_active_coords = batch[f"vertex_added_active_voxels_{res}"].to(device)
+ feats = vdf_encoder(
+ p=point_cloud,
+ sparse_coords=vertex_added_active_coords,
+ res=res,
+ bbox_size=(-0.5, 0.5),
+ )
+ z_gt, _ = vvae.encode(
+ SparseTensor(feats=feats, coords=vertex_added_active_coords.int()),
+ sample_posterior=args.sample_posterior,
+ )
+
+ min_active_coords = batch[f"active_voxels_{min_res}"].to(device)
+ noise = SparseTensor(
+ coords=min_active_coords.int(),
+ feats=torch.randn(min_active_coords.shape[0], latent_dim, device=device),
+ )
+ z_pred = sampler.sample(
+ model=vflow,
+ noise=noise,
+ cond=cond,
+ neg_cond=neg_cond,
+ steps=args.steps,
+ cfg_strength=args.cfg_strength,
+ rescale_t=args.rescale_t,
+ density=density,
+ )
+
+ pred_coords, pred_offsets = decode_vertices(
+ vvae, offset_head, z_pred, args.inference_threshold
+ )
+ if args.reconstruct:
+ recon_coords, recon_offsets = decode_vertices(
+ vvae, offset_head, z_gt, args.inference_threshold
+ )
+
+ gt_vox = batch[f"gt_vertex_voxels_{res}"]
+ gt_off = batch[f"gt_vertex_offsets_{res}"]
+ for b, name in enumerate(batch["name"]):
+ gt_sel = gt_vox[:, 0] == b
+ export_vertex(
+ args.out_dir,
+ name,
+ type_name="gt",
+ vert_int=gt_vox[gt_sel, 1:].numpy(),
+ vert_offsets=gt_off[gt_sel].numpy(),
+ resolution=res,
+ )
+ if args.reconstruct:
+ rec_sel = recon_coords[:, 0] == b
+ export_vertex(
+ args.out_dir,
+ name,
+ type_name="recon",
+ vert_int=recon_coords[rec_sel, 1:].numpy(),
+ vert_offsets=recon_offsets[rec_sel].numpy(),
+ resolution=res,
+ )
+ pred_sel = pred_coords[:, 0] == b
+ export_vertex(
+ args.out_dir,
+ name,
+ type_name="pred",
+ vert_int=pred_coords[pred_sel, 1:].numpy(),
+ vert_offsets=pred_offsets[pred_sel].numpy(),
+ resolution=res,
+ )
+ Image.fromarray(batch["image"][b]).save(
+ os.path.join(args.out_dir, f"{name}_render.png")
+ )
+
+ n_ok += 1
+
+ logging.info(
+ f"done: {n_ok} ok, {n_fail} failed in {time.time() - t_start:.0f}s -> {args.out_dir}"
+ )
+
+
+if __name__ == "__main__":
+ main()
diff --git a/scripts/vvae_inference.py b/scripts/vvae_inference.py
new file mode 100644
index 0000000000000000000000000000000000000000..48cc0e2681427384a0e0ee7464b487f0640330d7
--- /dev/null
+++ b/scripts/vvae_inference.py
@@ -0,0 +1,181 @@
+"""
+Outputs PLYs into --out_dir:
+ _gt_coords.ply GT vertex voxel coords in [0, 1024)
+ _gt.ply GT vertices (+GT offsets) in [-0.5, 0.5]
+ _recon_coords.ply reconstructed vertex voxel coords in [0, 1024)
+ _recon.ply reconstructed vertices (+offset head) in [-0.5, 0.5]
+
+Usage:
+ python scripts/vvae_inference.py --mesh_dir --out_dir outputs/vvae_run/ \
+ [--batch_size 2] [--num_samples 8]
+"""
+
+import argparse
+import os
+import sys
+import time
+from collections import Counter
+from functools import partial
+import tqdm
+
+ROOT = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
+sys.path.insert(0, ROOT)
+os.environ.setdefault("PYTORCH_CUDA_ALLOC_CONF", "expandable_segments:True")
+
+import numpy as np
+import torch
+from torch.utils.data import DataLoader
+
+from dataset.voxel_dataset import VoxelVertexDataset, collate_fn
+from models import OffsetHead, VDFEncoder, VertexVAE
+from modules.sparse import SparseTensor
+import utils.logging as logging
+from utils.export import export_vertex
+from utils.load import load_latov2_model
+
+
+def parse_args():
+ p = argparse.ArgumentParser(description="V-VAE vertex reconstruction inference")
+ p.add_argument("--mesh_dir", required=True, help="directory of input meshes")
+ p.add_argument("--out_dir", required=True, help="output directory for the PLYs")
+ p.add_argument("--vvae_ckpt", default=os.path.join(ROOT, "ckpt", "vvae.pt"))
+ p.add_argument(
+ "--vdf_encoder_ckpt", default=os.path.join(ROOT, "ckpt", "vdf_encoder.pt")
+ )
+ p.add_argument(
+ "--offset_head_ckpt", default=os.path.join(ROOT, "ckpt", "offset_head.pt")
+ )
+ p.add_argument("--batch_size", type=int, default=1)
+ p.add_argument(
+ "--num_samples", type=int, default=None, help="only run the first N meshes"
+ )
+ p.add_argument("--num_workers", type=int, default=4)
+ p.add_argument("--pc_sample_number", type=int, default=819200)
+ p.add_argument("--sample_type", choices=["dora", "uniform"], default="dora")
+ p.add_argument("--inference_threshold", type=float, default=0.5)
+ p.add_argument(
+ "--sample_posterior",
+ action=argparse.BooleanOptionalAction,
+ default=True,
+ help="sample the posterior (training-style) instead of taking its mode",
+ )
+ p.add_argument("--seed", type=int, default=42)
+ args = p.parse_args()
+ if args.num_samples is not None and args.num_samples <= 0:
+ args.num_samples = None # <= 0 means "all", not python slice semantics
+ return args
+
+
+def main():
+ logging.info("V-VAE inference starting...")
+
+ args = parse_args()
+ device = torch.device("cuda")
+ torch.manual_seed(args.seed)
+ np.random.seed(args.seed)
+ os.makedirs(args.out_dir, exist_ok=True)
+
+ vvae, vvae_cfg = load_latov2_model(VertexVAE, args.vvae_ckpt, device)
+ vdf_encoder, _ = load_latov2_model(VDFEncoder, args.vdf_encoder_ckpt, device)
+ offset_head, _ = load_latov2_model(OffsetHead, args.offset_head_ckpt, device)
+ res = vvae_cfg["resolution"]
+
+ dataset = VoxelVertexDataset(
+ root_dir=args.mesh_dir,
+ resolution=res,
+ pc_sample_number=args.pc_sample_number,
+ sample_type=args.sample_type,
+ num_samples=args.num_samples,
+ )
+ dupes = sorted(
+ s
+ for s, c in Counter(os.path.splitext(f)[0] for f in dataset.files).items()
+ if c > 1
+ )
+ if dupes:
+ logging.warning(
+ f"WARNING: {len(dupes)} duplicate mesh basename(s) — later samples will overwrite earlier outputs."
+ )
+ loader = DataLoader(
+ dataset,
+ batch_size=args.batch_size,
+ shuffle=False,
+ collate_fn=partial(collate_fn, resolution=res),
+ num_workers=args.num_workers,
+ pin_memory=True,
+ )
+ logging.info(
+ f"{len(dataset)} meshes from {args.mesh_dir} "
+ f"(batch_size={args.batch_size}, seed={args.seed}, "
+ f"sample_posterior={args.sample_posterior}) -> {args.out_dir}"
+ )
+
+ n_ok = n_fail = 0
+ t_start = time.time()
+ qbar = tqdm.tqdm(loader, desc="inference", unit="batch", dynamic_ncols=True)
+ for batch in qbar:
+ for err in batch["errors"]:
+ n_fail += 1
+ logging.error(
+ f"{err['name']}: FAILED during preprocessing: {err['error'].splitlines()[0]}"
+ )
+ if "name" not in batch:
+ continue
+
+ with torch.no_grad(), torch.autocast("cuda", dtype=torch.bfloat16):
+ point_cloud = batch[f"point_cloud_{res}"].to(device)
+ vertex_added_active_coords = batch[f"vertex_added_active_voxels_{res}"].to(device)
+
+ feats = vdf_encoder(
+ p=point_cloud,
+ sparse_coords=vertex_added_active_coords,
+ res=res,
+ bbox_size=(-0.5, 0.5),
+ )
+ z, _ = vvae.encode(
+ SparseTensor(feats=feats, coords=vertex_added_active_coords.int()),
+ sample_posterior=args.sample_posterior,
+ )
+ decoded = vvae.decode(
+ z,
+ gt_vertex_voxels_list=[],
+ training=False,
+ inference_threshold=args.inference_threshold,
+ )
+ recon_coords = decoded[-1]["coords"] # (N, 4) int, batch col first
+ recon_offsets = offset_head(
+ decoded[-1]["feats"]
+ ).float() # (N, 3) in (-1, 1)
+
+ gt_vox = batch[f"gt_vertex_voxels_{res}"]
+ gt_off = batch[f"gt_vertex_offsets_{res}"]
+ recon_coords = recon_coords.cpu()
+ recon_offsets = recon_offsets.cpu()
+ for b, name in enumerate(batch["name"]):
+ gt_sel = gt_vox[:, 0] == b
+ rec_sel = recon_coords[:, 0] == b
+ export_vertex(
+ args.out_dir,
+ name,
+ type_name="gt",
+ vert_int=gt_vox[gt_sel, 1:].numpy(),
+ vert_offsets=gt_off[gt_sel].numpy(),
+ resolution=res,
+ )
+ export_vertex(
+ args.out_dir,
+ name,
+ type_name="recon",
+ vert_int=recon_coords[rec_sel, 1:].numpy(),
+ vert_offsets=recon_offsets[rec_sel].numpy(),
+ resolution=res,
+ )
+ n_ok += 1
+
+ logging.info(
+ f"done: {n_ok} ok, {n_fail} failed in {time.time() - t_start:.0f}s -> {args.out_dir}"
+ )
+
+
+if __name__ == "__main__":
+ main()
diff --git a/setup.sh b/setup.sh
new file mode 100644
index 0000000000000000000000000000000000000000..2c43c68255bd6700fbfa27012ed1b11d4810ee2e
--- /dev/null
+++ b/setup.sh
@@ -0,0 +1,202 @@
+# copied and modified from https://github.com/microsoft/TRELLIS.2/blob/main/setup.sh
+
+# Read Arguments
+TEMP=`getopt -o h --long help,all,new-env,basic,flash-attn,nvdiffrast,cumesh,flexgemm,o-voxel,xformers -n 'setup.sh' -- "$@"`
+
+eval set -- "$TEMP"
+
+HELP=false
+ALL=false
+NEW_ENV=false
+BASIC=false
+FLASHATTN=false
+NVDIFFRAST=false
+CUMESH=false
+FLEXGEMM=false
+OVOXEL=false
+XFORMERS=false
+ERROR=false
+
+
+if [ "$#" -eq 1 ] ; then
+ HELP=true
+fi
+
+while true ; do
+ case "$1" in
+ -h|--help) HELP=true ; shift ;;
+ --all) ALL=true ; shift ;;
+ --new-env) NEW_ENV=true ; shift ;;
+ --basic) BASIC=true ; shift ;;
+ --flash-attn) FLASHATTN=true ; shift ;;
+ --nvdiffrast) NVDIFFRAST=true ; shift ;;
+ --cumesh) CUMESH=true ; shift ;;
+ --flexgemm) FLEXGEMM=true ; shift ;;
+ --o-voxel) OVOXEL=true ; shift ;;
+ --xformers) XFORMERS=true ; shift ;;
+ --) shift ; break ;;
+ *) ERROR=true ; break ;;
+ esac
+done
+
+if [ "$ERROR" = true ] ; then
+ echo "Error: Invalid argument"
+ HELP=true
+fi
+
+if [ "$HELP" = true ] ; then
+ echo "Usage: setup.sh [OPTIONS]"
+ echo "Options:"
+ echo " -h, --help Display this help message"
+ echo " --all Run every step below (full inference environment)"
+ echo " --new-env Create the 'lato2' conda env (python 3.10 + torch 2.6.0)"
+ echo " --basic Install core runtime deps (spconv, torch-scatter, open3d, ...)"
+ echo " --flash-attn Install flash-attention (default attention backend)"
+ echo " --nvdiffrast Build nvdiffrast (required at o_voxel import time)"
+ echo " --cumesh Build CuMesh (required at o_voxel import time)"
+ echo " --flexgemm Build FlexGEMM (required at o_voxel import time)"
+ echo " --o-voxel Build/install o_voxel (sparse-cloned from the TRELLIS.2 repo)"
+ echo " --xformers (optional) Install xformers (alt attn backend / DINOv2 speedup)"
+ return
+fi
+
+if [ "$ALL" = true ] ; then
+ NEW_ENV=true
+ BASIC=true
+ FLASHATTN=true
+ NVDIFFRAST=true
+ CUMESH=true
+ FLEXGEMM=true
+ OVOXEL=true
+fi
+
+WORKDIR=$(pwd)
+if command -v nvidia-smi > /dev/null; then
+ PLATFORM="cuda"
+elif command -v rocminfo > /dev/null; then
+ PLATFORM="hip"
+else
+ echo "Error: No supported GPU found"
+ return 1
+fi
+
+CONDA_EXE_PATH=""
+for c in "$CONDA_EXE" "$HOME/miniconda3/bin/conda" "$HOME/anaconda3/bin/conda" /opt/conda/bin/conda ; do
+ if [ -n "$c" ] && [ -x "$c" ] ; then CONDA_EXE_PATH="$c" ; break ; fi
+done
+if [ -z "$CONDA_EXE_PATH" ] ; then
+ _c=$(command -v conda 2>/dev/null)
+ if [ -x "$_c" ] ; then CONDA_EXE_PATH="$_c" ; fi
+fi
+if [ -n "$CONDA_EXE_PATH" ] ; then
+ CONDA_BASE_DIR=$("$CONDA_EXE_PATH" info --base 2>/dev/null)
+ if [ -n "$CONDA_BASE_DIR" ] && [ -f "$CONDA_BASE_DIR/etc/profile.d/conda.sh" ] ; then
+ source "$CONDA_BASE_DIR/etc/profile.d/conda.sh"
+ fi
+elif [ "$NEW_ENV" = true ] ; then
+ echo "Error: conda not found; cannot create a new env"
+ return 1
+fi
+
+ensure_cuda() {
+ if command -v nvcc > /dev/null 2>&1 ; then return 0 ; fi
+ for cu in "$CUDA_HOME" /usr/local/cuda /usr/local/cuda-12.4 ; do
+ if [ -n "$cu" ] && [ -x "$cu/bin/nvcc" ] ; then
+ export CUDA_HOME="$cu"
+ export PATH="$cu/bin:$PATH"
+ return 0
+ fi
+ done
+ return 1
+}
+
+ENV_NAME="${LATO_ENV:-lato2}"
+
+if [ "$NEW_ENV" = true ] ; then
+ conda create -n "$ENV_NAME" python=3.10 -y || { echo "Error: 'conda create' failed (a corrupted conda package cache is a common cause; try 'conda clean --all')"; return 1; }
+ conda activate "$ENV_NAME" || { echo "Error: 'conda activate $ENV_NAME' failed"; return 1; }
+ if [ "$PLATFORM" = "cuda" ] ; then
+ pip install torch==2.6.0 torchvision==0.21.0 --index-url https://download.pytorch.org/whl/cu124
+ elif [ "$PLATFORM" = "hip" ] ; then
+ pip install torch==2.6.0 torchvision==0.21.0 --index-url https://download.pytorch.org/whl/rocm6.2.4
+ fi
+fi
+
+if [ "$BASIC" = true ] || [ "$FLASHATTN" = true ] || [ "$NVDIFFRAST" = true ] || [ "$CUMESH" = true ] || [ "$FLEXGEMM" = true ] || [ "$OVOXEL" = true ] || [ "$XFORMERS" = true ] ; then
+ if [ "$CONDA_DEFAULT_ENV" != "$ENV_NAME" ] ; then
+ conda activate "$ENV_NAME" || { echo "Error: could not activate env '$ENV_NAME'. Create it first with --new-env."; return 1; }
+ fi
+fi
+
+if [ "$BASIC" = true ] ; then
+ pip install numpy trimesh tqdm pillow ninja psutil opencv-python-headless huggingface_hub open3d==0.19.0 plyfile zstandard easydict
+ if [ "$PLATFORM" = "cuda" ] ; then
+ pip install spconv-cu124==2.3.8
+ pip install torch-scatter -f https://data.pyg.org/whl/torch-2.6.0+cu124.html
+ elif [ "$PLATFORM" = "hip" ] ; then
+ echo "[BASIC] spconv/torch-scatter have no prebuilt ROCm wheels here; install manually."
+ fi
+fi
+
+if [ "$FLASHATTN" = true ] ; then
+ if [ "$PLATFORM" = "cuda" ] ; then
+ ensure_cuda || echo "[FLASHATTN] nvcc not found; ok if a prebuilt wheel is used, required for a source build."
+ pip install flash-attn==2.7.4.post1 --no-build-isolation --no-cache-dir
+ elif [ "$PLATFORM" = "hip" ] ; then
+ echo "[FLASHATTN] Prebuilt binaries not found. Building from source..."
+ mkdir -p /tmp/extensions
+ git clone --recursive https://github.com/ROCm/flash-attention.git /tmp/extensions/flash-attention
+ cd /tmp/extensions/flash-attention
+ git checkout tags/v2.7.3-cktile
+ GPU_ARCHS=gfx942 python setup.py install #MI300 series
+ cd $WORKDIR
+ else
+ echo "[FLASHATTN] Unsupported platform: $PLATFORM"
+ fi
+fi
+
+if [ "$NVDIFFRAST" = true ] ; then
+ if [ "$PLATFORM" = "cuda" ] ; then
+ ensure_cuda || echo "[NVDIFFRAST] nvcc not found; required to build."
+ mkdir -p /tmp/extensions
+ rm -rf /tmp/extensions/nvdiffrast
+ git clone -b v0.4.0 https://github.com/NVlabs/nvdiffrast.git /tmp/extensions/nvdiffrast
+ pip install /tmp/extensions/nvdiffrast --no-build-isolation --no-cache-dir
+ else
+ echo "[NVDIFFRAST] Unsupported platform: $PLATFORM"
+ fi
+fi
+
+if [ "$CUMESH" = true ] ; then
+ ensure_cuda || echo "[CUMESH] nvcc not found; required to build."
+ mkdir -p /tmp/extensions
+ rm -rf /tmp/extensions/CuMesh
+ git clone https://github.com/JeffreyXiang/CuMesh.git /tmp/extensions/CuMesh --recursive
+ pip install /tmp/extensions/CuMesh --no-build-isolation --no-cache-dir
+fi
+
+if [ "$FLEXGEMM" = true ] ; then
+ ensure_cuda || echo "[FLEXGEMM] nvcc not found; required to build."
+ mkdir -p /tmp/extensions
+ rm -rf /tmp/extensions/FlexGEMM
+ git clone https://github.com/JeffreyXiang/FlexGEMM.git /tmp/extensions/FlexGEMM --recursive
+ pip install /tmp/extensions/FlexGEMM --no-build-isolation --no-cache-dir
+fi
+
+if [ "$OVOXEL" = true ] ; then
+ mkdir -p /tmp/extensions
+ ensure_cuda || echo "[O_VOXEL] nvcc not found; set CUDA_HOME to a CUDA toolkit and retry."
+ rm -rf /tmp/extensions/TRELLIS.2
+ git clone --depth 1 --filter=blob:none --sparse https://github.com/microsoft/TRELLIS.2.git /tmp/extensions/TRELLIS.2
+ git -C /tmp/extensions/TRELLIS.2 sparse-checkout set o-voxel
+ git -C /tmp/extensions/TRELLIS.2 submodule update --init --recursive --depth 1
+ pip install /tmp/extensions/TRELLIS.2/o-voxel --no-build-isolation --no-cache-dir
+fi
+
+if [ "$XFORMERS" = true ] ; then
+ if [ "$PLATFORM" = "cuda" ] ; then
+ pip install xformers==0.0.29.post2 --index-url https://download.pytorch.org/whl/cu124
+ else
+ echo "[XFORMERS] Unsupported platform: $PLATFORM"
+ fi
+fi
diff --git a/utils/export.py b/utils/export.py
new file mode 100644
index 0000000000000000000000000000000000000000..e0aaf0aa20ddbe9afd46aebb7ae3d3ede087f3c2
--- /dev/null
+++ b/utils/export.py
@@ -0,0 +1,19 @@
+import os
+import numpy as np
+import trimesh
+
+
+def export_vertex(out_dir, base_name, type_name, vert_int, vert_offsets, resolution):
+ res = float(resolution)
+ vert_with_offset = (
+ vert_int.astype(np.float64) / res
+ - 0.5
+ + vert_offsets.astype(np.float64) / (res * 2.0)
+ )
+
+ trimesh.PointCloud(vert_with_offset).export(
+ os.path.join(out_dir, f"{base_name}_{type_name}.ply")
+ )
+ trimesh.PointCloud(vert_int).export(
+ os.path.join(out_dir, f"{base_name}_{type_name}_coords.ply")
+ )
diff --git a/utils/inference.py b/utils/inference.py
new file mode 100644
index 0000000000000000000000000000000000000000..1b975b84a38f7a4161d7ef9af0095ac44a26d571
--- /dev/null
+++ b/utils/inference.py
@@ -0,0 +1,122 @@
+import os
+import numpy as np
+import torch
+
+
+def worker_init(_worker_id):
+ import atexit
+
+ atexit.register(os._exit, 0)
+
+
+def compute_density(batch, args, density_max, device):
+ counts = []
+ for verts in batch["quantized_vertices"]:
+ if args.use_gt_vert_count:
+ counts.append(
+ min(
+ max(float(verts.shape[0]) * args.scaler, args.min_verts),
+ args.max_verts,
+ )
+ )
+ else:
+ counts.append(float(args.vert_num))
+ counts = torch.tensor(counts, dtype=torch.float32, device=device)
+ return counts / density_max * 1000.0
+
+
+def decode_vertices(vvae, offset_head, latent, inference_threshold):
+ decoded = vvae.decode(
+ latent,
+ gt_vertex_voxels_list=[],
+ training=False,
+ inference_threshold=inference_threshold,
+ )
+ coords = decoded[-1]["coords"]
+ offsets = offset_head(decoded[-1]["feats"]).float()
+ return coords.cpu(), offsets.cpu()
+
+
+def build_voxel_fields(voxel_coords_list, voxel_res, device):
+ bsz = len(voxel_coords_list)
+ field = torch.zeros(
+ (bsz, voxel_res, voxel_res, voxel_res), dtype=torch.float32, device=device
+ )
+ for b, coords in enumerate(voxel_coords_list):
+ if coords.numel() > 0:
+ c = coords.to(device=device, dtype=torch.long)
+ field[b, c[:, 0], c[:, 1], c[:, 2]] = 1.0
+ return field
+
+
+def pad_verts(verts_list, device):
+ bsz = len(verts_list)
+ lengths = [int(v.shape[0]) for v in verts_list]
+ n_max = max(lengths)
+ verts = torch.zeros(bsz, n_max, 3, dtype=torch.long, device=device)
+ mask = torch.zeros(bsz, n_max, dtype=torch.bool, device=device)
+ for b, (v, n) in enumerate(zip(verts_list, lengths)):
+ verts[b, :n] = v.to(device=device, dtype=torch.long)
+ mask[b, :n] = True
+ return verts, mask, lengths
+
+
+def triangulate_quad_rings(adj: np.ndarray) -> np.ndarray:
+ # thinks meshflow@CVPR2026!
+ num = int(adj.shape[0])
+ if num < 4:
+ return np.empty((0, 3), dtype=np.int32)
+ nbr_sets = [set(np.nonzero(adj[v])[0].tolist()) for v in range(num)]
+ new_faces: list[list[int]] = []
+ seen: set[tuple[int, int, int]] = set()
+ for a in range(num):
+ cand = sorted(v for v in nbr_sets[a] if v > a)
+ n_cand = len(cand)
+ if n_cand < 2:
+ continue
+ for i in range(n_cand):
+ b = cand[i]
+ set_b = nbr_sets[b]
+ for j in range(i + 1, n_cand):
+ d = cand[j]
+ if d in set_b:
+ continue
+ for c in set_b & nbr_sets[d]:
+ if c <= a or c in nbr_sets[a]:
+ continue
+ for tri in ([a, b, c], [a, c, d]):
+ key = tuple(sorted(tri))
+ if key in seen:
+ continue
+ seen.add(key)
+ new_faces.append(tri)
+ if not new_faces:
+ return np.empty((0, 3), dtype=np.int32)
+ return np.asarray(new_faces, dtype=np.int32)
+
+
+def edges_to_faces(
+ edges: np.ndarray, num_valid: int, fill_quad_rings: bool
+) -> np.ndarray:
+ adj = np.zeros((num_valid, num_valid), dtype=bool)
+ if edges.shape[0] > 0:
+ adj[edges[:, 0], edges[:, 1]] = True
+ adj[edges[:, 1], edges[:, 0]] = True
+ faces_list = []
+ for ei in range(edges.shape[0]):
+ u = int(edges[ei, 0])
+ v = int(edges[ei, 1])
+ common = np.where(adj[u] & adj[v])[0]
+ common = common[common > v]
+ for w in common:
+ faces_list.append([u, v, int(w)])
+ faces = (
+ np.asarray(faces_list, dtype=np.int32)
+ if faces_list
+ else np.empty((0, 3), dtype=np.int32)
+ )
+ if fill_quad_rings:
+ quad = triangulate_quad_rings(adj)
+ if quad.shape[0] > 0:
+ faces = np.concatenate([faces, quad], axis=0) if faces.shape[0] else quad
+ return faces
diff --git a/utils/load.py b/utils/load.py
new file mode 100644
index 0000000000000000000000000000000000000000..ed9edee0a1d31f03f80ca74b950af37f8331fc3a
--- /dev/null
+++ b/utils/load.py
@@ -0,0 +1,11 @@
+import torch
+
+from utils import logging
+
+
+def load_latov2_model(cls, ckpt_path, device):
+ ckpt = torch.load(ckpt_path, map_location="cpu", weights_only=False)
+ model = cls(**ckpt["config"]["args"])
+ model.load_state_dict(ckpt["state_dict"], strict=True)
+ logging.info(f"loaded {ckpt.get('id', cls.__name__)} from {ckpt_path}")
+ return model.to(device).eval(), ckpt["config"]
diff --git a/utils/logging.py b/utils/logging.py
new file mode 100644
index 0000000000000000000000000000000000000000..48fdf527a4f1e0cfd50adf7966fa3193cdb703b3
--- /dev/null
+++ b/utils/logging.py
@@ -0,0 +1,70 @@
+# copied and modified from https://github.com/LoHhhha/pmos_nn/blob/master/flowing/shower/Logger.py
+
+import os
+import inspect
+import time
+
+DEBUG_MSG = 0
+INFO_MSG = 1
+WARNING_MSG = 2
+ERROR_MSG = 3
+FAULT_MSG = 4
+PACKAGE_PATH = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
+
+LOGGER_LEVER = INFO_MSG
+
+MSG_TYPES_STR = [
+ "\033[1m\033[36mD\033[0m",
+ "\033[1m\033[32mI\033[0m",
+ "\033[1m\033[33mW\033[0m",
+ "\033[1m\033[35mE\033[0m",
+ "\033[1m\033[31mF\033[0m",
+]
+
+
+def debug(*msg):
+ __print_out(DEBUG_MSG, *msg)
+
+
+def info(*msg):
+ __print_out(INFO_MSG, *msg)
+
+
+def warning(*msg):
+ __print_out(WARNING_MSG, *msg)
+
+
+def error(*msg):
+ __print_out(ERROR_MSG, *msg)
+
+
+def fault(*msg):
+ __print_out(FAULT_MSG, *msg)
+
+
+def __print_out(msg_type: int, *msg) -> None:
+ if msg_type < LOGGER_LEVER:
+ return
+
+ # inspect.stack()[1] is info/warning/error
+ caller_frame = inspect.stack()[2]
+
+ caller_name = (
+ os.path.relpath(caller_frame.filename, PACKAGE_PATH)
+ .split(".")[0]
+ .replace("\\", "/")
+ .replace("/", ".")
+ )
+
+ print(
+ f"[{MSG_TYPES_STR[msg_type]}|{time.strftime('%H:%M:%S')}|{caller_name}:{caller_frame.lineno}]",
+ *msg,
+ )
+
+
+if __name__ == "__main__":
+ debug("debug msg")
+ info("info msg")
+ warning("warning msg")
+ error("error msg")
+ fault("fault msg")