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")