This commit is contained in:
smthemex
2025-10-15 15:35:42 +08:00
parent 86c5b51da5
commit 331f9e6a73
417 changed files with 62838 additions and 0 deletions
+44
View File
@@ -0,0 +1,44 @@
*.7z filter=lfs diff=lfs merge=lfs -text
*.arrow filter=lfs diff=lfs merge=lfs -text
*.bin filter=lfs diff=lfs merge=lfs -text
*.bz2 filter=lfs diff=lfs merge=lfs -text
*.ckpt filter=lfs diff=lfs merge=lfs -text
*.ftz filter=lfs diff=lfs merge=lfs -text
*.gz filter=lfs diff=lfs merge=lfs -text
*.h5 filter=lfs diff=lfs merge=lfs -text
*.joblib filter=lfs diff=lfs merge=lfs -text
*.lfs.* filter=lfs diff=lfs merge=lfs -text
*.mlmodel filter=lfs diff=lfs merge=lfs -text
*.model filter=lfs diff=lfs merge=lfs -text
*.msgpack filter=lfs diff=lfs merge=lfs -text
*.npy filter=lfs diff=lfs merge=lfs -text
*.npz filter=lfs diff=lfs merge=lfs -text
*.onnx filter=lfs diff=lfs merge=lfs -text
*.ot filter=lfs diff=lfs merge=lfs -text
*.parquet filter=lfs diff=lfs merge=lfs -text
*.pb filter=lfs diff=lfs merge=lfs -text
*.pickle filter=lfs diff=lfs merge=lfs -text
*.pkl filter=lfs diff=lfs merge=lfs -text
*.pt filter=lfs diff=lfs merge=lfs -text
*.pth filter=lfs diff=lfs merge=lfs -text
*.rar filter=lfs diff=lfs merge=lfs -text
*.safetensors filter=lfs diff=lfs merge=lfs -text
saved_model/**/* filter=lfs diff=lfs merge=lfs -text
*.tar.* filter=lfs diff=lfs merge=lfs -text
*.tar filter=lfs diff=lfs merge=lfs -text
*.tflite filter=lfs diff=lfs merge=lfs -text
*.tgz filter=lfs diff=lfs merge=lfs -text
*.wasm filter=lfs diff=lfs merge=lfs -text
*.xz filter=lfs diff=lfs merge=lfs -text
*.zip filter=lfs diff=lfs merge=lfs -text
*.zst filter=lfs diff=lfs merge=lfs -text
*tfevents* filter=lfs diff=lfs merge=lfs -text
assets/example_data/Batman.png filter=lfs diff=lfs merge=lfs -text
assets/example_data/astronaut.png filter=lfs diff=lfs merge=lfs -text
assets/example_data/car.png filter=lfs diff=lfs merge=lfs -text
assets/example_data/knight.png filter=lfs diff=lfs merge=lfs -text
assets/example_data/robot1.jpeg filter=lfs diff=lfs merge=lfs -text
assets/example_data/snake.png filter=lfs diff=lfs merge=lfs -text
assets/example_data/warhammer.png filter=lfs diff=lfs merge=lfs -text
modules/part_synthesis/representations/mesh/flexicubes/images/block_init.png filter=lfs diff=lfs merge=lfs -text
modules/part_synthesis/representations/mesh/flexicubes/images/teaser_top.png filter=lfs diff=lfs merge=lfs -text
+6
View File
@@ -0,0 +1,6 @@
__pycache__/
output/
ckpt/
.DS_Store
tmp/
debug_images/
View File
+22
View File
@@ -0,0 +1,22 @@
# MIT License
# 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
+15
View File
@@ -0,0 +1,15 @@
OmniPart
Copyright (c) 2025 VAST-AI-Research and contributors
This project includes code from the following open source projects:
RMBG
Copyright (c) BRIA AI
License: bria-rmbg-2.0
Source: https://huggingface.co/briaai/RMBG-2.0
This software contains code derived from 🤗 Diffusers (https://github.com/huggingface/diffusers), available under the Apache License 2.0.
This software contains code derived from TRELLIS (https://github.com/Microsoft/TRELLIS), available under the MIT License.
This software contains code derived from PartField (https://github.com/nv-tlabs/PartField), available under the NVIDIA Source Code License.
+99
View File
@@ -0,0 +1,99 @@
# OmniPart: Part-Aware 3D Generation with Semantic Decoupling and Structural Cohesion [SIGGRAPH Asia 2025]
<div align="center">
[![Project Page](https://img.shields.io/badge/🏠-Project%20Page-blue.svg)](https://omnipart.github.io/)
[![Paper](https://img.shields.io/badge/📑-Paper-green.svg)](https://arxiv.org/abs/2507.06165)
[![Model](https://img.shields.io/badge/🤗-Model-yellow.svg)](https://huggingface.co/omnipart)
[![Online Demo](https://img.shields.io/badge/%F0%9F%A4%97%20Hugging%20Face-Spaces-blue)](https://huggingface.co/spaces/omnipart/OmniPart)
</div>
![teaser](assets/doc/teaser.jpg)
## 🔥 Updates
### 📅 October 2025
- Pretrained models, interactive demo, training code and data processing.
## 🔨 Installation
Clone the repo:
```bash
git clone https://github.com/HKU-MMLab/OmniPart
cd OmniPart
```
Create a conda environment (optional):
```bash
conda create -n omnipart python=3.10
conda activate omnipart
```
Install dependencies:
```bash
pip install -r requirements.txt
```
## 💡 Usage
### Launch Demo
```bash
python app.py
```
### Inference Scripts
If running OmniPart with command lines, you need to obtain the segmentation mask of the input image first. The mask is saved as a .exr file with the shape [h, w, 3], where the last dimension contains the 2D part_id replicated across all three channels.
```bash
python -m scripts.inference_omnipart --image_input {IMAGE_PATH} --mask_input {MASK_PATH}
```
The required model weights will be automatically downloaded:
- OmniPart model from [OmniPart](https://huggingface.co/omnipart) → local directory `ckpt/`
### Training
#### Data processing
Step 1: Render multi-view images of parts and overall shapes, following [TRELLIS Step 4](https://github.com/microsoft/TRELLIS/blob/main/DATASET.md#step-4-render-multiview-images).
Step 2: Voxelize parts and overall shapes with `dataset_toolkits/voxelize_part.py` and `dataset_toolkits/voxelize_overall.py`.
Step 3: Extract DINO features of parts and overall shapes, following [TRELLIS Step 6](https://github.com/microsoft/TRELLIS/blob/main/DATASET.md#step-6-extract-dino-features).
Step 4: Encode SLat of parts and overall shapes, following [TRELLIS Step 8](https://github.com/microsoft/TRELLIS/blob/main/DATASET.md#step-8-encode-slat).
Step 5: Merge SLat of parts and overall shapes with `dataset_toolkits/merge_slat.py`.
Step 6: Render image and mask conditions with `dataset_toolkits/blender_render_img_mask.py`.
#### Training code
Fill in the values for `data_root`, `train_mesh_list`, `val_mesh_list` and `denoiser` in `configs/training_part_synthesis.json`. The `denoiser` field requires the path to a diffusion model checkpoint in `.pt` format (using `training/utils/transfer_st_pt.py`) that you wish to finetune, for example: `ckpt/slat_flow_img_dit_L_64l8p2_fp16.pt`.
```bash
python train.py --config configs/training_part_synthesis.json --output_dir {OUTPUT_PATH} --data_dir {SLat_PATH}
```
## ⭐ Acknowledgements
We would like to thank the following open-source projects and research works that made OmniPart possible:
- [TRELLIS](https://github.com/microsoft/TRELLIS)
- [PartField](https://github.com/nv-tlabs/PartField)
- [FlexiCubes](https://github.com/nv-tlabs/FlexiCubes)
We are grateful to the broader research community for their open exploration and contributions to the field of 3D generation.
## 📚 Citation
```
@article{yang2025omnipart,
title={Omnipart: Part-aware 3d generation with semantic decoupling and structural cohesion},
author={Yang, Yunhan and Zhou, Yufan and Guo, Yuan-Chen and Zou, Zi-Xin and Huang, Yukun and Liu, Ying-Tian and Xu, Hao and Liang, Ding and Cao, Yan-Pei and Liu, Xihui},
journal={arXiv preprint arXiv:2507.06165},
year={2025}
}
```
+185
View File
@@ -0,0 +1,185 @@
import gradio as gr
import spaces
import os
import shutil
os.environ['SPCONV_ALGO'] = 'native'
from huggingface_hub import hf_hub_download
from app_utils import (
generate_parts,
prepare_models,
process_image,
apply_merge,
DEFAULT_SIZE_TH,
TMP_ROOT,
)
EXAMPLES = [
["assets/example_data/knight.png", 1800, "6,0,26,20,7;13,1,22,11,12,2,21,27,3,24,23;5,18;4,17;19,16,14,25,28", 42],
["assets/example_data/car.png", 2000, "12,10,2,11;1,7", 42],
["assets/example_data/warhammer.png", 1800, "7,1,0,8", 0],
["assets/example_data/snake.png", 3000, "2,3;0,1;4,5,6,7", 42],
["assets/example_data/Batman.png", 1800, "4,5", 42],
["assets/example_data/robot1.jpeg", 1600, "0,5;10,14,3;1,12,2;13,11,4;7,15", 42],
["assets/example_data/astronaut.png", 2000, "0,4,6;1,8,9,7;2,5", 42],
["assets/example_data/crossbow.jpg", 2000, "2,9;10,12,0,7,11,8,13;4,3", 42],
["assets/example_data/robot.jpg", 1600, "7,19;15,0;6,18", 42],
["assets/example_data/robot_dog.jpg", 1000, "21,9;2,12,10,15,17;11,7;1,0;13,19;4,16", 0],
["assets/example_data/crossbow.jpg", 1600, "9,2;10,15,13;7,14,8,11;0,12,16;5,3,1", 42],
["assets/example_data/robot.jpg", 1800, "1,2,3,5,4,16,17;11,7,19;10,14;18,6,0,15;13,9;12,8", 0],
["assets/example_data/robot_dog.jpg", 1000, "2,12,10,15,17,8,3,5,13,19,6,14;11,7;1,0,21,9,11;4,16", 0],
]
HEADER = """
# OmniPart: Part-Aware 3D Generation with Semantic Decoupling and Structural Cohesion
🔮 Generate **part-aware 3D content** from a single 2D image with **2D mask control**.
## How to Use
**🚀 Quick Start**: Select an example below and click **"▶️ Run Example"**
**📋 Custom Image Processing**:
1. **Upload Image** - Select your image file
2. **Click "Segment Image"** - Get initial 2D segmentation
3. **Merge Segments** - Enter merge groups like `0,1;3,4` and click **"Apply Merge"** (Recommend keeping **2-15 parts**)
4. **Click "Generate 3D Model"** - Create the final 3D results
"""
def start_session(req: gr.Request):
user_dir = os.path.join(TMP_ROOT, str(req.session_hash))
os.makedirs(user_dir, exist_ok=True)
def end_session(req: gr.Request):
user_dir = os.path.join(TMP_ROOT, str(req.session_hash))
shutil.rmtree(user_dir)
with gr.Blocks(title="OmniPart") as demo:
gr.Markdown(HEADER)
state = gr.State({})
with gr.Row():
with gr.Column(scale=1):
gr.Markdown("<div style='text-align: center'>\n\n## Input\n\n</div>")
input_image = gr.Image(label="Upload Image", type="filepath", height=250, width=250)
with gr.Row():
segment_btn = gr.Button("Segment Image", variant="primary", size="lg")
run_example_btn = gr.Button("▶️ Run Example", variant="secondary", size="lg")
size_threshold = gr.Slider(
minimum=600,
maximum=4000,
value=DEFAULT_SIZE_TH,
step=200,
label="Minimum Segment Size (pixels)",
info="Segments smaller than this will be ignored"
)
gr.Markdown("### Merge Controls")
merge_input = gr.Textbox(
label="Merge Groups",
placeholder="0,1;3,4",
lines=2,
info="Specify which segments to merge (e.g., '0,1;3,4' merges segments 0&1 together and 3&4 together)"
)
merge_btn = gr.Button("Apply Merge", variant="primary", size="lg")
gr.Markdown("### 3D Generation Controls")
seed_slider = gr.Slider(
minimum=0,
maximum=10000,
value=42,
step=1,
label="Generation Seed",
info="Random seed for 3D model generation"
)
cfg_slider = gr.Slider(
minimum=0.0,
maximum=15.0,
value=7.5,
step=0.5,
label="CFG Strength",
info="Classifier-Free Guidance strength"
)
generate_mesh_btn = gr.Button("Generate 3D Model", variant="secondary", size="lg")
with gr.Column(scale=2):
gr.Markdown("<div style='text-align: center'>\n\n## Results Display\n\n</div>")
with gr.Row():
initial_seg = gr.Image(label="Init Seg", height=220, width=220)
pre_merge_vis = gr.Image(label="Pre-merge", height=220, width=220)
merged_seg = gr.Image(label="Merged Seg", height=220, width=220)
with gr.Row():
bbox_mesh = gr.Model3D(label="Bounding Boxes", height=350)
whole_mesh = gr.Model3D(label="Combined Parts", height=350)
exploded_mesh = gr.Model3D(label="Exploded Parts", height=350)
with gr.Row():
combined_gs = gr.Model3D(label="Combined 3D Gaussians", clear_color=(0.0, 0.0, 0.0, 0.0), height=350)
exploded_gs = gr.Model3D(label="Exploded 3D Gaussians", clear_color=(0.0, 0.0, 0.0, 0.0), height=350)
with gr.Row():
examples = gr.Examples(
examples=EXAMPLES,
inputs=[input_image, size_threshold, merge_input, seed_slider],
cache_examples=False,
)
demo.load(start_session)
demo.unload(end_session)
segment_btn.click(
process_image,
inputs=[input_image, size_threshold],
outputs=[initial_seg, pre_merge_vis, state]
)
merge_btn.click(
apply_merge,
inputs=[merge_input, state],
outputs=[merged_seg, state]
)
generate_mesh_btn.click(
generate_parts,
inputs=[state, seed_slider, cfg_slider],
outputs=[bbox_mesh, whole_mesh, exploded_mesh, combined_gs, exploded_gs]
)
run_example_btn.click(
fn=process_image,
inputs=[input_image, size_threshold],
outputs=[initial_seg, pre_merge_vis, state]
).then(
fn=apply_merge,
inputs=[merge_input, state],
outputs=[merged_seg, state]
).then(
fn=generate_parts,
inputs=[state, seed_slider, cfg_slider],
outputs=[bbox_mesh, whole_mesh, exploded_mesh, combined_gs, exploded_gs]
)
if __name__ == "__main__":
os.makedirs("ckpt", exist_ok=True)
sam_ckpt_path = hf_hub_download(repo_id="omnipart/OmniPart_modules", filename="sam_vit_h_4b8939.pth", local_dir="ckpt")
partfield_ckpt_path = hf_hub_download(repo_id="omnipart/OmniPart_modules", filename="partfield_encoder.ckpt", local_dir="ckpt")
bbox_gen_ckpt_path = hf_hub_download(repo_id="omnipart/OmniPart_modules", filename="bbox_gen.ckpt", local_dir="ckpt")
prepare_models(sam_ckpt_path, partfield_ckpt_path, bbox_gen_ckpt_path)
port = int(os.getenv("PORT", "8080"))
demo.launch(share=False, server_name="0.0.0.0", server_port=port)
+412
View File
@@ -0,0 +1,412 @@
import gradio as gr
import spaces
import os
import numpy as np
import trimesh
import time
import traceback
import torch
from PIL import Image
import cv2
import shutil
from segment_anything import SamAutomaticMaskGenerator, build_sam
from omegaconf import OmegaConf
from modules.bbox_gen.models.autogressive_bbox_gen import BboxGen
from modules.part_synthesis.process_utils import save_parts_outputs
from modules.inference_utils import load_img_mask, prepare_bbox_gen_input, prepare_part_synthesis_input, gen_mesh_from_bounds, vis_voxel_coords, merge_parts
from modules.part_synthesis.pipelines import OmniPartImageTo3DPipeline
from modules.label_2d_mask.visualizer import Visualizer
from transformers import AutoModelForImageSegmentation
from modules.label_2d_mask.label_parts import (
prepare_image,
get_sam_mask,
get_mask,
clean_segment_edges,
resize_and_pad_to_square,
size_th as DEFAULT_SIZE_TH
)
# Constants
DEVICE = "cuda" if torch.cuda.is_available() else "cpu"
DTYPE = torch.float16
MAX_SEED = np.iinfo(np.int32).max
TMP_ROOT = os.path.join(os.path.dirname(os.path.abspath(__file__)), "tmp")
os.makedirs(TMP_ROOT, exist_ok=True)
sam_mask_generator = None
rmbg_model = None
bbox_gen_model = None
part_synthesis_pipeline = None
size_th = DEFAULT_SIZE_TH
def prepare_models(sam_ckpt_path, partfield_ckpt_path, bbox_gen_ckpt_path):
global sam_mask_generator, rmbg_model, bbox_gen_model, part_synthesis_pipeline
if sam_mask_generator is None:
print("Loading SAM model...")
sam_model = build_sam(checkpoint=sam_ckpt_path).to(device=DEVICE)
sam_mask_generator = SamAutomaticMaskGenerator(sam_model)
if rmbg_model is None:
print("Loading BriaRMBG 2.0 model...")
rmbg_model = AutoModelForImageSegmentation.from_pretrained('briaai/RMBG-2.0', trust_remote_code=True)
rmbg_model.to(DEVICE)
rmbg_model.eval()
if part_synthesis_pipeline is None:
print("Loading PartSynthesis model...")
part_synthesis_pipeline = OmniPartImageTo3DPipeline.from_pretrained('omnipart/OmniPart')
part_synthesis_pipeline.to(DEVICE)
if bbox_gen_model is None:
print("Loading BboxGen model...")
bbox_gen_config = OmegaConf.load("configs/bbox_gen.yaml").model.args
bbox_gen_config.partfield_encoder_path = partfield_ckpt_path
bbox_gen_model = BboxGen(bbox_gen_config)
bbox_gen_model.load_state_dict(torch.load(bbox_gen_ckpt_path), strict=False)
bbox_gen_model.to(DEVICE)
bbox_gen_model.eval().half()
print("Models ready")
@spaces.GPU
def process_image(image_path, threshold, req: gr.Request):
"""Process image and generate initial segmentation"""
global size_th
user_dir = os.path.join(TMP_ROOT, str(req.session_hash))
os.makedirs(user_dir, exist_ok=True)
img_name = os.path.basename(image_path).split(".")[0]
size_th = threshold
img = Image.open(image_path).convert("RGB")
processed_image = prepare_image(img, rmbg_net=rmbg_model.to(DEVICE))
processed_image = resize_and_pad_to_square(processed_image)
white_bg = Image.new("RGBA", processed_image.size, (255, 255, 255, 255))
white_bg_img = Image.alpha_composite(white_bg, processed_image.convert("RGBA"))
image = np.array(white_bg_img.convert('RGB'))
rgba_path = os.path.join(user_dir, f"{img_name}_processed.png")
processed_image.save(rgba_path)
print("Generating raw SAM masks without post-processing...")
raw_masks = sam_mask_generator.generate(image)
raw_sam_vis = np.copy(image)
raw_sam_vis = np.ones_like(image) * 255
sorted_masks = sorted(raw_masks, key=lambda x: x["area"], reverse=True)
for i, mask_data in enumerate(sorted_masks):
if mask_data["area"] < size_th:
continue
color_r = (i * 50 + 80) % 256
color_g = (i * 120 + 40) % 256
color_b = (i * 180 + 20) % 256
color = np.array([color_r, color_g, color_b])
mask = mask_data["segmentation"]
raw_sam_vis[mask] = color
visual = Visualizer(image)
group_ids, pre_merge_im = get_sam_mask(
image,
sam_mask_generator,
visual,
merge_groups=None,
rgba_image=processed_image,
img_name=img_name,
save_dir=user_dir,
size_threshold=size_th
)
pre_merge_path = os.path.join(user_dir, f"{img_name}_mask_pre_merge.png")
Image.fromarray(pre_merge_im).save(pre_merge_path)
pre_split_vis = np.ones_like(image) * 255
unique_ids = np.unique(group_ids)
unique_ids = unique_ids[unique_ids >= 0]
for i, unique_id in enumerate(unique_ids):
color_r = (i * 50 + 80) % 256
color_g = (i * 120 + 40) % 256
color_b = (i * 180 + 20) % 256
color = np.array([color_r, color_g, color_b])
mask = (group_ids == unique_id)
pre_split_vis[mask] = color
y_indices, x_indices = np.where(mask)
if len(y_indices) > 0:
center_y = int(np.mean(y_indices))
center_x = int(np.mean(x_indices))
cv2.putText(pre_split_vis, str(unique_id),
(center_x, center_y), cv2.FONT_HERSHEY_SIMPLEX,
0.5, (0, 0, 0), 1, cv2.LINE_AA)
pre_split_path = os.path.join(user_dir, f"{img_name}_pre_split.png")
Image.fromarray(pre_split_vis).save(pre_split_path)
print(f"Pre-split segmentation (before disconnected parts handling) saved to {pre_split_path}")
get_mask(group_ids, image, ids=2, img_name=img_name, save_dir=user_dir)
init_seg_path = os.path.join(user_dir, f"{img_name}_mask_segments_2.png")
seg_img = Image.open(init_seg_path)
if seg_img.mode == 'RGBA':
white_bg = Image.new('RGBA', seg_img.size, (255, 255, 255, 255))
seg_img = Image.alpha_composite(white_bg, seg_img)
seg_img.save(init_seg_path)
state = {
"image": image.tolist(),
"processed_image": rgba_path,
"group_ids": group_ids.tolist() if isinstance(group_ids, np.ndarray) else group_ids,
"original_group_ids": group_ids.tolist() if isinstance(group_ids, np.ndarray) else group_ids,
"img_name": img_name,
"pre_split_path": pre_split_path,
}
return init_seg_path, pre_merge_path, state
def apply_merge(merge_input, state, req: gr.Request):
"""Apply merge parameters and generate merged segmentation"""
global sam_mask_generator
if not state:
return None, None, state
user_dir = os.path.join(TMP_ROOT, str(req.session_hash))
# Convert back from list to numpy array
image = np.array(state["image"])
# Use original group IDs instead of the most recent ones
group_ids = np.array(state["original_group_ids"])
img_name = state["img_name"]
# Load processed image from path
processed_image = Image.open(state["processed_image"])
# Display the original IDs before merging, SORTED for easier reading
unique_ids = np.unique(group_ids)
unique_ids = unique_ids[unique_ids >= 0] # Exclude background
print(f"Original segment IDs (used for merging): {sorted(unique_ids.tolist())}")
# Parse merge groups
merge_groups = None
try:
if merge_input:
merge_groups = []
group_sets = merge_input.split(';')
for group_set in group_sets:
ids = [int(x) for x in group_set.split(',')]
if ids:
# Validate if these IDs exist in the segmentation
existing_ids = [id for id in ids if id in unique_ids]
missing_ids = [id for id in ids if id not in unique_ids]
if missing_ids:
print(f"Warning: These IDs don't exist in the segmentation: {missing_ids}")
# Only add group if it has valid IDs
if existing_ids:
merge_groups.append(ids)
print(f"Valid merge group: {ids} (missing: {missing_ids if missing_ids else 'none'})")
else:
print(f"Skipping merge group with no valid IDs: {ids}")
print(f"Using merge groups: {merge_groups}")
except Exception as e:
print(f"Error parsing merge groups: {e}")
return None, None, state
# Initialize visualizer
visual = Visualizer(image)
# Generate merged segmentation starting from original IDs
# Add skip_split=True to prevent splitting after merging
new_group_ids, merged_im = get_sam_mask(
image,
sam_mask_generator,
visual,
merge_groups=merge_groups,
existing_group_ids=group_ids,
rgba_image=processed_image,
skip_split=True,
img_name=img_name,
save_dir=user_dir,
size_threshold=size_th
)
# Display the new IDs after merging for future reference
new_unique_ids = np.unique(new_group_ids)
new_unique_ids = new_unique_ids[new_unique_ids >= 0] # Exclude background
print(f"New segment IDs (after merging): {new_unique_ids.tolist()}")
# Clean edges
new_group_ids = clean_segment_edges(new_group_ids)
# Save merged segmentation visualization
get_mask(new_group_ids, image, ids=3, img_name=img_name, save_dir=user_dir)
# Path to merged segmentation
merged_seg_path = os.path.join(user_dir, f"{img_name}_mask_segments_3.png")
save_mask = new_group_ids + 1
save_mask = save_mask.reshape(518, 518, 1).repeat(3, axis=-1)
cv2.imwrite(os.path.join(user_dir, f"{img_name}_mask.exr"), save_mask.astype(np.float32))
# Update state with the new group IDs but keep original IDs unchanged
state["group_ids"] = new_group_ids.tolist() if isinstance(new_group_ids, np.ndarray) else new_group_ids
state["save_mask_path"] = os.path.join(user_dir, f"{img_name}_mask.exr")
return merged_seg_path, state
def explode_mesh(mesh, explosion_scale=0.4):
if isinstance(mesh, trimesh.Scene):
scene = mesh
elif isinstance(mesh, trimesh.Trimesh):
print("Warning: Single mesh provided, can't create exploded view")
scene = trimesh.Scene(mesh)
return scene
else:
print(f"Warning: Unexpected mesh type: {type(mesh)}")
scene = mesh
if len(scene.geometry) <= 1:
print("Only one geometry found - nothing to explode")
return scene
print(f"[EXPLODE_MESH] Starting mesh explosion with scale {explosion_scale}")
print(f"[EXPLODE_MESH] Processing {len(scene.geometry)} parts")
exploded_scene = trimesh.Scene()
part_centers = []
geometry_names = []
for geometry_name, geometry in scene.geometry.items():
if hasattr(geometry, 'vertices'):
transform = scene.graph[geometry_name][0]
vertices_global = trimesh.transformations.transform_points(
geometry.vertices, transform)
center = np.mean(vertices_global, axis=0)
part_centers.append(center)
geometry_names.append(geometry_name)
print(f"[EXPLODE_MESH] Part {geometry_name}: center = {center}")
if not part_centers:
print("No valid geometries with vertices found")
return scene
part_centers = np.array(part_centers)
global_center = np.mean(part_centers, axis=0)
print(f"[EXPLODE_MESH] Global center: {global_center}")
for i, (geometry_name, geometry) in enumerate(scene.geometry.items()):
if hasattr(geometry, 'vertices'):
if i < len(part_centers):
part_center = part_centers[i]
direction = part_center - global_center
direction_norm = np.linalg.norm(direction)
if direction_norm > 1e-6:
direction = direction / direction_norm
else:
direction = np.random.randn(3)
direction = direction / np.linalg.norm(direction)
offset = direction * explosion_scale
else:
offset = np.zeros(3)
original_transform = scene.graph[geometry_name][0].copy()
new_transform = original_transform.copy()
new_transform[:3, 3] = new_transform[:3, 3] + offset
exploded_scene.add_geometry(
geometry,
transform=new_transform,
geom_name=geometry_name
)
print(f"[EXPLODE_MESH] Part {geometry_name}: moved by {np.linalg.norm(offset):.4f}")
print("[EXPLODE_MESH] Mesh explosion complete")
return exploded_scene
@spaces.GPU(duration=90)
def generate_parts(state, seed, cfg_strength, req: gr.Request):
explode_factor=0.3
img_path = state["processed_image"]
mask_path = state["save_mask_path"]
user_dir = os.path.join(TMP_ROOT, str(req.session_hash))
img_white_bg, img_black_bg, ordered_mask_input, img_mask_vis = load_img_mask(img_path, mask_path)
img_mask_vis.save(os.path.join(user_dir, "img_mask_vis.png"))
voxel_coords = part_synthesis_pipeline.get_coords(img_black_bg, num_samples=1, seed=seed, sparse_structure_sampler_params={"steps": 25, "cfg_strength": 7.5})
voxel_coords = voxel_coords.cpu().numpy()
np.save(os.path.join(user_dir, "voxel_coords.npy"), voxel_coords)
voxel_coords_ply = vis_voxel_coords(voxel_coords)
voxel_coords_ply.export(os.path.join(user_dir, "voxel_coords_vis.ply"))
print("[INFO] Voxel coordinates saved")
bbox_gen_input = prepare_bbox_gen_input(os.path.join(user_dir, "voxel_coords.npy"), img_white_bg, ordered_mask_input)
bbox_gen_output = bbox_gen_model.generate(bbox_gen_input)
np.save(os.path.join(user_dir, "bboxes.npy"), bbox_gen_output['bboxes'][0])
bboxes_vis = gen_mesh_from_bounds(bbox_gen_output['bboxes'][0])
bboxes_vis.export(os.path.join(user_dir, "bboxes_vis.glb"))
print("[INFO] BboxGen output saved")
part_synthesis_input = prepare_part_synthesis_input(os.path.join(user_dir, "voxel_coords.npy"), os.path.join(user_dir, "bboxes.npy"), ordered_mask_input)
torch.cuda.empty_cache()
part_synthesis_output = part_synthesis_pipeline.get_slat(
img_black_bg,
part_synthesis_input['coords'],
[part_synthesis_input['part_layouts']],
part_synthesis_input['masks'],
seed=seed,
slat_sampler_params={"steps": 25, "cfg_strength": cfg_strength},
formats=['mesh', 'gaussian'],
preprocess_image=False,
)
save_parts_outputs(
part_synthesis_output,
output_dir=user_dir,
simplify_ratio=0.0,
save_video=False,
save_glb=True,
textured=False,
)
merge_parts(user_dir)
print("[INFO] PartSynthesis output saved")
bbox_mesh_path = os.path.join(user_dir, "bboxes_vis.glb")
whole_mesh_path = os.path.join(user_dir, "mesh_segment.glb")
combined_mesh = trimesh.load(whole_mesh_path)
exploded_mesh_result = explode_mesh(combined_mesh, explosion_scale=explode_factor)
exploded_mesh_result.export(os.path.join(user_dir, "exploded_parts.glb"))
exploded_mesh_path = os.path.join(user_dir, "exploded_parts.glb")
combined_gs_path = os.path.join(user_dir, "merged_gs.ply")
exploded_gs_path = os.path.join(user_dir, "exploded_gs.ply")
return bbox_mesh_path, whole_mesh_path, exploded_mesh_path, combined_gs_path, exploded_gs_path
Binary file not shown.

After

Width:  |  Height:  |  Size: 1.1 MiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 20 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 83 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 70 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 28 KiB

+11
View File
@@ -0,0 +1,11 @@
from transformers import PretrainedConfig
class BiRefNetConfig(PretrainedConfig):
model_type = "SegformerForSemanticSegmentation"
def __init__(
self,
bb_pretrained=False,
**kwargs
):
self.bb_pretrained = bb_pretrained
super().__init__(**kwargs)
File diff suppressed because it is too large Load Diff
+20
View File
@@ -0,0 +1,20 @@
{
"_name_or_path": "ZhengPeng7/BiRefNet",
"architectures": [
"BiRefNet"
],
"auto_map": {
"AutoConfig": "BiRefNet_config.BiRefNetConfig",
"AutoModelForImageSegmentation": "birefnet.BiRefNet"
},
"custom_pipelines": {
"image-segmentation": {
"pt": [
"AutoModelForImageSegmentation"
],
"tf": [],
"type": "image"
}
},
"bb_pretrained": false
}
+34
View File
@@ -0,0 +1,34 @@
model:
name: bbox_gen
args:
encoder_dim_feat: 448
encoder_dim: 64
encoder_heads: 4
encoder_token_num: 2048
encoder_qkv_bias: false
encoder_use_ln_post: true
encoder_use_checkpoint: true
encoder_num_embed_freqs: 8
encoder_embed_include_pi: false
encoder_init_scale: 0.25
encoder_random_fps: true
encoder_learnable_query: true
encoder_layers: 8
max_group_size: 50
vocab_size: 67
decoder_hidden_size: 1024
decoder_num_hidden_layers: 24
decoder_ffn_dim: 4096
decoder_heads: 16
decoder_use_flash_attention: true
decoder_gradient_checkpointing: false
bins: 64
BOS_id: 64
EOS_id: 65
PAD_id: 66
max_length: 2187
voxel_token_length: 1886
voxel_token_placeholder: -1
@@ -0,0 +1,50 @@
{
"apply_layernorm": true,
"architectures": [
"Dinov2WithRegistersModel"
],
"attention_probs_dropout_prob": 0.0,
"drop_path_rate": 0.0,
"hidden_act": "gelu",
"hidden_dropout_prob": 0.0,
"hidden_size": 1024,
"image_size": 518,
"initializer_range": 0.02,
"interpolate_antialias": true,
"interpolate_offset": 0.0,
"layer_norm_eps": 1e-06,
"layerscale_value": 1.0,
"mlp_ratio": 4,
"model_type": "dinov2_with_registers",
"num_attention_heads": 16,
"num_channels": 3,
"num_hidden_layers": 24,
"num_register_tokens": 4,
"out_features": [
"stage12"
],
"out_indices": [
12
],
"patch_size": 14,
"qkv_bias": true,
"reshape_hidden_states": true,
"stage_names": [
"stem",
"stage1",
"stage2",
"stage3",
"stage4",
"stage5",
"stage6",
"stage7",
"stage8",
"stage9",
"stage10",
"stage11",
"stage12"
],
"torch_dtype": "float32",
"transformers_version": "4.48.0.dev0",
"use_swiglu_ffn": false
}
@@ -0,0 +1,27 @@
{
"crop_size": {
"height": 224,
"width": 224
},
"do_center_crop": true,
"do_convert_rgb": true,
"do_normalize": true,
"do_rescale": true,
"do_resize": true,
"image_mean": [
0.485,
0.456,
0.406
],
"image_processor_type": "BitImageProcessor",
"image_std": [
0.229,
0.224,
0.225
],
"resample": 3,
"rescale_factor": 0.00392156862745098,
"size": {
"shortest_edge": 256
}
}
@@ -0,0 +1,114 @@
{
"models": {
"denoiser": {
"name": "ElasticSLatFlowModel",
"args": {
"resolution": 64,
"in_channels": 9,
"out_channels": 9,
"model_channels": 1024,
"cond_channels": 1024,
"num_blocks": 24,
"num_heads": 16,
"mlp_ratio": 4,
"patch_size": 2,
"num_io_res_blocks": 2,
"io_block_channels": [
128
],
"pe_mode": "ape",
"qk_rms_norm": true,
"use_fp16": true
}
}
},
"dataset": {
"name": "ImageConditionedSLat",
"args": {
"data_root": "",
"train_mesh_list": "",
"val_mesh_list": "",
"aug_bbox": "0 2",
"latent_model": "dinov2_vitl14_reg_slat_enc_swin8_B_64l8_fp16",
"min_aesthetic_score": 4.5,
"max_num_voxels": 65536,
"image_size": 518,
"normalization": {
"mean": [
-2.1687545776367188,
-0.004347046371549368,
-0.13352349400520325,
-0.08418072760105133,
-0.5271206498146057,
0.7238689064979553,
-1.1414450407028198,
1.2039363384246826,
0.0
],
"std": [
2.377650737762451,
2.386378288269043,
2.124418020248413,
2.1748552322387695,
2.663944721221924,
2.371192216873169,
2.6217446327209473,
2.684523105621338,
1.0
]
},
"pretrained_slat_dec": "JeffreyXiang/TRELLIS-image-large/ckpts/slat_dec_gs_swin8_B_64l8gs32_fp16"
}
},
"trainer": {
"name": "ImageConditionedSparseFlowMatchingCFGTrainer",
"args": {
"max_steps": 1000000,
"batch_size_per_gpu": 12,
"batch_split": 4,
"optimizer": {
"name": "AdamW",
"args": {
"lr": 0.0001,
"weight_decay": 0.0
}
},
"ema_rate": [
0.9999
],
"fp16_mode": "inflat_all",
"fp16_scale_growth": 0.001,
"elastic": {
"name": "LinearMemoryController",
"args": {
"target_ratio": 0.75,
"max_mem_ratio_start": 0.5
}
},
"grad_clip": {
"name": "AdaptiveGradClipper",
"args": {
"max_norm": 1.0,
"clip_percentile": 95
}
},
"i_print": 50,
"i_log": 500,
"i_sample": 2000,
"i_save": 10000,
"p_uncond": 0.1,
"t_schedule": {
"name": "logitNormal",
"args": {
"mean": 1.0,
"std": 1.0
}
},
"sigma_min": 1e-5,
"image_cond_model": "dinov2_vitl14_reg",
"finetune_ckpt": {
"denoiser": "TRELLIS-image-large-pt/slat_flow_img_dit_L_64l8p2_fp16.pt"
}
}
}
}
@@ -0,0 +1,646 @@
import bpy
import random
import sys
from mathutils import Vector
import os
os.environ["OPENCV_IO_ENABLE_OPENEXR"] = "1"
import glob
import math
import json
import mathutils
from mathutils import Matrix, Vector, Euler
from typing import List, Optional
import numpy as np
import cv2
from PIL import Image
META_FILENAME = "meta.json"
RESOLUTION = 512
RENDER_SAMPLES = 64
ENGINE = "EEVEE"
CLOSE_SHADOW = False
RENDER_DEPTH = False
if os.environ.get("ENGINE"):
ENGINE = os.environ.get("ENGINE")
if os.environ.get("RENDER_SAMPLES"):
RENDER_SAMPLES = int(os.environ.get("RENDER_SAMPLES"))
if os.environ.get("RESOLUTION"):
RESOLUTION = int(os.environ.get("RESOLUTION"))
def generate_random(left, right):
while True:
val = random.gauss(0, (right-left)/5)
if val >= left and val <= right:
return val
def generate_frames(name):
render_type_list = [
{'name': "depth", "suffix": ".exr", "enable": RENDER_DEPTH},
{'name': "render_opaque", "suffix": ".webp",
"enable": RENDER_DEPTH},
]
param = [
{
"type": "render",
"name": "{}",
"height": RESOLUTION,
"width": RESOLUTION,
}
]
variables = ["render_"+name+".webp"]
for render_type in render_type_list:
if render_type["enable"]:
param.append({
"type": render_type["name"],
"name": "{}",
"height": RESOLUTION,
"width": RESOLUTION,
})
variables.append(render_type["name"] +
"_"+name+render_type["suffix"])
for i in range(len(variables)):
param[i]["name"] = variables[i]
return param
def build_transformation_mat(translation,
rotation) -> np.ndarray:
""" Build a transformation matrix from translation and rotation parts.
:param translation: A (3,) vector representing the translation part.
:param rotation: A 3x3 rotation matrix or Euler angles of shape (3,).
:return: The 4x4 transformation matrix.
"""
translation = np.array(translation)
rotation = np.array(rotation)
mat = np.eye(4)
if translation.shape[0] == 3:
mat[:3, 3] = translation
else:
raise RuntimeError(
f"Translation has invalid shape: {translation.shape}. Must be (3,) or (3,1) vector.")
if rotation.shape == (3, 3):
mat[:3, :3] = rotation
elif rotation.shape[0] == 3:
mat[:3, :3] = np.array(Euler(rotation).to_matrix())
else:
raise RuntimeError(f"Rotation has invalid shape: {rotation.shape}. Must be rotation matrix of shape "
f"(3,3) or Euler angles of shape (3,) or (3,1).")
return mat
def reset_keyframes() -> None:
""" Removes registered keyframes from all objects and resets frame_start and frame_end """
bpy.context.scene.frame_start = 0
bpy.context.scene.frame_end = 0
for a in bpy.data.actions:
bpy.data.actions.remove(a)
def get_local2world_mat(blender_obj) -> np.ndarray:
""" Returns the pose of the object in the form of a local2world matrix.
:return: The 4x4 local2world matrix.
"""
obj = blender_obj
# Start with local2parent matrix (if obj has no parent, that equals local2world)
matrix_world = obj.matrix_basis
# Go up the scene graph along all parents
while obj.parent is not None:
# Add transformation to parent frame
matrix_world = obj.parent.matrix_basis @ obj.matrix_parent_inverse @ matrix_world
obj = obj.parent
return np.array(matrix_world)
def add_camera(cam2world_matrix,camera_params) -> int:
if not isinstance(cam2world_matrix, Matrix):
cam2world_matrix = Matrix(cam2world_matrix)
bpy.ops.object.camera_add(location=(0, 0, 0))
cam_ob = bpy.context.object
cam_ob.matrix_world = cam2world_matrix
cam_ob_data = cam_ob.data
cam_ob_data.type = camera_params['camera_type']
cam_ob_data.sensor_width = camera_params['camera_sensor_width']
if camera_params['camera_type'] == 'ORTHO':
cam_ob_data.ortho_scale = camera_params['camera_ortho_scale']
elif camera_params['camera_type'] == 'PERSP':
cam_ob_data.lens = camera_params['camera_lens']
def add_camera_pose(cam2world_matrix, camera_params) -> int:
if not isinstance(cam2world_matrix, Matrix):
cam2world_matrix = Matrix(cam2world_matrix)
cam_ob = bpy.context.scene.camera
cam_ob.matrix_world = cam2world_matrix
cam_ob_data = cam_ob.data
cam_ob_data.type = camera_params['camera_type']
cam_ob_data.sensor_width = camera_params['camera_sensor_width']
if camera_params['camera_type'] == 'ORTHO':
cam_ob_data.ortho_scale = camera_params['camera_ortho_scale']
elif camera_params['camera_type'] == 'PERSP':
cam_ob_data.lens = camera_params['camera_lens']
frame = bpy.context.scene.frame_end
if bpy.context.scene.frame_end < frame + 1:
bpy.context.scene.frame_end = frame + 1
cam_ob.keyframe_insert(data_path='location', frame=frame)
cam_ob.keyframe_insert(data_path='rotation_euler', frame=frame)
cam_ob_data.keyframe_insert(data_path='type', frame=frame)
cam_ob_data.keyframe_insert(data_path='sensor_width', frame=frame)
if camera_params['camera_type'] == 'ORTHO':
cam_ob_data.keyframe_insert(data_path='ortho_scale', frame=frame)
elif camera_params['camera_type'] == 'PERSP':
cam_ob_data.keyframe_insert(data_path='lens', frame=frame)
return frame
def clear_normal_map():
for material in bpy.data.materials:
material.use_nodes = True
node_tree = material.node_tree
try:
bsdf = node_tree.nodes["Principled BSDF"]
if bsdf.inputs["Normal"].is_linked:
for link in bsdf.inputs["Normal"].links:
node_tree.links.remove(link)
except:
pass
def enable_depth_output(output_dir: Optional[str] = '', file_prefix: str = "depth_"):
bpy.context.scene.render.use_compositing = True
bpy.context.scene.use_nodes = True
tree = bpy.context.scene.node_tree
links = tree.links
if "Render Layers" not in tree.nodes:
rl = tree.nodes.new('CompositorNodeRLayers')
else:
rl = tree.nodes["Render Layers"]
bpy.context.view_layer.use_pass_z = True
depth_output = tree.nodes.new('CompositorNodeOutputFile')
depth_output.base_path = output_dir
depth_output.name = 'DepthOutput'
depth_output.format.file_format = 'OPEN_EXR'
depth_output.format.color_depth = '32'
depth_output.file_slots.values()[0].path = file_prefix
links.new(rl.outputs["Depth"], depth_output.inputs['Image'])
def scene_mesh_objects(override_context: bool = False):
bpy.ops.object.select_all(action="DESELECT")
for obj in bpy.context.scene.objects:
if obj.type == 'MESH' and obj.visible_get() is True and obj.hide_get() is False:
yield obj
def enable_mask_output(output_dir: Optional[str] = '', file_prefix: str = "mask_"):
bpy.context.view_layer.use_pass_cryptomatte_object = True
tree = bpy.context.scene.node_tree
links = tree.links
for obj in scene_mesh_objects():
crypto_node = tree.nodes.new('CompositorNodeCryptomatteV2')
crypto_node.name = f"Cryptomatte_{obj.name}"
crypto_node.matte_id = obj.name
id_mask_output_node = tree.nodes.new('CompositorNodeOutputFile')
id_mask_output_node.name = f'Mask Output_{obj.name}'
id_mask_output_node.format.file_format = 'WEBP'
id_mask_output_node.format.quality = 100
id_mask_output_node.format.color_depth = '8'
id_mask_output_node.base_path = output_dir
id_mask_output_node.file_slots.values()[0].path = f"{file_prefix}{obj.name}_"
links.new(crypto_node.outputs['Matte'], id_mask_output_node.inputs['Image'])
def format_mask_output(output_dir):
scene = bpy.context.scene
h, w = scene.render.resolution_y, scene.render.resolution_x
webp_list = os.listdir(output_dir)
mask_list = [img for img in webp_list if img.startswith("mask_") and img.endswith(".webp")]
frame_groups = {}
for mask_file in mask_list:
parts = mask_file.replace(".webp", "").split("_")
frame_idx = int(parts[-1])
part_name = "_".join(parts[1:-1])
if frame_idx not in frame_groups:
frame_groups[frame_idx] = []
frame_groups[frame_idx].append((part_name, mask_file))
for frame_idx in sorted(frame_groups.keys()):
save_mask = np.zeros((h, w, 3), dtype=np.float32)
part_to_index = {}
parts_in_frame = sorted(frame_groups[frame_idx], key=lambda x: x[0])
for i, (part_name, mask_file) in enumerate(parts_in_frame):
mask_index = i + 1
part_to_index[part_name] = mask_index
mask_path = os.path.join(output_dir, mask_file)
mask_img = np.array(Image.open(mask_path))
if mask_img.shape[-1] == 4:
valid_region = mask_img[..., 3] > 0
else:
valid_region = mask_img.sum(axis=-1) > 0
save_mask[valid_region] = [mask_index, mask_index, mask_index]
output_file = os.path.join(output_dir, f"mask_{frame_idx:04d}.exr")
cv2.imwrite(output_file, save_mask)
print(f"Mask output saved to {output_dir}")
for mask_path in mask_list:
os.remove(os.path.join(output_dir, mask_path))
def render():
bpy.context.scene.render.use_compositing = True
bpy.context.scene.use_nodes = True
tree = bpy.context.scene.node_tree
links = tree.links
if "Render Layers" not in tree.nodes:
rl = tree.nodes.new('CompositorNodeRLayers')
else:
rl = tree.nodes["Render Layers"]
if bpy.context.scene.frame_end != bpy.context.scene.frame_start:
bpy.context.scene.frame_end -= 1
bpy.ops.render.render(animation=True, write_still=True)
bpy.context.scene.frame_end += 1
else:
raise RuntimeError("No camera poses have been registered, therefore nothing can be rendered. A camera "
"pose can be registered via bproc.camera.add_camera_pose().")
def convert_position(location, center):
position = ""
axis = ['x', 'y', 'z']
sub = location-center
for i in range(len(axis)):
if sub[i] > 0:
position = position + "+" + axis[i]
elif sub[i] < 0:
position = position + "-" + axis[i]
return position
def set_color_output(output_dir: Optional[str] = '', file_prefix: str = "render_"):
scene = bpy.context.scene
scene.render.use_compositing = True
scene.use_nodes = True
scene.render.resolution_x = RESOLUTION
scene.render.resolution_y = RESOLUTION
scene.render.image_settings.file_format = 'WEBP'
scene.render.image_settings.quality = 100
scene.render.image_settings.color_mode = 'RGBA'
# scene.render.image_settings.color_depth = '16'
scene.render.film_transparent = True
scene.render.filepath = os.path.join(output_dir, file_prefix)
pass
def eevee_init():
bpy.context.scene.render.engine = 'BLENDER_EEVEE'
bpy.context.scene.eevee.taa_render_samples = RENDER_SAMPLES
if CLOSE_SHADOW == False:
bpy.context.scene.eevee.use_gtao = True
bpy.context.scene.eevee.use_ssr = True
bpy.context.scene.render.use_high_quality_normals = True
def clear_scene(NOT_CLEAR_LIGHT=False):
bpy.ops.object.select_all(action="DESELECT")
if NOT_CLEAR_LIGHT:
for obj in bpy.data.objects:
if obj.type not in {"CAMERA", "LIGHT"}:
bpy.data.objects.remove(obj, do_unlink=True)
else:
bpy.ops.object.select_all(action='SELECT')
bpy.ops.object.delete()
bpy.context.scene.use_nodes = True
node_tree = bpy.context.scene.node_tree
for node in node_tree.nodes:
node_tree.nodes.remove(node)
reset_keyframes()
def import_models(filepath, types):
if types == "glb":
bpy.ops.import_scene.gltf(
filepath=filepath)
elif types == "obj":
forward_axis = os.environ.get('FORWARD_AXIS', 'NEGATIVE_Z')
up_axis = os.environ.get('UP_AXIS', 'Y')
bpy.ops.wm.obj_import(filepath=filepath, directory=os.path.dirname(filepath),forward_axis=forward_axis, up_axis=up_axis)
def rotation_matrix(x_left, x_right, y_left, y_right):
x_rotation = math.radians(generate_random(x_left, x_right))
y_rotation = math.radians(generate_random(y_left, y_right))
x_rotation_matrix = mathutils.Matrix.Rotation(x_rotation, 4, 'X')
y_rotation_matrix = mathutils.Matrix.Rotation(y_rotation, 4, 'Y')
final_rotation_matrix = y_rotation_matrix @ x_rotation_matrix
return final_rotation_matrix
def scene_bbox(objects=None, ignore_small_obj=False, ignore_matrix=False):
bbox_min = (math.inf,) * 3
bbox_max = (-math.inf,) * 3
found = False
for obj in objects:
# print(max(obj.dimensions*100))
if max(obj.dimensions*100) < 0.1 and ignore_small_obj:
print("ignore_small_obj", obj.name,max(obj.dimensions*100))
continue
found = True
for coord in obj.bound_box:
# print(coord[0], coord[1], coord[2])
coord = Vector(coord)
if not ignore_matrix:
coord = obj.matrix_world @ coord
bbox_min = Vector(
(min(bbox_min[i], coord[i]) for i in range(3)))
bbox_max = Vector(
(max(bbox_max[i], coord[i]) for i in range(3)))
if not found:
raise RuntimeError("no objects in scene to compute bounding box for")
return Vector(bbox_min), Vector(bbox_max)
def scene_root_objects():
for obj in bpy.context.scene.objects.values():
if not obj.parent:
yield obj
def set_global_light(env_light=0.5):
world_tree = bpy.context.scene.world.node_tree
back_node = world_tree.nodes["Background"]
back_node.inputs["Color"].default_value = Vector(
[env_light, env_light, env_light, 1.0]
)
back_node.inputs["Strength"].default_value = 1.0
def normalize_scene(normalization_range, objects):
bpy.ops.object.empty_add(type='PLAIN_AXES')
root_object = bpy.context.object
for obj in scene_root_objects():
if obj != root_object:
_matrix_world = obj.matrix_world.copy()
obj.parent = root_object
obj.matrix_world = _matrix_world
bpy.context.view_layer.update()
bbox_min, bbox_max = scene_bbox(objects)
scale = normalization_range / max(bbox_max - bbox_min)
root_object.scale *= scale
bpy.context.view_layer.update()
bbox_min, bbox_max = scene_bbox(objects,True)
mesh_offset = - (bbox_min + bbox_max) / 2
root_object.matrix_local.translation = mesh_offset
bpy.context.view_layer.update()
bpy.ops.object.select_all(action="DESELECT")
return root_object, bbox_max - bbox_min, scale, mesh_offset
def compute_bounding_box(mesh_objects):
min_coords = Vector((float('inf'), float('inf'), float('inf')))
max_coords = Vector((float('-inf'), float('-inf'), float('-inf')))
for obj in mesh_objects:
matrix_world = obj.matrix_world
mesh = obj.data
for vert in mesh.vertices:
global_coord = matrix_world @ vert.co
min_coords = Vector(
(min(min_coords[i], global_coord[i]) for i in range(3)))
max_coords = Vector(
(max(max_coords[i], global_coord[i]) for i in range(3)))
bbox_center = (min_coords + max_coords) / 2
bbox_size = max_coords - min_coords
return bbox_center, bbox_size
def change_material_blend_mode():
for material in bpy.data.materials:
material.use_nodes = True
node_tree = material.node_tree
material.blend_method = 'OPAQUE'
def change_material_blend_show_transparent(value):
for material in bpy.data.materials:
material.use_nodes = True
if material.blend_method == 'BLEND':
material.show_transparent_back = value
def get_random_points_on_sphere(center, radius, num_points=8):
points = []
for i in range(num_points):
r = radius
theta = random.uniform(0, 2*math.pi)
phi = random.uniform(0, 0.5*math.pi)
x = center[0] + r * math.sin(phi) * math.cos(theta)
y = center[1] + r * math.sin(phi) * math.sin(theta)
z = center[2] + r * math.cos(phi)
flag = -1 if bool(random.randint(0, 1)) else 1
points.append(Vector((x, y, z*flag)))
return points
def get_solid_points_on_sphere(center, radius):
points = []
elev_list = [25., 25., 25., 25., 25., 25., 25., 25.]
azim_list = [0, 45, 90, 135, 180, 225, 270, 315]
for i in range(len(elev_list)):
x = center[0] + radius * math.cos(math.radians(elev_list[i])) * math.cos(math.radians(azim_list[i]))
y = center[1] + radius * math.cos(math.radians(elev_list[i])) * math.sin(math.radians(azim_list[i]))
z = center[2] + radius * math.sin(math.radians(elev_list[i]))
points.append(Vector((x, y, z)))
return points
def listify_matrix(matrix):
matrix_list = []
for row in matrix:
matrix_list.append(list(row))
return matrix_list
def process(filepath, types, output_path):
random.seed()
eevee_init()
clear_scene()
import_models(filepath, types)
reset_keyframes()
bpy.ops.object.select_by_type(type='MESH')
os.makedirs(output_path, exist_ok=True)
mesh_objects = []
for obj in bpy.context.scene.objects:
if obj.type == 'MESH' and obj.visible_get() == True and obj.hide_get() == False:
mesh_objects.append(obj)
bpy.ops.object.select_all(action="DESELECT")
bpy.ops.object.select_pattern(pattern=obj.name)
for obj in mesh_objects:
obj.data.use_auto_smooth = True
obj.data.auto_smooth_angle = np.deg2rad(30)
for obj in bpy.data.objects:
if obj.animation_data is not None:
obj.animation_data_clear()
clear_normal_map()
change_material_blend_show_transparent(False)
normalization_range = 1.0
root_object, bbox_size, scale, mesh_offset = normalize_scene(normalization_range,mesh_objects)
bpy.context.view_layer.update()
root_object.rotation_euler[2] = math.radians(int(os.environ.get("FORCE_ROTATION", 0)))
bbox_center = Vector((0,0,0))
bpy.ops.object.camera_add(location=(0, 0, 0))
bpy.context.scene.camera = bpy.context.object
default_camera_lens = 50
default_camera_senser_width = 36
default_camera_ortho_scale = 1.4
ratio = 1
distance = ratio * default_camera_lens / default_camera_senser_width * \
math.sqrt(bbox_size.x**2 + bbox_size.y**2+bbox_size.z**2)
idx = 0
env_texture = "null"
set_global_light(env_light=0.5)
camera_angle_x = 2.0*math.atan(default_camera_senser_width/2/default_camera_lens)
out_data = {
'camera_angle_x': camera_angle_x,
'camera_lens': default_camera_lens,
'sensor_width': default_camera_senser_width,
'env_texture': env_texture,
'bbox_size': list(bbox_size),
'scaling_factor': scale,
'mesh_offset': list(mesh_offset),
'transforms': []
}
parent_matrix_list = [
rotation_matrix(0, 0, 0, 0),
]
camera_locations = get_solid_points_on_sphere(
bbox_center, distance)
camera_locations_random = get_random_points_on_sphere(bbox_center, distance)
camera_locations = camera_locations + camera_locations_random
positon_tag = [convert_position(camera_location,bbox_center) for camera_location in camera_locations]
for parent_matrix in parent_matrix_list:
camera_idx = 0
for camera_location in camera_locations:
_lens = 50
_camera_location = camera_location * \
(_lens / default_camera_lens)
_rotation_euler = (
bbox_center - _camera_location).to_track_quat('-Z', 'Y').to_euler()
cam_matrix = build_transformation_mat(
_camera_location, _rotation_euler)
cam_matrix = listify_matrix(parent_matrix) @ cam_matrix
camera_params = {
'camera_type': 'PERSP',
'camera_lens': _lens,
'camera_sensor_width': default_camera_senser_width,
}
# add_camera(cam_matrix,camera_params)
add_camera_pose(cam_matrix, camera_params)
index = "{0:04d}".format(idx)
out_data['transforms'].append(listify_matrix(cam_matrix))
idx += 1
camera_idx += 1
set_color_output(output_dir=output_path)
enable_mask_output(output_dir=output_path, file_prefix="mask_")
render()
format_mask_output(output_dir=output_path)
render_opaque_flag = RENDER_DEPTH
if render_opaque_flag:
change_material_blend_mode()
set_color_output(output_dir=output_path, file_prefix="render_opaque_")
if RENDER_DEPTH:
enable_depth_output(output_dir=output_path)
if render_opaque_flag:
render()
with open(os.path.join(output_path, META_FILENAME), 'w') as out_file:
json.dump(out_data, out_file, indent=4)
file_prefix = "render_opaque_"
pattern = os.path.join(output_path, f'{file_prefix}*')
files_to_delete = glob.glob(pattern)
for file in files_to_delete:
os.remove(file)
if __name__ == "__main__":
if len(sys.argv) < 1:
print(
"Usage: [path_to_blender] -b -P blender_render_img_mask.py [mesh_path] [types] [output_path]")
exit(1)
else:
mesh_path = sys.argv[4]
types = sys.argv[5]
output_path = sys.argv[6]
ret = process(mesh_path, types, output_path)
if ret:
exit(0)
else:
exit(1)
+51
View File
@@ -0,0 +1,51 @@
from tqdm import tqdm
from multiprocessing import Pool
import os
import numpy as np
def merge_part_latents(args):
uuid, output_dir = args
part_dir = os.path.join(output_dir, uuid[:2], uuid)
valid_part_id_path = os.path.join(part_dir, 'overall', 'part_id.txt')
with open(valid_part_id_path, 'r') as f:
part_ids = f.read().strip().splitlines()
all_data_coord = []
all_data_feat = []
all_data_offset = [0]
overall_save_path = os.path.join(part_dir, 'overall', 'latent.npz')
overall_latent = np.load(overall_save_path)
all_data_coord.append(overall_latent['coords'])
all_data_feat.append(overall_latent['feats'])
all_data_offset.append(overall_latent['coords'].shape[0])
for part_id in part_ids:
part_latent_path = os.path.join(part_dir, part_id, 'latent.npz')
part_latent = np.load(part_latent_path)
all_data_coord.append(part_latent['coords'])
all_data_feat.append(part_latent['feats'])
all_data_offset.append(all_data_offset[-1] + part_latent['coords'].shape[0])
all_data_coord = np.concatenate(all_data_coord, axis=0)
all_data_feat = np.concatenate(all_data_feat, axis=0)
all_data_offset = np.array(all_data_offset)
save_dict = {
'coords': all_data_coord,
'feats': all_data_feat,
'offsets': all_data_offset
}
save_path = os.path.join(part_dir, 'all_latent.npz')
np.savez_compressed(save_path, **save_dict)
if __name__ == "__main__":
valid_uuid_path = ''
output_dir = ''
with open(valid_uuid_path, 'r') as f:
valid_uuids = [line.strip() for line in f.readlines()]
args_list = []
for uuid in valid_uuids:
args_list.append((uuid, output_dir))
with Pool(64) as p:
results = list(tqdm(p.imap(merge_part_latents, args_list), total=len(args_list)))
@@ -0,0 +1,68 @@
import os
import copy
import sys
import importlib
import argparse
from easydict import EasyDict as edict
import pandas as pd
from functools import partial
import numpy as np
import open3d as o3d
import utils3d
from multiprocessing import Pool
from tqdm import tqdm
import json
import glob
import utils3d
def _voxelize(args):
sha256, output_dir = args
uuid_path = os.path.join(output_dir, sha256[:2], sha256)
pattern = os.path.join(uuid_path, "[0-9][0-9][0-9][0-9]")
matching_dirs = glob.glob(pattern)
part_id_list = [os.path.basename(dir_path) for dir_path in matching_dirs]
part_num = len(part_id_list)
filter_small_part_th = 5
part_voxel_list = []
new_part_id_list = []
for i, (dir_path, part_id) in enumerate(zip(matching_dirs, part_id_list)):
part_path = os.path.join(dir_path, 'voxel.ply')
voxel = utils3d.io.read_ply(part_path)[0]
if len(voxel) <= filter_small_part_th:
continue
part_voxel_list.append(voxel)
new_part_id_list.append(part_id)
if len(part_voxel_list) == 0:
print(f"Error: No valid parts found for {sha256}")
return
combined = list(zip(part_voxel_list, new_part_id_list))
sorted_combined = sorted(combined, key=lambda x: x[0].min(axis=0)[2])
part_voxel_list, new_part_id_list = zip(*sorted_combined)
overall_voxel = np.vstack(part_voxel_list)
overall_voxel = np.unique(overall_voxel, axis=0)
utils3d.io.write_ply(os.path.join(uuid_path, 'overall', f'voxel.ply'), overall_voxel)
with open(os.path.join(uuid_path, 'overall', f'part_id.txt'), 'w') as f:
for part_id in new_part_id_list:
f.write(f"{part_id}\n")
if __name__ == '__main__':
data_path = ''
output_dir = ''
with open(data_path, 'r') as f:
data_list = [json.loads(line.strip()) for line in f.readlines()]
args_list = [(sha256, output_dir) for sha256 in data_list]
with Pool(128) as p:
results = list(tqdm(p.imap(_voxelize, args_list), total=len(args_list)))
@@ -0,0 +1,49 @@
import os
import copy
import sys
import importlib
import argparse
from easydict import EasyDict as edict
import pandas as pd
from functools import partial
import numpy as np
import open3d as o3d
import utils3d
from multiprocessing import Pool
from tqdm import tqdm
import json
import glob
def _voxelize(args):
sha256, output_dir = args
uuid_path = os.path.join(output_dir, sha256[:2], sha256)
pattern = os.path.join(uuid_path, "[0-9][0-9][0-9][0-9]")
matching_dirs = glob.glob(pattern)
for dir_path in matching_dirs:
mesh = o3d.io.read_triangle_mesh(os.path.join(dir_path, 'mesh.stl'))
# clamp vertices to the range [-0.5, 0.5]
vertices = np.clip(np.asarray(mesh.vertices), -0.5 + 1e-6, 0.5 - 1e-6)
assert len(vertices)>0, "Error loading mesh.stl, no vertices found"
mesh.vertices = o3d.utility.Vector3dVector(vertices)
voxel_grid = o3d.geometry.VoxelGrid.create_from_triangle_mesh_within_bounds(mesh, voxel_size=1/64, min_bound=(-0.5, -0.5, -0.5), max_bound=(0.5, 0.5, 0.5))
vertices = np.array([voxel.grid_index for voxel in voxel_grid.get_voxels()])
assert np.all(vertices >= 0) and np.all(vertices < 64), "Some vertices are out of bounds"
vertices = (vertices + 0.5) / 64 - 0.5
utils3d.io.write_ply(os.path.join(dir_path, f'voxel.ply'), vertices)
# return {'sha256': sha256, 'voxelized': True, 'num_voxels': len(vertices)}
if __name__ == '__main__':
data_path = ''
output_dir = ''
with open(data_path, 'r') as f:
data_list = [json.loads(line.strip()) for line in f.readlines()]
args_list = [(sha256, output_dir) for sha256 in data_list]
with Pool(128) as p:
results = list(tqdm(p.imap(_voxelize, args_list), total=len(args_list)))
@@ -0,0 +1,44 @@
result_name: partfield_features/correspondence_demo
continue_ckpt: model/model.ckpt
triplane_channels_low: 128
triplane_channels_high: 512
triplane_resolution: 128
vertex_feature: True
n_point_per_face: 1000
n_sample_each: 10000
is_pc: True
remesh_demo: False
correspondence_demo: True
preprocess_mesh: True
dataset:
type: "Mix"
data_path: data/DenseCorr3D
train_batch_size: 1
val_batch_size: 1
train_num_workers: 8
all_files:
# pairs of example to run correspondence
- animals/071b8_toy_animals_017/simple_mesh.obj
- animals/bdfd0_toy_animals_016/simple_mesh.obj
- animals/2d6b3_toy_animals_009/simple_mesh.obj
- animals/96615_toy_animals_018/simple_mesh.obj
- chairs/063d1_chair_006/simple_mesh.obj
- chairs/bea57_chair_012/simple_mesh.obj
- chairs/fe0fe_chair_004/simple_mesh.obj
- chairs/288dc_chair_011/simple_mesh.obj
# consider decimating animals/../color_mesh.obj yourself for better mesh topology than the provided simple_mesh.obj
# (e.g. <50k vertices for functional map efficiency).
loss:
triplet: 1.0
use_2d_feat: False
pvcnn:
point_encoder_type: 'pvcnn'
z_triplane_channels: 256
z_triplane_resolution: 128
@@ -0,0 +1,28 @@
result_name: demo_test
continue_ckpt: model/model.ckpt
triplane_channels_low: 128
triplane_channels_high: 512
triplane_resolution: 128
n_point_per_face: 1000
n_sample_each: 10000
is_pc : True
remesh_demo : False
dataset:
type: "Mix"
data_path: "objaverse_data"
train_batch_size: 1
val_batch_size: 1
train_num_workers: 8
loss:
triplet: 1.0
use_2d_feat: False
pvcnn:
point_encoder_type: 'pvcnn'
z_triplane_channels: 256
z_triplane_resolution: 128
@@ -0,0 +1,26 @@
import argparse
import os.path as osp
from datetime import datetime
import pytz
def default_argument_parser(add_help=True, default_config_file=""):
parser = argparse.ArgumentParser(add_help=add_help)
parser.add_argument("--config-file", '-c', default=default_config_file, metavar="FILE", help="path to config file")
parser.add_argument(
"--opts",
help="Modify config options using the command-line",
default=None,
nargs=argparse.REMAINDER,
)
return parser
def setup(args, freeze=True):
from .defaults import _C as cfg
cfg = cfg.clone()
cfg.merge_from_file(args.config_file)
cfg.merge_from_list(args.opts)
dt = datetime.now(pytz.timezone('America/Los_Angeles')).strftime('%y%m%d-%H%M%S')
cfg.output_dir = osp.join(cfg.output_dir, cfg.name, dt)
if freeze:
cfg.freeze()
return cfg
@@ -0,0 +1,92 @@
from yacs.config import CfgNode as CN
_C = CN()
_C.seed = 0
_C.output_dir = "results"
_C.result_name = "test_all"
_C.triplet_sampling = "random"
_C.load_original_mesh = False
_C.num_pos = 64
_C.num_neg_random = 256
_C.num_neg_hard_pc = 128
_C.num_neg_hard_emb = 128
_C.vertex_feature = False # if true, sample feature on vertices; if false, sample feature on faces
_C.n_point_per_face = 2000
_C.n_sample_each = 10000
_C.preprocess_mesh = False
_C.regress_2d_feat = False
_C.is_pc = False
_C.cut_manifold = False
_C.remesh_demo = False
_C.correspondence_demo = False
_C.save_every_epoch = 10
_C.training_epochs = 30
_C.continue_training = False
_C.continue_ckpt = None
_C.epoch_selected = "epoch=50.ckpt"
_C.triplane_resolution = 128
_C.triplane_channels_low = 128
_C.triplane_channels_high = 512
_C.lr = 1e-3
_C.train = True
_C.test = False
_C.inference_save_pred_sdf_to_mesh=True
_C.inference_save_feat_pca=True
_C.name = "test"
_C.test_subset = False
_C.test_corres = False
_C.test_partobjaversetiny = False
_C.dataset = CN()
_C.dataset.type = "Demo_Dataset"
_C.dataset.data_path = "objaverse_data/"
_C.dataset.train_num_workers = 64
_C.dataset.val_num_workers = 32
_C.dataset.train_batch_size = 2
_C.dataset.val_batch_size = 2
_C.dataset.all_files = [] # only used for correspondence demo
_C.voxel2triplane = CN()
_C.voxel2triplane.transformer_dim = 1024
_C.voxel2triplane.transformer_layers = 6
_C.voxel2triplane.transformer_heads = 8
_C.voxel2triplane.triplane_low_res = 32
_C.voxel2triplane.triplane_high_res = 256
_C.voxel2triplane.triplane_dim = 64
_C.voxel2triplane.normalize_vox_feat = False
_C.loss = CN()
_C.loss.triplet = 0.0
_C.loss.sdf = 1.0
_C.loss.feat = 10.0
_C.loss.l1 = 0.0
_C.use_pvcnn = False
_C.use_pvcnnonly = True
_C.pvcnn = CN()
_C.pvcnn.point_encoder_type = 'pvcnn'
_C.pvcnn.use_point_scatter = True
_C.pvcnn.z_triplane_channels = 64
_C.pvcnn.z_triplane_resolution = 256
_C.pvcnn.unet_cfg = CN()
_C.pvcnn.unet_cfg.depth = 3
_C.pvcnn.unet_cfg.enabled = True
_C.pvcnn.unet_cfg.rolled = True
_C.pvcnn.unet_cfg.use_3d_aware = True
_C.pvcnn.unet_cfg.start_hidden_channels = 32
_C.pvcnn.unet_cfg.use_initial_conv = False
_C.use_2d_feat = False
_C.inference_metrics_only = False
@@ -0,0 +1,251 @@
"""
Taken from gensdf
https://github.com/princeton-computational-imaging/gensdf
"""
import torch
import torch.nn as nn
import torch.nn.functional as F
# from dnnlib.util import printarr
try:
from torch_scatter import scatter_mean, scatter_max
except:
pass
# from .unet import UNet
import torch
import torch.nn as nn
import torch.nn.functional as F
# Resnet Blocks
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.ReLU()
if size_in == size_out:
self.shortcut = None
else:
self.shortcut = nn.Linear(size_in, size_out, bias=False)
# Initialization
nn.init.zeros_(self.fc_1.weight)
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 ConvPointnet(nn.Module):
''' PointNet-based encoder network with ResNet blocks for each point.
Number of input points are fixed.
Args:
c_dim (int): dimension of latent code c
dim (int): input points dimension
hidden_dim (int): hidden dimension of the network
scatter_type (str): feature aggregation when doing local pooling
unet (bool): weather to use U-Net
unet_kwargs (str): U-Net parameters
plane_resolution (int): defined resolution for plane feature
plane_type (str): feature type, 'xz' - 1-plane, ['xz', 'xy', 'yz'] - 3-plane, ['grid'] - 3D grid volume
padding (float): conventional padding paramter of ONet for unit cube, so [-0.5, 0.5] -> [-0.55, 0.55]
n_blocks (int): number of blocks ResNetBlockFC layers
'''
def __init__(self, c_dim=128, dim=3, hidden_dim=128, scatter_type='max',
# unet=False, unet_kwargs=None,
plane_resolution=None, plane_type=['xz', 'xy', 'yz'], padding=0.1, n_blocks=5):
super().__init__()
self.c_dim = c_dim
self.fc_pos = nn.Linear(dim, 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, c_dim)
self.actvn = nn.ReLU()
self.hidden_dim = hidden_dim
# if unet:
# self.unet = UNet(c_dim, in_channels=c_dim, **unet_kwargs)
# else:
# self.unet = None
self.reso_plane = plane_resolution
self.plane_type = plane_type
self.padding = padding
if scatter_type == 'max':
self.scatter = scatter_max
elif scatter_type == 'mean':
self.scatter = scatter_mean
# takes in "p": point cloud and "query": sdf_xyz
# sample plane features for unlabeled_query as well
def forward(self, p):#, query2):
batch_size, T, D = p.size()
# acquire the index for each point
coord = {}
index = {}
if 'xz' in self.plane_type:
coord['xz'] = self.normalize_coordinate(p.clone(), plane='xz', padding=self.padding)
index['xz'] = self.coordinate2index(coord['xz'], self.reso_plane)
if 'xy' in self.plane_type:
coord['xy'] = self.normalize_coordinate(p.clone(), plane='xy', padding=self.padding)
index['xy'] = self.coordinate2index(coord['xy'], self.reso_plane)
if 'yz' in self.plane_type:
coord['yz'] = self.normalize_coordinate(p.clone(), plane='yz', padding=self.padding)
index['yz'] = self.coordinate2index(coord['yz'], self.reso_plane)
net = self.fc_pos(p)
net = self.blocks[0](net)
for block in self.blocks[1:]:
pooled = self.pool_local(coord, index, net)
net = torch.cat([net, pooled], dim=2)
net = block(net)
c = self.fc_c(net)
fea = {}
plane_feat_sum = 0
#second_sum = 0
if 'xz' in self.plane_type:
fea['xz'] = self.generate_plane_features(p, c, plane='xz') # shape: batch, latent size, resolution, resolution (e.g. 16, 256, 64, 64)
# plane_feat_sum += self.sample_plane_feature(query, fea['xz'], 'xz')
#second_sum += self.sample_plane_feature(query2, fea['xz'], 'xz')
if 'xy' in self.plane_type:
fea['xy'] = self.generate_plane_features(p, c, plane='xy')
# plane_feat_sum += self.sample_plane_feature(query, fea['xy'], 'xy')
#second_sum += self.sample_plane_feature(query2, fea['xy'], 'xy')
if 'yz' in self.plane_type:
fea['yz'] = self.generate_plane_features(p, c, plane='yz')
# plane_feat_sum += self.sample_plane_feature(query, fea['yz'], 'yz')
#second_sum += self.sample_plane_feature(query2, fea['yz'], 'yz')
return fea
# return plane_feat_sum.transpose(2,1)#, second_sum.transpose(2,1)
def normalize_coordinate(self, p, padding=0.1, plane='xz'):
''' Normalize coordinate to [0, 1] for unit cube experiments
Args:
p (tensor): point
padding (float): conventional padding paramter of ONet for unit cube, so [-0.5, 0.5] -> [-0.55, 0.55]
plane (str): plane feature type, ['xz', 'xy', 'yz']
'''
if plane == 'xz':
xy = p[:, :, [0, 2]]
elif plane =='xy':
xy = p[:, :, [0, 1]]
else:
xy = p[:, :, [1, 2]]
xy_new = xy / (1 + padding + 10e-6) # (-0.5, 0.5)
xy_new = xy_new + 0.5 # range (0, 1)
# f there are outliers out of the range
if xy_new.max() >= 1:
xy_new[xy_new >= 1] = 1 - 10e-6
if xy_new.min() < 0:
xy_new[xy_new < 0] = 0.0
return xy_new
def coordinate2index(self, x, reso):
''' Normalize coordinate to [0, 1] for unit cube experiments.
Corresponds to our 3D model
Args:
x (tensor): coordinate
reso (int): defined resolution
coord_type (str): coordinate type
'''
x = (x * reso).long()
index = x[:, :, 0] + reso * x[:, :, 1]
index = index[:, None, :]
return index
# xy is the normalized coordinates of the point cloud of each plane
# I'm pretty sure the keys of xy are the same as those of index, so xy isn't needed here as input
def pool_local(self, xy, index, c):
bs, fea_dim = c.size(0), c.size(2)
keys = xy.keys()
c_out = 0
for key in keys:
# scatter plane features from points
fea = self.scatter(c.permute(0, 2, 1), index[key], dim_size=self.reso_plane**2)
if self.scatter == scatter_max:
fea = fea[0]
# gather feature back to points
fea = fea.gather(dim=2, index=index[key].expand(-1, fea_dim, -1))
c_out += fea
return c_out.permute(0, 2, 1)
def generate_plane_features(self, p, c, plane='xz'):
# acquire indices of features in plane
xy = self.normalize_coordinate(p.clone(), plane=plane, padding=self.padding) # normalize to the range of (0, 1)
index = self.coordinate2index(xy, self.reso_plane)
# scatter plane features from points
fea_plane = c.new_zeros(p.size(0), self.c_dim, self.reso_plane**2)
c = c.permute(0, 2, 1) # B x 512 x T
fea_plane = scatter_mean(c, index, out=fea_plane) # B x 512 x reso^2
fea_plane = fea_plane.reshape(p.size(0), self.c_dim, self.reso_plane, self.reso_plane) # sparce matrix (B x 512 x reso x reso)
# printarr(fea_plane, c, p, xy, index)
# import pdb; pdb.set_trace()
# process the plane features with UNet
# if self.unet is not None:
# fea_plane = self.unet(fea_plane)
return fea_plane
# sample_plane_feature function copied from /src/conv_onet/models/decoder.py
# uses values from plane_feature and pixel locations from vgrid to interpolate feature
def sample_plane_feature(self, query, plane_feature, plane):
xy = self.normalize_coordinate(query.clone(), plane=plane, padding=self.padding)
xy = xy[:, :, None].float()
vgrid = 2.0 * xy - 1.0 # normalize to (-1, 1)
sampled_feat = F.grid_sample(plane_feature, vgrid, padding_mode='border', align_corners=True, mode='bilinear').squeeze(-1)
return sampled_feat
File diff suppressed because it is too large Load Diff
@@ -0,0 +1,243 @@
# Copyright (c) 2022, NVIDIA CORPORATION & AFFILIATES. All rights reserved.
#
# NVIDIA CORPORATION & AFFILIATES and its licensors retain all intellectual property
# and proprietary rights in and to this software, related documentation
# and any modifications thereto. Any use, reproduction, disclosure or
# distribution of this software and related documentation without an express
# license agreement from NVIDIA CORPORATION & AFFILIATES is strictly prohibited.
from ast import Dict
import math
import numpy as np
import torch
from torch import nn
import torch.nn.functional as F
from torch_scatter import scatter_mean #, scatter_max
from .unet_3daware import setup_unet #UNetTriplane3dAware
from .conv_pointnet import ConvPointnet
from .pc_encoder import PVCNNEncoder #PointNet
import einops
from .dnnlib_util import ScopedTorchProfiler, printarr
def generate_plane_features(p, c, resolution, plane='xz'):
"""
Args:
p: (B,3,n_p)
c: (B,C,n_p)
"""
padding = 0.
c_dim = c.size(1)
# acquire indices of features in plane
xy = normalize_coordinate(p.clone(), plane=plane, padding=padding) # normalize to the range of (0, 1)
index = coordinate2index(xy, resolution)
# scatter plane features from points
fea_plane = c.new_zeros(p.size(0), c_dim, resolution**2)
fea_plane = scatter_mean(c, index, out=fea_plane) # B x 512 x reso^2
fea_plane = fea_plane.reshape(p.size(0), c_dim, resolution, resolution) # sparce matrix (B x 512 x reso x reso)
return fea_plane
def normalize_coordinate(p, padding=0.1, plane='xz'):
''' Normalize coordinate to [0, 1] for unit cube experiments
Args:
p (tensor): point
padding (float): conventional padding paramter of ONet for unit cube, so [-0.5, 0.5] -> [-0.55, 0.55]
plane (str): plane feature type, ['xz', 'xy', 'yz']
'''
if plane == 'xz':
xy = p[:, :, [0, 2]]
elif plane =='xy':
xy = p[:, :, [0, 1]]
else:
xy = p[:, :, [1, 2]]
xy_new = xy / (1 + padding + 10e-6) # (-0.5, 0.5)
xy_new = xy_new + 0.5 # range (0, 1)
# if there are outliers out of the range
if xy_new.max() >= 1:
xy_new[xy_new >= 1] = 1 - 10e-6
if xy_new.min() < 0:
xy_new[xy_new < 0] = 0.0
return xy_new
def coordinate2index(x, resolution):
''' Normalize coordinate to [0, 1] for unit cube experiments.
Corresponds to our 3D model
Args:
x (tensor): coordinate
reso (int): defined resolution
coord_type (str): coordinate type
'''
x = (x * resolution).long()
index = x[:, :, 0] + resolution * x[:, :, 1]
index = index[:, None, :]
return index
def softclip(x, min, max, hardness=5):
# Soft clipping for the logsigma
x = min + F.softplus(hardness*(x - min))/hardness
x = max - F.softplus(-hardness*(x - max))/hardness
return x
def sample_triplane_feat(feature_triplane, normalized_pos):
'''
normalized_pos [-1, 1]
'''
tri_plane = torch.unbind(feature_triplane, dim=1)
x_feat = F.grid_sample(
tri_plane[0],
torch.cat(
[normalized_pos[:, :, 0:1], normalized_pos[:, :, 1:2]],
dim=-1).unsqueeze(dim=1), padding_mode='border',
align_corners=True)
y_feat = F.grid_sample(
tri_plane[1],
torch.cat(
[normalized_pos[:, :, 1:2], normalized_pos[:, :, 2:3]],
dim=-1).unsqueeze(dim=1), padding_mode='border',
align_corners=True)
z_feat = F.grid_sample(
tri_plane[2],
torch.cat(
[normalized_pos[:, :, 0:1], normalized_pos[:, :, 2:3]],
dim=-1).unsqueeze(dim=1), padding_mode='border',
align_corners=True)
final_feat = (x_feat + y_feat + z_feat)
final_feat = final_feat.squeeze(dim=2).permute(0, 2, 1) # 32dimension
return final_feat
# @persistence.persistent_class
class TriPlanePC2Encoder(torch.nn.Module):
# Encoder that encode point cloud to triplane feature vector similar to ConvOccNet
def __init__(
self,
cfg,
device='cuda',
shape_min=-1.0,
shape_length=2.0,
use_2d_feat=False,
# point_encoder='pvcnn',
# use_point_scatter=False
):
"""
Outputs latent triplane from PC input
Configs:
max_logsigma: (float) Soft clip upper range for logsigm
min_logsigma: (float)
point_encoder_type: (str) one of ['pvcnn', 'pointnet']
pvcnn_flatten_voxels: (bool) for pvcnn whether to reduce voxel
features (instead of scattering point features)
unet_cfg: (dict)
z_triplane_channels: (int) output latent triplane
z_triplane_resolution: (int)
Args:
"""
# assert img_resolution >= 4 and img_resolution & (img_resolution - 1) == 0
super().__init__()
self.device = device
self.cfg = cfg
self.shape_min = shape_min
self.shape_length = shape_length
self.z_triplane_resolution = cfg.z_triplane_resolution
z_triplane_channels = cfg.z_triplane_channels
point_encoder_out_dim = z_triplane_channels #* 2
in_channels = 6
# self.resample_filter=[1, 3, 3, 1]
if cfg.point_encoder_type == 'pvcnn':
self.pc_encoder = PVCNNEncoder(point_encoder_out_dim,
device=self.device, in_channels=in_channels, use_2d_feat=use_2d_feat) # Encode it to a volume vector.
elif cfg.point_encoder_type == 'pointnet':
# TODO the pointnet was buggy, investigate
self.pc_encoder = ConvPointnet(c_dim=point_encoder_out_dim,
dim=in_channels, hidden_dim=32,
plane_resolution=self.z_triplane_resolution,
padding=0)
else:
raise NotImplementedError(f"Point encoder {cfg.point_encoder_type} not implemented")
if cfg.unet_cfg.enabled:
self.unet_encoder = setup_unet(
output_channels=point_encoder_out_dim,
input_channels=point_encoder_out_dim,
unet_cfg=cfg.unet_cfg)
else:
self.unet_encoder = None
# @ScopedTorchProfiler('encode')
def encode(self, point_cloud_xyz, point_cloud_feature, mv_feat=None, pc2pc_idx=None) -> Dict:
# output = AttrDict()
point_cloud_xyz = (point_cloud_xyz - self.shape_min) / self.shape_length # [0, 1]
point_cloud_xyz = point_cloud_xyz - 0.5 # [-0.5, 0.5]
point_cloud = torch.cat([point_cloud_xyz, point_cloud_feature], dim=-1)
if self.cfg.point_encoder_type == 'pvcnn':
if mv_feat is not None:
pc_feat, points_feat = self.pc_encoder(point_cloud, mv_feat, pc2pc_idx)
else:
pc_feat, points_feat = self.pc_encoder(point_cloud) # 3D feature volume: BxDx32x32x32
if self.cfg.use_point_scatter:
# Scattering from PVCNN point features
points_feat_ = points_feat[0]
# shape: batch, latent size, resolution, resolution (e.g. 16, 256, 64, 64)
pc_feat_1 = generate_plane_features(point_cloud_xyz, points_feat_,
resolution=self.z_triplane_resolution, plane='xy')
pc_feat_2 = generate_plane_features(point_cloud_xyz, points_feat_,
resolution=self.z_triplane_resolution, plane='yz')
pc_feat_3 = generate_plane_features(point_cloud_xyz, points_feat_,
resolution=self.z_triplane_resolution, plane='xz')
pc_feat = pc_feat[0]
else:
pc_feat = pc_feat[0]
sf = self.z_triplane_resolution//32 # 32 is PVCNN's voxel dim
pc_feat_1 = torch.mean(pc_feat, dim=-1) #xy_plane, normalize in z plane
pc_feat_2 = torch.mean(pc_feat, dim=-3) #yz_plane, normalize in x plane
pc_feat_3 = torch.mean(pc_feat, dim=-2) #xz_plane, normalize in y plane
# nearest upsample
pc_feat_1 = einops.repeat(pc_feat_1, 'b c h w -> b c (h hm ) (w wm)', hm = sf, wm = sf)
pc_feat_2 = einops.repeat(pc_feat_2, 'b c h w -> b c (h hm) (w wm)', hm = sf, wm = sf)
pc_feat_3 = einops.repeat(pc_feat_3, 'b c h w -> b c (h hm) (w wm)', hm = sf, wm = sf)
elif self.cfg.point_encoder_type == 'pointnet':
assert self.cfg.use_point_scatter
# Run ConvPointnet
pc_feat = self.pc_encoder(point_cloud)
pc_feat_1 = pc_feat['xy'] #
pc_feat_2 = pc_feat['yz']
pc_feat_3 = pc_feat['xz']
else:
raise NotImplementedError()
if self.unet_encoder is not None:
# TODO eval adding a skip connection
# Unet expects B, 3, C, H, W
pc_feat_tri_plane_stack_pre = torch.stack([pc_feat_1, pc_feat_2, pc_feat_3], dim=1)
# dpc_feat_tri_plane_stack = self.unet_encoder(pc_feat_tri_plane_stack_pre)
# pc_feat_tri_plane_stack = pc_feat_tri_plane_stack_pre + dpc_feat_tri_plane_stack
pc_feat_tri_plane_stack = self.unet_encoder(pc_feat_tri_plane_stack_pre)
pc_feat_1, pc_feat_2, pc_feat_3 = torch.unbind(pc_feat_tri_plane_stack, dim=1)
return torch.stack([pc_feat_1, pc_feat_2, pc_feat_3], dim=1)
def forward(self, point_cloud_xyz, point_cloud_feature=None, mv_feat=None, pc2pc_idx=None):
return self.encode(point_cloud_xyz, point_cloud_feature=point_cloud_feature, mv_feat=mv_feat, pc2pc_idx=pc2pc_idx)
@@ -0,0 +1,90 @@
import torch
import torch.nn as nn
import torch.nn.functional as F
import functools
from .pv_module import SharedMLP, PVConv
def create_pointnet_components(
blocks, in_channels, with_se=False, normalize=True, eps=0,
width_multiplier=1, voxel_resolution_multiplier=1, scale_pvcnn=False, device='cuda'):
r, vr = width_multiplier, voxel_resolution_multiplier
layers, concat_channels = [], 0
for out_channels, num_blocks, voxel_resolution in blocks:
out_channels = int(r * out_channels)
if voxel_resolution is None:
block = functools.partial(SharedMLP, device=device)
else:
block = functools.partial(
PVConv, kernel_size=3, resolution=int(vr * voxel_resolution),
with_se=with_se, normalize=normalize, eps=eps, scale_pvcnn=scale_pvcnn, device=device)
for _ in range(num_blocks):
layers.append(block(in_channels, out_channels))
in_channels = out_channels
concat_channels += out_channels
return layers, in_channels, concat_channels
class PCMerger(nn.Module):
# merge surface sampled PC and rendering backprojected PC (w/ 2D features):
def __init__(self, in_channels=204, device="cuda"):
super(PCMerger, self).__init__()
self.mlp_normal = SharedMLP(3, [128, 128], device=device)
self.mlp_rgb = SharedMLP(3, [128, 128], device=device)
self.mlp_sam = SharedMLP(204 - 6, [128, 128], device=device)
def forward(self, feat, mv_feat, pc2pc_idx):
mv_feat_normal = self.mlp_normal(mv_feat[:, :3, :])
mv_feat_rgb = self.mlp_rgb(mv_feat[:, 3:6, :])
mv_feat_sam = self.mlp_sam(mv_feat[:, 6:, :])
mv_feat_normal = mv_feat_normal.permute(0, 2, 1)
mv_feat_rgb = mv_feat_rgb.permute(0, 2, 1)
mv_feat_sam = mv_feat_sam.permute(0, 2, 1)
feat = feat.permute(0, 2, 1)
for i in range(mv_feat.shape[0]):
mask = (pc2pc_idx[i] != -1).reshape(-1)
idx = pc2pc_idx[i][mask].reshape(-1)
feat[i][mask] += mv_feat_normal[i][idx] + mv_feat_rgb[i][idx] + mv_feat_sam[i][idx]
return feat.permute(0, 2, 1)
class PVCNNEncoder(nn.Module):
def __init__(self, pvcnn_feat_dim, device='cuda', in_channels=3, use_2d_feat=False):
super(PVCNNEncoder, self).__init__()
self.device = device
self.blocks = ((pvcnn_feat_dim, 1, 32), (128, 2, 16), (256, 1, 8))
self.use_2d_feat=use_2d_feat
if in_channels == 6:
self.append_channel = 2
elif in_channels == 3:
self.append_channel = 1
else:
raise NotImplementedError
layers, channels_point, concat_channels_point = create_pointnet_components(
blocks=self.blocks, in_channels=in_channels + self.append_channel, with_se=False, normalize=False,
width_multiplier=1, voxel_resolution_multiplier=1, scale_pvcnn=True,
device=device
)
self.encoder = nn.ModuleList(layers)#.to(self.device)
if self.use_2d_feat:
self.merger = PCMerger()
def forward(self, input_pc, mv_feat=None, pc2pc_idx=None):
features = input_pc.permute(0, 2, 1) * 2 # make point cloud [-1, 1]
coords = features[:, :3, :]
out_features_list = []
voxel_feature_list = []
zero_padding = torch.zeros(features.shape[0], self.append_channel, features.shape[-1], device=features.device, dtype=features.dtype)
features = torch.cat([features, zero_padding], dim=1)##################
for i in range(len(self.encoder)):
features, _, voxel_feature = self.encoder[i]((features, coords))
if i == 0 and mv_feat is not None:
features = self.merger(features, mv_feat.permute(0, 2, 1), pc2pc_idx)
out_features_list.append(features)
voxel_feature_list.append(voxel_feature)
return voxel_feature_list, out_features_list
@@ -0,0 +1,2 @@
from .pvconv import PVConv
from .shared_mlp import SharedMLP
@@ -0,0 +1,34 @@
import torch
import torch.nn as nn
from . import functional as F
__all__ = ['BallQuery']
class BallQuery(nn.Module):
def __init__(self, radius, num_neighbors, include_coordinates=True):
super().__init__()
self.radius = radius
self.num_neighbors = num_neighbors
self.include_coordinates = include_coordinates
def forward(self, points_coords, centers_coords, points_features=None):
points_coords = points_coords.contiguous()
centers_coords = centers_coords.contiguous()
neighbor_indices = F.ball_query(centers_coords, points_coords, self.radius, self.num_neighbors)
neighbor_coordinates = F.grouping(points_coords, neighbor_indices)
neighbor_coordinates = neighbor_coordinates - centers_coords.unsqueeze(-1)
if points_features is None:
assert self.include_coordinates, 'No Features For Grouping'
neighbor_features = neighbor_coordinates
else:
neighbor_features = F.grouping(points_features, neighbor_indices)
if self.include_coordinates:
neighbor_features = torch.cat([neighbor_coordinates, neighbor_features], dim=1)
return neighbor_features
def extra_repr(self):
return 'radius={}, num_neighbors={}{}'.format(
self.radius, self.num_neighbors, ', include coordinates' if self.include_coordinates else '')
@@ -0,0 +1,141 @@
import numpy as np
import torch
import torch.nn as nn
import torch.nn.functional as F
from . import functional as PF
__all__ = ['FrustumPointNetLoss', 'get_box_corners_3d']
class FrustumPointNetLoss(nn.Module):
def __init__(
self, num_heading_angle_bins, num_size_templates, size_templates, box_loss_weight=1.0,
corners_loss_weight=10.0, heading_residual_loss_weight=20.0, size_residual_loss_weight=20.0):
super().__init__()
self.box_loss_weight = box_loss_weight
self.corners_loss_weight = corners_loss_weight
self.heading_residual_loss_weight = heading_residual_loss_weight
self.size_residual_loss_weight = size_residual_loss_weight
self.num_heading_angle_bins = num_heading_angle_bins
self.num_size_templates = num_size_templates
self.register_buffer('size_templates', size_templates.view(self.num_size_templates, 3))
self.register_buffer(
'heading_angle_bin_centers', torch.arange(0, 2 * np.pi, 2 * np.pi / self.num_heading_angle_bins)
)
def forward(self, inputs, targets):
mask_logits = inputs['mask_logits'] # (B, 2, N)
center_reg = inputs['center_reg'] # (B, 3)
center = inputs['center'] # (B, 3)
heading_scores = inputs['heading_scores'] # (B, NH)
heading_residuals_normalized = inputs['heading_residuals_normalized'] # (B, NH)
heading_residuals = inputs['heading_residuals'] # (B, NH)
size_scores = inputs['size_scores'] # (B, NS)
size_residuals_normalized = inputs['size_residuals_normalized'] # (B, NS, 3)
size_residuals = inputs['size_residuals'] # (B, NS, 3)
mask_logits_target = targets['mask_logits'] # (B, N)
center_target = targets['center'] # (B, 3)
heading_bin_id_target = targets['heading_bin_id'] # (B, )
heading_residual_target = targets['heading_residual'] # (B, )
size_template_id_target = targets['size_template_id'] # (B, )
size_residual_target = targets['size_residual'] # (B, 3)
batch_size = center.size(0)
batch_id = torch.arange(batch_size, device=center.device)
# Basic Classification and Regression losses
mask_loss = F.cross_entropy(mask_logits, mask_logits_target)
heading_loss = F.cross_entropy(heading_scores, heading_bin_id_target)
size_loss = F.cross_entropy(size_scores, size_template_id_target)
center_loss = PF.huber_loss(torch.norm(center_target - center, dim=-1), delta=2.0)
center_reg_loss = PF.huber_loss(torch.norm(center_target - center_reg, dim=-1), delta=1.0)
# Refinement losses for size/heading
heading_residuals_normalized = heading_residuals_normalized[batch_id, heading_bin_id_target] # (B, )
heading_residual_normalized_target = heading_residual_target / (np.pi / self.num_heading_angle_bins)
heading_residual_normalized_loss = PF.huber_loss(
heading_residuals_normalized - heading_residual_normalized_target, delta=1.0
)
size_residuals_normalized = size_residuals_normalized[batch_id, size_template_id_target] # (B, 3)
size_residual_normalized_target = size_residual_target / self.size_templates[size_template_id_target]
size_residual_normalized_loss = PF.huber_loss(
torch.norm(size_residual_normalized_target - size_residuals_normalized, dim=-1), delta=1.0
)
# Bounding box losses
heading = (heading_residuals[batch_id, heading_bin_id_target]
+ self.heading_angle_bin_centers[heading_bin_id_target]) # (B, )
# Warning: in origin code, size_residuals are added twice (issue #43 and #49 in charlesq34/frustum-pointnets)
size = (size_residuals[batch_id, size_template_id_target]
+ self.size_templates[size_template_id_target]) # (B, 3)
corners = get_box_corners_3d(centers=center, headings=heading, sizes=size, with_flip=False) # (B, 3, 8)
heading_target = self.heading_angle_bin_centers[heading_bin_id_target] + heading_residual_target # (B, )
size_target = self.size_templates[size_template_id_target] + size_residual_target # (B, 3)
corners_target, corners_target_flip = get_box_corners_3d(
centers=center_target, headings=heading_target,
sizes=size_target, with_flip=True) # (B, 3, 8)
corners_loss = PF.huber_loss(
torch.min(
torch.norm(corners - corners_target, dim=1), torch.norm(corners - corners_target_flip, dim=1)
), delta=1.0)
# Summing up
loss = mask_loss + self.box_loss_weight * (
center_loss + center_reg_loss + heading_loss + size_loss
+ self.heading_residual_loss_weight * heading_residual_normalized_loss
+ self.size_residual_loss_weight * size_residual_normalized_loss
+ self.corners_loss_weight * corners_loss
)
return loss
def get_box_corners_3d(centers, headings, sizes, with_flip=False):
"""
:param centers: coords of box centers, FloatTensor[N, 3]
:param headings: heading angles, FloatTensor[N, ]
:param sizes: box sizes, FloatTensor[N, 3]
:param with_flip: bool, whether to return flipped box (headings + np.pi)
:return:
coords of box corners, FloatTensor[N, 3, 8]
NOTE: corner points are in counter clockwise order, e.g.,
2--1
3--0 5
7--4
"""
l = sizes[:, 0] # (N,)
w = sizes[:, 1] # (N,)
h = sizes[:, 2] # (N,)
x_corners = torch.stack([l / 2, l / 2, -l / 2, -l / 2, l / 2, l / 2, -l / 2, -l / 2], dim=1) # (N, 8)
y_corners = torch.stack([h / 2, h / 2, h / 2, h / 2, -h / 2, -h / 2, -h / 2, -h / 2], dim=1) # (N, 8)
z_corners = torch.stack([w / 2, -w / 2, -w / 2, w / 2, w / 2, -w / 2, -w / 2, w / 2], dim=1) # (N, 8)
c = torch.cos(headings) # (N,)
s = torch.sin(headings) # (N,)
o = torch.ones_like(headings) # (N,)
z = torch.zeros_like(headings) # (N,)
centers = centers.unsqueeze(-1) # (B, 3, 1)
corners = torch.stack([x_corners, y_corners, z_corners], dim=1) # (N, 3, 8)
R = torch.stack([c, z, s, z, o, z, -s, z, c], dim=1).view(-1, 3, 3) # roty matrix: (N, 3, 3)
if with_flip:
R_flip = torch.stack([-c, z, -s, z, o, z, s, z, -c], dim=1).view(-1, 3, 3)
return torch.matmul(R, corners) + centers, torch.matmul(R_flip, corners) + centers
else:
return torch.matmul(R, corners) + centers
# centers = centers.unsqueeze(1) # (B, 1, 3)
# corners = torch.stack([x_corners, y_corners, z_corners], dim=-1) # (N, 8, 3)
# RT = torch.stack([c, z, -s, z, o, z, s, z, c], dim=1).view(-1, 3, 3) # (N, 3, 3)
# if with_flip:
# RT_flip = torch.stack([-c, z, s, z, o, z, -s, z, -c], dim=1).view(-1, 3, 3) # (N, 3, 3)
# return torch.matmul(corners, RT) + centers, torch.matmul(corners, RT_flip) + centers # (N, 8, 3)
# else:
# return torch.matmul(corners, RT) + centers # (N, 8, 3)
# corners = torch.stack([x_corners, y_corners, z_corners], dim=1) # (N, 3, 8)
# R = torch.stack([c, z, s, z, o, z, -s, z, c], dim=1).view(-1, 3, 3) # (N, 3, 3)
# corners = torch.matmul(R, corners) + centers.unsqueeze(2) # (N, 3, 8)
# corners = corners.transpose(1, 2) # (N, 8, 3)
@@ -0,0 +1 @@
from .devoxelization import trilinear_devoxelize
@@ -0,0 +1,12 @@
from torch.autograd import Function
import torch
import torch.nn.functional as F
__all__ = ['trilinear_devoxelize']
def trilinear_devoxelize(c, coords, r, training=None):
coords = (coords * 2 + 1.0) / r - 1.0
coords = coords.permute(0, 2, 1).reshape(c.shape[0], 1, 1, -1, 3)
f = F.grid_sample(input=c, grid=coords, padding_mode='border', align_corners=False)
f = f.squeeze(dim=2).squeeze(dim=2)
return f
@@ -0,0 +1,10 @@
import torch.nn as nn
from . import functional as F
__all__ = ['KLLoss']
class KLLoss(nn.Module):
def forward(self, x, y):
return F.kl_loss(x, y)
@@ -0,0 +1,113 @@
import torch
import torch.nn as nn
from . import functional as F
from .ball_query import BallQuery
from .shared_mlp import SharedMLP
__all__ = ['PointNetAModule', 'PointNetSAModule', 'PointNetFPModule']
class PointNetAModule(nn.Module):
def __init__(self, in_channels, out_channels, include_coordinates=True):
super().__init__()
if not isinstance(out_channels, (list, tuple)):
out_channels = [[out_channels]]
elif not isinstance(out_channels[0], (list, tuple)):
out_channels = [out_channels]
mlps = []
total_out_channels = 0
for _out_channels in out_channels:
mlps.append(
SharedMLP(
in_channels=in_channels + (3 if include_coordinates else 0),
out_channels=_out_channels, dim=1)
)
total_out_channels += _out_channels[-1]
self.include_coordinates = include_coordinates
self.out_channels = total_out_channels
self.mlps = nn.ModuleList(mlps)
def forward(self, inputs):
features, coords = inputs
if self.include_coordinates:
features = torch.cat([features, coords], dim=1)
coords = torch.zeros((coords.size(0), 3, 1), device=coords.device)
if len(self.mlps) > 1:
features_list = []
for mlp in self.mlps:
features_list.append(mlp(features).max(dim=-1, keepdim=True).values)
return torch.cat(features_list, dim=1), coords
else:
return self.mlps[0](features).max(dim=-1, keepdim=True).values, coords
def extra_repr(self):
return f'out_channels={self.out_channels}, include_coordinates={self.include_coordinates}'
class PointNetSAModule(nn.Module):
def __init__(self, num_centers, radius, num_neighbors, in_channels, out_channels, include_coordinates=True):
super().__init__()
if not isinstance(radius, (list, tuple)):
radius = [radius]
if not isinstance(num_neighbors, (list, tuple)):
num_neighbors = [num_neighbors] * len(radius)
assert len(radius) == len(num_neighbors)
if not isinstance(out_channels, (list, tuple)):
out_channels = [[out_channels]] * len(radius)
elif not isinstance(out_channels[0], (list, tuple)):
out_channels = [out_channels] * len(radius)
assert len(radius) == len(out_channels)
groupers, mlps = [], []
total_out_channels = 0
for _radius, _out_channels, _num_neighbors in zip(radius, out_channels, num_neighbors):
groupers.append(
BallQuery(radius=_radius, num_neighbors=_num_neighbors, include_coordinates=include_coordinates)
)
mlps.append(
SharedMLP(
in_channels=in_channels + (3 if include_coordinates else 0),
out_channels=_out_channels, dim=2)
)
total_out_channels += _out_channels[-1]
self.num_centers = num_centers
self.out_channels = total_out_channels
self.groupers = nn.ModuleList(groupers)
self.mlps = nn.ModuleList(mlps)
def forward(self, inputs):
features, coords = inputs
centers_coords = F.furthest_point_sample(coords, self.num_centers)
features_list = []
for grouper, mlp in zip(self.groupers, self.mlps):
features_list.append(mlp(grouper(coords, centers_coords, features)).max(dim=-1).values)
if len(features_list) > 1:
return torch.cat(features_list, dim=1), centers_coords
else:
return features_list[0], centers_coords
def extra_repr(self):
return f'num_centers={self.num_centers}, out_channels={self.out_channels}'
class PointNetFPModule(nn.Module):
def __init__(self, in_channels, out_channels):
super().__init__()
self.mlp = SharedMLP(in_channels=in_channels, out_channels=out_channels, dim=1)
def forward(self, inputs):
if len(inputs) == 3:
points_coords, centers_coords, centers_features = inputs
points_features = None
else:
points_coords, centers_coords, centers_features, points_features = inputs
interpolated_features = F.nearest_neighbor_interpolate(points_coords, centers_coords, centers_features)
if points_features is not None:
interpolated_features = torch.cat(
[interpolated_features, points_features], dim=1
)
return self.mlp(interpolated_features), points_coords
@@ -0,0 +1,38 @@
import torch.nn as nn
from . import functional as F
from .voxelization import Voxelization
from .shared_mlp import SharedMLP
import torch
__all__ = ['PVConv']
class PVConv(nn.Module):
def __init__(
self, in_channels, out_channels, kernel_size, resolution, with_se=False, normalize=True, eps=0, scale_pvcnn=False,
device='cuda'):
super().__init__()
self.in_channels = in_channels
self.out_channels = out_channels
self.kernel_size = kernel_size
self.resolution = resolution
self.voxelization = Voxelization(resolution, normalize=normalize, eps=eps, scale_pvcnn=scale_pvcnn)
voxel_layers = [
nn.Conv3d(in_channels, out_channels, kernel_size, stride=1, padding=kernel_size // 2, device=device),
nn.InstanceNorm3d(out_channels, eps=1e-4, device=device),
nn.LeakyReLU(0.1, True),
nn.Conv3d(out_channels, out_channels, kernel_size, stride=1, padding=kernel_size // 2, device=device),
nn.InstanceNorm3d(out_channels, eps=1e-4, device=device),
nn.LeakyReLU(0.1, True),
]
self.voxel_layers = nn.Sequential(*voxel_layers)
self.point_features = SharedMLP(in_channels, out_channels, device=device)
def forward(self, inputs):
features, coords = inputs
voxel_features, voxel_coords = self.voxelization(features, coords)
voxel_features = self.voxel_layers(voxel_features)
devoxel_features = F.trilinear_devoxelize(voxel_features, voxel_coords, self.resolution, self.training)
fused_features = devoxel_features + self.point_features(features)
return fused_features, coords, voxel_features
@@ -0,0 +1,35 @@
import torch.nn as nn
__all__ = ['SharedMLP']
class SharedMLP(nn.Module):
def __init__(self, in_channels, out_channels, dim=1, device='cuda'):
super().__init__()
# print('==> SharedMLP device: ', device)
if dim == 1:
conv = nn.Conv1d
bn = nn.InstanceNorm1d
elif dim == 2:
conv = nn.Conv2d
bn = nn.InstanceNorm1d
else:
raise ValueError
if not isinstance(out_channels, (list, tuple)):
out_channels = [out_channels]
layers = []
for oc in out_channels:
layers.extend(
[
conv(in_channels, oc, 1, device=device),
bn(oc, device=device),
nn.ReLU(True),
])
in_channels = oc
self.layers = nn.Sequential(*layers)
def forward(self, inputs):
if isinstance(inputs, (list, tuple)):
return (self.layers(inputs[0]), *inputs[1:])
else:
return self.layers(inputs)
@@ -0,0 +1,80 @@
import torch
import torch.nn as nn
from . import functional as F
__all__ = ['Voxelization']
def my_voxelization(features, coords, resolution):
b, c, _ = features.shape
result = torch.zeros(b, c + 1, resolution * resolution * resolution, device=features.device, dtype=features.dtype)
r = resolution
r2 = resolution * resolution
coords = coords.long()
indices = coords[:, 0] * r2 + coords[:, 1] * r + coords[:, 2]
# print(r, r2, coords[:, 0].max(), coords[:, 1].max(), coords[:, 2].max())
# print(f"Resolution: {resolution}")
# print(f"Coords shape: {coords.shape}")
# print(f"Coords max per dim: x={coords[:, 0].max()}, y={coords[:, 1].max()}, z={coords[:, 2].max()}")
# print(f"Coords min per dim: x={coords[:, 0].min()}, y={coords[:, 1].min()}, z={coords[:, 2].min()}")
# print(f"Indices shape: {indices.shape}")
# print(f"Indices max: {indices.max()}, min: {indices.min()}")
# print(f"Expected max index: {resolution * resolution * resolution - 1}")
# # 检查是否有越界的索引
# max_valid_index = resolution * resolution * resolution - 1
# invalid_mask = (indices > max_valid_index) | (indices < 0)
# if invalid_mask.any():
# print(f"Found {invalid_mask.sum()} invalid indices!")
# print(f"Invalid indices: {indices[invalid_mask]}")
# # 找到对应的坐标
# invalid_coords = coords[:, :, invalid_mask.any(dim=0)]
# print(f"Invalid coords shape: {invalid_coords.shape}")
# if invalid_coords.numel() > 0:
# print(f"Sample invalid coords: {invalid_coords[:, :, :5]}") # 显示前5个无效坐标
indices = indices.unsqueeze(dim=1).expand(-1, result.shape[1], -1)
features = torch.cat([features, torch.ones(features.shape[0], 1, features.shape[2], device=features.device, dtype=features.dtype)], dim=1)
out_feature = result.scatter_(index=indices.long(), src=features, dim=2, reduce='add')
cnt = out_feature[:, -1:, :]
zero_mask = (cnt == 0).to(features.dtype)
cnt = cnt * (1 - zero_mask) + zero_mask * 1e-5
vox_feature = out_feature[:, :-1, :] / cnt
return vox_feature.view(b, c, resolution, resolution, resolution)
class Voxelization(nn.Module):
def __init__(self, resolution, normalize=True, eps=0, scale_pvcnn=False):
super().__init__()
self.r = int(resolution)
self.normalize = normalize
self.eps = eps
self.scale_pvcnn = scale_pvcnn
assert not normalize
def forward(self, features, coords):
# import pdb; pdb.set_trace()
with torch.no_grad():
coords = coords.detach()
if self.normalize:
norm_coords = norm_coords / (norm_coords.norm(dim=1, keepdim=True).max(dim=2, keepdim=True).values * 2.0 + self.eps) + 0.5
else:
if self.scale_pvcnn:
norm_coords = (coords + 1) / 2.0 # [0, 1]
# print(norm_coords.shape, norm_coords.max(), norm_coords.min())
else:
# norm_coords = (norm_coords + 1) / 2.0
norm_coords = (coords + 1) / 2.0
norm_coords = torch.clamp(norm_coords * self.r, 0, self.r - 1)
# print(norm_coords.shape, norm_coords.max(), norm_coords.min())
vox_coords = torch.round(norm_coords)
# print(vox_coords.shape, vox_coords.max(), vox_coords.min())
# print(features.shape)
new_vox_feat = my_voxelization(features, vox_coords, self.r)
return new_vox_feat, norm_coords
def extra_repr(self):
return 'resolution={}{}'.format(self.r, ', normalized eps = {}'.format(self.eps) if self.normalize else '')
@@ -0,0 +1,427 @@
import numpy as np
import torch
import torch.nn as nn
import torch.nn.functional as F
from torch.nn import init
import einops
def conv3x3(in_channels, out_channels, stride=1,
padding=1, bias=True, groups=1):
return nn.Conv2d(
in_channels,
out_channels,
kernel_size=3,
stride=stride,
padding=padding,
bias=bias,
groups=groups)
def upconv2x2(in_channels, out_channels, mode='transpose'):
if mode == 'transpose':
return nn.ConvTranspose2d(
in_channels,
out_channels,
kernel_size=2,
stride=2)
else:
# out_channels is always going to be the same
# as in_channels
return nn.Sequential(
nn.Upsample(mode='bilinear', scale_factor=2),
conv1x1(in_channels, out_channels))
def conv1x1(in_channels, out_channels, groups=1):
return nn.Conv2d(
in_channels,
out_channels,
kernel_size=1,
groups=groups,
stride=1)
class ConvTriplane3dAware(nn.Module):
""" 3D aware triplane conv (as described in RODIN) """
def __init__(self, internal_conv_f, in_channels, out_channels, order='xz'):
"""
Args:
internal_conv_f: function that should return a 2D convolution Module
given in and out channels
order: if triplane input is in 'xz' order
"""
super(ConvTriplane3dAware, self).__init__()
# Need 3 seperate convolutions
self.in_channels = in_channels
self.out_channels = out_channels
assert order in ['xz', 'zx']
self.order = order
# Going to stack from other planes
self.plane_convs = nn.ModuleList([
internal_conv_f(3*self.in_channels, self.out_channels) for _ in range(3)])
def forward(self, triplanes_list):
"""
Args:
triplanes_list: [(B,Ci,H,W)]*3 in xy,yz,(zx or xz) depending on order
Returns:
out_triplanes_list: [(B,Co,H,W)]*3 in xy,yz,(zx or xz) depending on order
"""
inps = list(triplanes_list)
xp = 1 #(yz)
yp = 2 #(zx)
zp = 0 #(xy)
if self.order == 'xz':
# get into zx order
inps[yp] = einops.rearrange(inps[yp], 'b c x z -> b c z x')
oplanes = [None]*3
# order shouldn't matter
for iplane in [zp, xp, yp]:
# i_plane -> (j,k)
# need to average out i and convert to (j,k)
# j_plane -> (k,i)
# k_plane -> (i,j)
jplane = (iplane+1)%3
kplane = (iplane+2)%3
ifeat = inps[iplane]
# need to average out nonshared dim
# Average pool across
# j_plane -> (k,i) -> (k,1) -> (1,k) -> (j,k)
# b c k i -> b c k 1
jpool = torch.mean(inps[jplane], dim=3 ,keepdim=True)
jpool = einops.rearrange(jpool, 'b c k 1 -> b c 1 k')
jpool = einops.repeat(jpool, 'b c 1 k -> b c j k', j=ifeat.size(2))
# k_plane -> (i,j) -> (1,j) -> (j,1) -> (j,k)
# b c i j -> b c 1 j
kpool = torch.mean(inps[kplane], dim=2 ,keepdim=True)
kpool = einops.rearrange(kpool, 'b c 1 j -> b c j 1')
kpool = einops.repeat(kpool, 'b c j 1 -> b c j k', k=ifeat.size(3))
# b c h w
# jpool = jpool.expand_as(ifeat)
# kpool = kpool.expand_as(ifeat)
# concat and conv on feature dim
catfeat = torch.cat([ifeat, jpool, kpool], dim=1)
oplane = self.plane_convs[iplane](catfeat)
oplanes[iplane] = oplane
if self.order == 'xz':
# get back into xz order
oplanes[yp] = einops.rearrange(oplanes[yp], 'b c z x -> b c x z')
return oplanes
def roll_triplanes(triplanes_list):
# B, C, tri, h, w
tristack = torch.stack((triplanes_list),dim=2)
return einops.rearrange(tristack, 'b c tri h w -> b c (tri h) w', tri=3)
def unroll_triplanes(rolled_triplane):
# B, C, tri*h, w
tristack = einops.rearrange(rolled_triplane, 'b c (tri h) w -> b c tri h w', tri=3)
return torch.unbind(tristack, dim=2)
def conv1x1triplane3daware(in_channels, out_channels, order='xz', **kwargs):
return ConvTriplane3dAware(lambda inp, out: conv1x1(inp,out,**kwargs),
in_channels, out_channels,order=order)
def Normalize(in_channels, num_groups=32):
num_groups = min(in_channels, num_groups) # avoid error if in_channels < 32
return torch.nn.GroupNorm(num_groups=num_groups, num_channels=in_channels, eps=1e-6, affine=True)
def nonlinearity(x):
# return F.relu(x)
# Swish
return x*torch.sigmoid(x)
class Upsample(nn.Module):
def __init__(self, in_channels, with_conv):
super().__init__()
self.with_conv = with_conv
if self.with_conv:
self.conv = torch.nn.Conv2d(in_channels,
in_channels,
kernel_size=3,
stride=1,
padding=1)
def forward(self, x):
x = torch.nn.functional.interpolate(x, scale_factor=2.0, mode="nearest")
if self.with_conv:
x = self.conv(x)
return x
class Downsample(nn.Module):
def __init__(self, in_channels, with_conv):
super().__init__()
self.with_conv = with_conv
if self.with_conv:
# no asymmetric padding in torch conv, must do it ourselves
self.conv = torch.nn.Conv2d(in_channels,
in_channels,
kernel_size=3,
stride=2,
padding=0)
def forward(self, x):
if self.with_conv:
pad = (0,1,0,1)
x = torch.nn.functional.pad(x, pad, mode="constant", value=0)
x = self.conv(x)
else:
x = torch.nn.functional.avg_pool2d(x, kernel_size=2, stride=2)
return x
class ResnetBlock3dAware(nn.Module):
def __init__(self, in_channels, out_channels=None):
#, conv_shortcut=False):
super().__init__()
self.in_channels = in_channels
out_channels = in_channels if out_channels is None else out_channels
self.out_channels = out_channels
# self.use_conv_shortcut = conv_shortcut
self.norm1 = Normalize(in_channels)
self.conv1 = conv3x3(self.in_channels, self.out_channels)
self.norm_mid = Normalize(out_channels)
self.conv_3daware = conv1x1triplane3daware(self.out_channels, self.out_channels)
self.norm2 = Normalize(out_channels)
self.conv2 = conv3x3(self.out_channels, self.out_channels)
if self.in_channels != self.out_channels:
self.nin_shortcut = torch.nn.Conv2d(in_channels,
out_channels,
kernel_size=1,
stride=1,
padding=0)
def forward(self, x):
# 3x3 plane comm
h = x
h = self.norm1(h)
h = nonlinearity(h)
h = self.conv1(h)
# 1x1 3d aware, crossplane comm
h = self.norm_mid(h)
h = nonlinearity(h)
h = unroll_triplanes(h)
h = self.conv_3daware(h)
h = roll_triplanes(h)
# 3x3 plane comm
h = self.norm2(h)
h = nonlinearity(h)
h = self.conv2(h)
if self.in_channels != self.out_channels:
x = self.nin_shortcut(x)
return x+h
class DownConv3dAware(nn.Module):
"""
A helper Module that performs 2 convolutions and 1 MaxPool.
A ReLU activation follows each convolution.
"""
def __init__(self, in_channels, out_channels, downsample=True, with_conv=False):
super(DownConv3dAware, self).__init__()
self.in_channels = in_channels
self.out_channels = out_channels
self.block = ResnetBlock3dAware(in_channels=in_channels,
out_channels=out_channels)
self.do_downsample = downsample
self.downsample = Downsample(out_channels, with_conv=with_conv)
def forward(self, x):
"""
rolled input, rolled output
Args:
x: rolled (b c (tri*h) w)
"""
x = self.block(x)
before_pool = x
# if self.pooling:
# x = self.pool(x)
if self.do_downsample:
# unroll and cat channel-wise (to prevent pooling across triplane boundaries)
x = einops.rearrange(x, 'b c (tri h) w -> b (c tri) h w', tri=3)
x = self.downsample(x)
# undo
x = einops.rearrange(x, 'b (c tri) h w -> b c (tri h) w', tri=3)
return x, before_pool
class UpConv3dAware(nn.Module):
"""
A helper Module that performs 2 convolutions and 1 UpConvolution.
A ReLU activation follows each convolution.
"""
def __init__(self, in_channels, out_channels,
merge_mode='concat', with_conv=False): #up_mode='transpose', ):
super(UpConv3dAware, self).__init__()
self.in_channels = in_channels
self.out_channels = out_channels
self.merge_mode = merge_mode
self.upsample = Upsample(in_channels, with_conv)
if self.merge_mode == 'concat':
self.norm1 = Normalize(in_channels+out_channels)
self.block = ResnetBlock3dAware(in_channels=in_channels+out_channels,
out_channels=out_channels)
else:
self.norm1 = Normalize(in_channels)
self.block = ResnetBlock3dAware(in_channels=in_channels,
out_channels=out_channels)
def forward(self, from_down, from_up):
""" Forward pass
rolled inputs, rolled output
rolled (b c (tri*h) w)
Arguments:
from_down: tensor from the encoder pathway
from_up: upconv'd tensor from the decoder pathway
"""
# from_up = self.upconv(from_up)
from_up = self.upsample(from_up)
if self.merge_mode == 'concat':
x = torch.cat((from_up, from_down), 1)
else:
x = from_up + from_down
x = self.norm1(x)
x = self.block(x)
return x
class UNetTriplane3dAware(nn.Module):
def __init__(self, out_channels, in_channels=3, depth=5,
start_filts=64,# up_mode='transpose',
use_initial_conv=False,
merge_mode='concat', **kwargs):
"""
Arguments:
in_channels: int, number of channels in the input tensor.
Default is 3 for RGB images.
depth: int, number of MaxPools in the U-Net.
start_filts: int, number of convolutional filters for the
first conv.
"""
super(UNetTriplane3dAware, self).__init__()
self.out_channels = out_channels
self.in_channels = in_channels
self.start_filts = start_filts
self.depth = depth
self.use_initial_conv = use_initial_conv
if use_initial_conv:
self.conv_initial = conv1x1(self.in_channels, self.start_filts)
self.down_convs = []
self.up_convs = []
# create the encoder pathway and add to a list
for i in range(depth):
if i == 0:
ins = self.start_filts if use_initial_conv else self.in_channels
else:
ins = outs
outs = self.start_filts*(2**i)
downsamp_it = True if i < depth-1 else False
down_conv = DownConv3dAware(ins, outs, downsample = downsamp_it)
self.down_convs.append(down_conv)
for i in range(depth-1):
ins = outs
outs = ins // 2
up_conv = UpConv3dAware(ins, outs,
merge_mode=merge_mode)
self.up_convs.append(up_conv)
# add the list of modules to current module
self.down_convs = nn.ModuleList(self.down_convs)
self.up_convs = nn.ModuleList(self.up_convs)
self.norm_out = Normalize(outs)
self.conv_final = conv1x1(outs, self.out_channels)
self.reset_params()
@staticmethod
def weight_init(m):
if isinstance(m, nn.Conv2d):
# init.xavier_normal_(m.weight, gain=0.1)
init.xavier_normal_(m.weight)
init.constant_(m.bias, 0)
def reset_params(self):
for i, m in enumerate(self.modules()):
self.weight_init(m)
def forward(self, x):
"""
Args:
x: Stacked triplane expected to be in (B,3,C,H,W)
"""
# Roll
x = einops.rearrange(x, 'b tri c h w -> b c (tri h) w', tri=3)
if self.use_initial_conv:
x = self.conv_initial(x)
encoder_outs = []
# encoder pathway, save outputs for merging
for i, module in enumerate(self.down_convs):
x, before_pool = module(x)
encoder_outs.append(before_pool)
# Spend a block in the middle
# x = self.block_mid(x)
for i, module in enumerate(self.up_convs):
before_pool = encoder_outs[-(i+2)]
x = module(before_pool, x)
x = self.norm_out(x)
# No softmax is used. This means you need to use
# nn.CrossEntropyLoss is your training script,
# as this module includes a softmax already.
x = self.conv_final(nonlinearity(x))
# Unroll
x = einops.rearrange(x, 'b c (tri h) w -> b tri c h w', tri=3)
return x
def setup_unet(output_channels, input_channels, unet_cfg):
if unet_cfg['use_3d_aware']:
assert(unet_cfg['rolled'])
unet = UNetTriplane3dAware(
out_channels=output_channels,
in_channels=input_channels,
depth=unet_cfg['depth'],
use_initial_conv=unet_cfg['use_initial_conv'],
start_filts=unet_cfg['start_hidden_channels'],)
else:
raise NotImplementedError
return unet
@@ -0,0 +1,546 @@
#https://github.com/wolny/pytorch-3dunet/blob/master/pytorch3dunet/unet3d/buildingblocks.py
# MIT License
# Copyright (c) 2018 Adrian Wolny
# 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 functools import partial
import torch
from torch import nn as nn
from torch.nn import functional as F
# from pytorch3dunet.unet3d.se import ChannelSELayer3D, ChannelSpatialSELayer3D, SpatialSELayer3D
def create_conv(in_channels, out_channels, kernel_size, order, num_groups, padding,
dropout_prob, is3d):
"""
Create a list of modules with together constitute a single conv layer with non-linearity
and optional batchnorm/groupnorm.
Args:
in_channels (int): number of input channels
out_channels (int): number of output channels
kernel_size(int or tuple): size of the convolving kernel
order (string): order of things, e.g.
'cr' -> conv + ReLU
'gcr' -> groupnorm + conv + ReLU
'cl' -> conv + LeakyReLU
'ce' -> conv + ELU
'bcr' -> batchnorm + conv + ReLU
'cbrd' -> conv + batchnorm + ReLU + dropout
'cbrD' -> conv + batchnorm + ReLU + dropout2d
num_groups (int): number of groups for the GroupNorm
padding (int or tuple): add zero-padding added to all three sides of the input
dropout_prob (float): dropout probability
is3d (bool): is3d (bool): if True use Conv3d, otherwise use Conv2d
Return:
list of tuple (name, module)
"""
assert 'c' in order, "Conv layer MUST be present"
assert order[0] not in 'rle', 'Non-linearity cannot be the first operation in the layer'
modules = []
for i, char in enumerate(order):
if char == 'r':
modules.append(('ReLU', nn.ReLU(inplace=True)))
elif char == 'l':
modules.append(('LeakyReLU', nn.LeakyReLU(inplace=True)))
elif char == 'e':
modules.append(('ELU', nn.ELU(inplace=True)))
elif char == 'c':
# add learnable bias only in the absence of batchnorm/groupnorm
bias = not ('g' in order or 'b' in order)
if is3d:
conv = nn.Conv3d(in_channels, out_channels, kernel_size, padding=padding, bias=bias)
else:
conv = nn.Conv2d(in_channels, out_channels, kernel_size, padding=padding, bias=bias)
modules.append(('conv', conv))
elif char == 'g':
is_before_conv = i < order.index('c')
if is_before_conv:
num_channels = in_channels
else:
num_channels = out_channels
# use only one group if the given number of groups is greater than the number of channels
if num_channels < num_groups:
num_groups = 1
assert num_channels % num_groups == 0, f'Expected number of channels in input to be divisible by num_groups. num_channels={num_channels}, num_groups={num_groups}'
modules.append(('groupnorm', nn.GroupNorm(num_groups=num_groups, num_channels=num_channels)))
elif char == 'b':
is_before_conv = i < order.index('c')
if is3d:
bn = nn.BatchNorm3d
else:
bn = nn.BatchNorm2d
if is_before_conv:
modules.append(('batchnorm', bn(in_channels)))
else:
modules.append(('batchnorm', bn(out_channels)))
elif char == 'd':
modules.append(('dropout', nn.Dropout(p=dropout_prob)))
elif char == 'D':
modules.append(('dropout2d', nn.Dropout2d(p=dropout_prob)))
else:
raise ValueError(f"Unsupported layer type '{char}'. MUST be one of ['b', 'g', 'r', 'l', 'e', 'c', 'd', 'D']")
return modules
class SingleConv(nn.Sequential):
"""
Basic convolutional module consisting of a Conv3d, non-linearity and optional batchnorm/groupnorm. The order
of operations can be specified via the `order` parameter
Args:
in_channels (int): number of input channels
out_channels (int): number of output channels
kernel_size (int or tuple): size of the convolving kernel
order (string): determines the order of layers, e.g.
'cr' -> conv + ReLU
'crg' -> conv + ReLU + groupnorm
'cl' -> conv + LeakyReLU
'ce' -> conv + ELU
num_groups (int): number of groups for the GroupNorm
padding (int or tuple): add zero-padding
dropout_prob (float): dropout probability, default 0.1
is3d (bool): if True use Conv3d, otherwise use Conv2d
"""
def __init__(self, in_channels, out_channels, kernel_size=3, order='gcr', num_groups=8,
padding=1, dropout_prob=0.1, is3d=True):
super(SingleConv, self).__init__()
for name, module in create_conv(in_channels, out_channels, kernel_size, order,
num_groups, padding, dropout_prob, is3d):
self.add_module(name, module)
class DoubleConv(nn.Sequential):
"""
A module consisting of two consecutive convolution layers (e.g. BatchNorm3d+ReLU+Conv3d).
We use (Conv3d+ReLU+GroupNorm3d) by default.
This can be changed however by providing the 'order' argument, e.g. in order
to change to Conv3d+BatchNorm3d+ELU use order='cbe'.
Use padded convolutions to make sure that the output (H_out, W_out) is the same
as (H_in, W_in), so that you don't have to crop in the decoder path.
Args:
in_channels (int): number of input channels
out_channels (int): number of output channels
encoder (bool): if True we're in the encoder path, otherwise we're in the decoder
kernel_size (int or tuple): size of the convolving kernel
order (string): determines the order of layers, e.g.
'cr' -> conv + ReLU
'crg' -> conv + ReLU + groupnorm
'cl' -> conv + LeakyReLU
'ce' -> conv + ELU
num_groups (int): number of groups for the GroupNorm
padding (int or tuple): add zero-padding added to all three sides of the input
upscale (int): number of the convolution to upscale in encoder if DoubleConv, default: 2
dropout_prob (float or tuple): dropout probability for each convolution, default 0.1
is3d (bool): if True use Conv3d instead of Conv2d layers
"""
def __init__(self, in_channels, out_channels, encoder, kernel_size=3, order='gcr',
num_groups=8, padding=1, upscale=2, dropout_prob=0.1, is3d=True):
super(DoubleConv, self).__init__()
if encoder:
# we're in the encoder path
conv1_in_channels = in_channels
if upscale == 1:
conv1_out_channels = out_channels
else:
conv1_out_channels = out_channels // 2
if conv1_out_channels < in_channels:
conv1_out_channels = in_channels
conv2_in_channels, conv2_out_channels = conv1_out_channels, out_channels
else:
# we're in the decoder path, decrease the number of channels in the 1st convolution
conv1_in_channels, conv1_out_channels = in_channels, out_channels
conv2_in_channels, conv2_out_channels = out_channels, out_channels
# check if dropout_prob is a tuple and if so
# split it for different dropout probabilities for each convolution.
if isinstance(dropout_prob, list) or isinstance(dropout_prob, tuple):
dropout_prob1 = dropout_prob[0]
dropout_prob2 = dropout_prob[1]
else:
dropout_prob1 = dropout_prob2 = dropout_prob
# conv1
self.add_module('SingleConv1',
SingleConv(conv1_in_channels, conv1_out_channels, kernel_size, order, num_groups,
padding=padding, dropout_prob=dropout_prob1, is3d=is3d))
# conv2
self.add_module('SingleConv2',
SingleConv(conv2_in_channels, conv2_out_channels, kernel_size, order, num_groups,
padding=padding, dropout_prob=dropout_prob2, is3d=is3d))
class ResNetBlock(nn.Module):
"""
Residual block that can be used instead of standard DoubleConv in the Encoder module.
Motivated by: https://arxiv.org/pdf/1706.00120.pdf
Notice we use ELU instead of ReLU (order='cge') and put non-linearity after the groupnorm.
"""
def __init__(self, in_channels, out_channels, kernel_size=3, order='cge', num_groups=8, is3d=True, **kwargs):
super(ResNetBlock, self).__init__()
if in_channels != out_channels:
# conv1x1 for increasing the number of channels
if is3d:
self.conv1 = nn.Conv3d(in_channels, out_channels, 1)
else:
self.conv1 = nn.Conv2d(in_channels, out_channels, 1)
else:
self.conv1 = nn.Identity()
self.conv2 = SingleConv(in_channels, out_channels, kernel_size=kernel_size, order=order, num_groups=num_groups,
is3d=is3d)
# remove non-linearity from the 3rd convolution since it's going to be applied after adding the residual
n_order = order
for c in 'rel':
n_order = n_order.replace(c, '')
self.conv3 = SingleConv(out_channels, out_channels, kernel_size=kernel_size, order=n_order,
num_groups=num_groups, is3d=is3d)
# create non-linearity separately
if 'l' in order:
self.non_linearity = nn.LeakyReLU(negative_slope=0.1, inplace=True)
elif 'e' in order:
self.non_linearity = nn.ELU(inplace=True)
else:
self.non_linearity = nn.ReLU(inplace=True)
def forward(self, x):
# apply first convolution to bring the number of channels to out_channels
residual = self.conv1(x)
out = self.conv2(x)
out = self.conv3(out)
out += residual
out = self.non_linearity(out)
return out
class Encoder(nn.Module):
"""
A single module from the encoder path consisting of the optional max
pooling layer (one may specify the MaxPool kernel_size to be different
from the standard (2,2,2), e.g. if the volumetric data is anisotropic
(make sure to use complementary scale_factor in the decoder path) followed by
a basic module (DoubleConv or ResNetBlock).
Args:
in_channels (int): number of input channels
out_channels (int): number of output channels
conv_kernel_size (int or tuple): size of the convolving kernel
apply_pooling (bool): if True use MaxPool3d before DoubleConv
pool_kernel_size (int or tuple): the size of the window
pool_type (str): pooling layer: 'max' or 'avg'
basic_module(nn.Module): either ResNetBlock or DoubleConv
conv_layer_order (string): determines the order of layers
in `DoubleConv` module. See `DoubleConv` for more info.
num_groups (int): number of groups for the GroupNorm
padding (int or tuple): add zero-padding added to all three sides of the input
upscale (int): number of the convolution to upscale in encoder if DoubleConv, default: 2
dropout_prob (float or tuple): dropout probability, default 0.1
is3d (bool): use 3d or 2d convolutions/pooling operation
"""
def __init__(self, in_channels, out_channels, conv_kernel_size=3, apply_pooling=True,
pool_kernel_size=2, pool_type='max', basic_module=DoubleConv, conv_layer_order='gcr',
num_groups=8, padding=1, upscale=2, dropout_prob=0.1, is3d=True):
super(Encoder, self).__init__()
assert pool_type in ['max', 'avg']
if apply_pooling:
if pool_type == 'max':
if is3d:
self.pooling = nn.MaxPool3d(kernel_size=pool_kernel_size)
else:
self.pooling = nn.MaxPool2d(kernel_size=pool_kernel_size)
else:
if is3d:
self.pooling = nn.AvgPool3d(kernel_size=pool_kernel_size)
else:
self.pooling = nn.AvgPool2d(kernel_size=pool_kernel_size)
else:
self.pooling = None
self.basic_module = basic_module(in_channels, out_channels,
encoder=True,
kernel_size=conv_kernel_size,
order=conv_layer_order,
num_groups=num_groups,
padding=padding,
upscale=upscale,
dropout_prob=dropout_prob,
is3d=is3d)
def forward(self, x):
if self.pooling is not None:
x = self.pooling(x)
x = self.basic_module(x)
return x
class Decoder(nn.Module):
"""
A single module for decoder path consisting of the upsampling layer
(either learned ConvTranspose3d or nearest neighbor interpolation)
followed by a basic module (DoubleConv or ResNetBlock).
Args:
in_channels (int): number of input channels
out_channels (int): number of output channels
conv_kernel_size (int or tuple): size of the convolving kernel
scale_factor (int or tuple): used as the multiplier for the image H/W/D in
case of nn.Upsample or as stride in case of ConvTranspose3d, must reverse the MaxPool3d operation
from the corresponding encoder
basic_module(nn.Module): either ResNetBlock or DoubleConv
conv_layer_order (string): determines the order of layers
in `DoubleConv` module. See `DoubleConv` for more info.
num_groups (int): number of groups for the GroupNorm
padding (int or tuple): add zero-padding added to all three sides of the input
upsample (str): algorithm used for upsampling:
InterpolateUpsampling: 'nearest' | 'linear' | 'bilinear' | 'trilinear' | 'area'
TransposeConvUpsampling: 'deconv'
No upsampling: None
Default: 'default' (chooses automatically)
dropout_prob (float or tuple): dropout probability, default 0.1
"""
def __init__(self, in_channels, out_channels, conv_kernel_size=3, scale_factor=2, basic_module=DoubleConv,
conv_layer_order='gcr', num_groups=8, padding=1, upsample='default',
dropout_prob=0.1, is3d=True):
super(Decoder, self).__init__()
# perform concat joining per default
concat = True
# don't adapt channels after join operation
adapt_channels = False
if upsample is not None and upsample != 'none':
if upsample == 'default':
if basic_module == DoubleConv:
upsample = 'nearest' # use nearest neighbor interpolation for upsampling
concat = True # use concat joining
adapt_channels = False # don't adapt channels
elif basic_module == ResNetBlock: #or basic_module == ResNetBlockSE:
upsample = 'deconv' # use deconvolution upsampling
concat = False # use summation joining
adapt_channels = True # adapt channels after joining
# perform deconvolution upsampling if mode is deconv
if upsample == 'deconv':
self.upsampling = TransposeConvUpsampling(in_channels=in_channels, out_channels=out_channels,
kernel_size=conv_kernel_size, scale_factor=scale_factor,
is3d=is3d)
else:
self.upsampling = InterpolateUpsampling(mode=upsample)
else:
# no upsampling
self.upsampling = NoUpsampling()
# concat joining
self.joining = partial(self._joining, concat=True)
# perform joining operation
self.joining = partial(self._joining, concat=concat)
# adapt the number of in_channels for the ResNetBlock
if adapt_channels is True:
in_channels = out_channels
self.basic_module = basic_module(in_channels, out_channels,
encoder=False,
kernel_size=conv_kernel_size,
order=conv_layer_order,
num_groups=num_groups,
padding=padding,
dropout_prob=dropout_prob,
is3d=is3d)
def forward(self, encoder_features, x):
x = self.upsampling(encoder_features=encoder_features, x=x)
x = self.joining(encoder_features, x)
x = self.basic_module(x)
return x
@staticmethod
def _joining(encoder_features, x, concat):
if concat:
return torch.cat((encoder_features, x), dim=1)
else:
return encoder_features + x
def create_encoders(in_channels, f_maps, basic_module, conv_kernel_size, conv_padding,
conv_upscale, dropout_prob,
layer_order, num_groups, pool_kernel_size, is3d):
# create encoder path consisting of Encoder modules. Depth of the encoder is equal to `len(f_maps)`
encoders = []
for i, out_feature_num in enumerate(f_maps):
if i == 0:
# apply conv_coord only in the first encoder if any
encoder = Encoder(in_channels, out_feature_num,
apply_pooling=False, # skip pooling in the firs encoder
basic_module=basic_module,
conv_layer_order=layer_order,
conv_kernel_size=conv_kernel_size,
num_groups=num_groups,
padding=conv_padding,
upscale=conv_upscale,
dropout_prob=dropout_prob,
is3d=is3d)
else:
encoder = Encoder(f_maps[i - 1], out_feature_num,
basic_module=basic_module,
conv_layer_order=layer_order,
conv_kernel_size=conv_kernel_size,
num_groups=num_groups,
pool_kernel_size=pool_kernel_size,
padding=conv_padding,
upscale=conv_upscale,
dropout_prob=dropout_prob,
is3d=is3d)
encoders.append(encoder)
return nn.ModuleList(encoders)
def create_decoders(f_maps, basic_module, conv_kernel_size, conv_padding, layer_order,
num_groups, upsample, dropout_prob, is3d):
# create decoder path consisting of the Decoder modules. The length of the decoder list is equal to `len(f_maps) - 1`
decoders = []
reversed_f_maps = list(reversed(f_maps[1:]))
for i in range(len(reversed_f_maps) - 1):
if basic_module == DoubleConv and upsample != 'deconv':
in_feature_num = reversed_f_maps[i] + reversed_f_maps[i + 1]
else:
in_feature_num = reversed_f_maps[i]
out_feature_num = reversed_f_maps[i + 1]
decoder = Decoder(in_feature_num, out_feature_num,
basic_module=basic_module,
conv_layer_order=layer_order,
conv_kernel_size=conv_kernel_size,
num_groups=num_groups,
padding=conv_padding,
upsample=upsample,
dropout_prob=dropout_prob,
is3d=is3d)
decoders.append(decoder)
return nn.ModuleList(decoders)
class AbstractUpsampling(nn.Module):
"""
Abstract class for upsampling. A given implementation should upsample a given 5D input tensor using either
interpolation or learned transposed convolution.
"""
def __init__(self, upsample):
super(AbstractUpsampling, self).__init__()
self.upsample = upsample
def forward(self, encoder_features, x):
# get the spatial dimensions of the output given the encoder_features
output_size = encoder_features.size()[2:]
# upsample the input and return
return self.upsample(x, output_size)
class InterpolateUpsampling(AbstractUpsampling):
"""
Args:
mode (str): algorithm used for upsampling:
'nearest' | 'linear' | 'bilinear' | 'trilinear' | 'area'. Default: 'nearest'
used only if transposed_conv is False
"""
def __init__(self, mode='nearest'):
upsample = partial(self._interpolate, mode=mode)
super().__init__(upsample)
@staticmethod
def _interpolate(x, size, mode):
return F.interpolate(x, size=size, mode=mode)
class TransposeConvUpsampling(AbstractUpsampling):
"""
Args:
in_channels (int): number of input channels for transposed conv
used only if transposed_conv is True
out_channels (int): number of output channels for transpose conv
used only if transposed_conv is True
kernel_size (int or tuple): size of the convolving kernel
used only if transposed_conv is True
scale_factor (int or tuple): stride of the convolution
used only if transposed_conv is True
is3d (bool): if True use ConvTranspose3d, otherwise use ConvTranspose2d
"""
class Upsample(nn.Module):
"""
Workaround the 'ValueError: requested an output size...' in the `_output_padding` method in
transposed convolution. It performs transposed conv followed by the interpolation to the correct size if necessary.
"""
def __init__(self, conv_transposed, is3d):
super().__init__()
self.conv_transposed = conv_transposed
self.is3d = is3d
def forward(self, x, size):
x = self.conv_transposed(x)
return F.interpolate(x, size=size)
def __init__(self, in_channels, out_channels, kernel_size=3, scale_factor=2, is3d=True):
# make sure that the output size reverses the MaxPool3d from the corresponding encoder
if is3d is True:
conv_transposed = nn.ConvTranspose3d(in_channels, out_channels, kernel_size=kernel_size,
stride=scale_factor, padding=1, bias=False)
else:
conv_transposed = nn.ConvTranspose2d(in_channels, out_channels, kernel_size=kernel_size,
stride=scale_factor, padding=1, bias=False)
upsample = self.Upsample(conv_transposed, is3d)
super().__init__(upsample)
class NoUpsampling(AbstractUpsampling):
def __init__(self):
super().__init__(self._no_upsampling)
@staticmethod
def _no_upsampling(x, size):
return x
@@ -0,0 +1,170 @@
# https://github.com/wolny/pytorch-3dunet/blob/master/pytorch3dunet/unet3d/buildingblocks.py
# MIT License
# Copyright (c) 2018 Adrian Wolny
# 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.nn as nn
from partfield.model.UNet.buildingblocks import DoubleConv, ResNetBlock, \
create_decoders, create_encoders
def number_of_features_per_level(init_channel_number, num_levels):
return [init_channel_number * 2 ** k for k in range(num_levels)]
class AbstractUNet(nn.Module):
"""
Base class for standard and residual UNet.
Args:
in_channels (int): number of input channels
out_channels (int): number of output segmentation masks;
Note that the of out_channels might correspond to either
different semantic classes or to different binary segmentation mask.
It's up to the user of the class to interpret the out_channels and
use the proper loss criterion during training (i.e. CrossEntropyLoss (multi-class)
or BCEWithLogitsLoss (two-class) respectively)
f_maps (int, tuple): number of feature maps at each level of the encoder; if it's an integer the number
of feature maps is given by the geometric progression: f_maps ^ k, k=1,2,3,4
final_sigmoid (bool): if True apply element-wise nn.Sigmoid after the final 1x1 convolution,
otherwise apply nn.Softmax. In effect only if `self.training == False`, i.e. during validation/testing
basic_module: basic model for the encoder/decoder (DoubleConv, ResNetBlock, ....)
layer_order (string): determines the order of layers in `SingleConv` module.
E.g. 'crg' stands for GroupNorm3d+Conv3d+ReLU. See `SingleConv` for more info
num_groups (int): number of groups for the GroupNorm
num_levels (int): number of levels in the encoder/decoder path (applied only if f_maps is an int)
default: 4
is_segmentation (bool): if True and the model is in eval mode, Sigmoid/Softmax normalization is applied
after the final convolution; if False (regression problem) the normalization layer is skipped
conv_kernel_size (int or tuple): size of the convolving kernel in the basic_module
pool_kernel_size (int or tuple): the size of the window
conv_padding (int or tuple): add zero-padding added to all three sides of the input
conv_upscale (int): number of the convolution to upscale in encoder if DoubleConv, default: 2
upsample (str): algorithm used for decoder upsampling:
InterpolateUpsampling: 'nearest' | 'linear' | 'bilinear' | 'trilinear' | 'area'
TransposeConvUpsampling: 'deconv'
No upsampling: None
Default: 'default' (chooses automatically)
dropout_prob (float or tuple): dropout probability, default: 0.1
is3d (bool): if True the model is 3D, otherwise 2D, default: True
"""
def __init__(self, in_channels, out_channels, final_sigmoid, basic_module, f_maps=64, layer_order='gcr',
num_groups=8, num_levels=4, is_segmentation=False, conv_kernel_size=3, pool_kernel_size=2,
conv_padding=1, conv_upscale=2, upsample='default', dropout_prob=0.1, is3d=True, encoder_only=False):
super(AbstractUNet, self).__init__()
if isinstance(f_maps, int):
f_maps = number_of_features_per_level(f_maps, num_levels=num_levels)
assert isinstance(f_maps, list) or isinstance(f_maps, tuple)
assert len(f_maps) > 1, "Required at least 2 levels in the U-Net"
if 'g' in layer_order:
assert num_groups is not None, "num_groups must be specified if GroupNorm is used"
# create encoder path
self.encoders = create_encoders(in_channels, f_maps, basic_module, conv_kernel_size,
conv_padding, conv_upscale, dropout_prob,
layer_order, num_groups, pool_kernel_size, is3d)
self.encoder_only = encoder_only
if encoder_only == False:
# create decoder path
self.decoders = create_decoders(f_maps, basic_module, conv_kernel_size, conv_padding,
layer_order, num_groups, upsample, dropout_prob,
is3d)
# in the last layer a 1×1 convolution reduces the number of output channels to the number of labels
if is3d:
self.final_conv = nn.Conv3d(f_maps[1], out_channels, 1)
else:
self.final_conv = nn.Conv2d(f_maps[1], out_channels, 1)
if is_segmentation:
# semantic segmentation problem
if final_sigmoid:
self.final_activation = nn.Sigmoid()
else:
self.final_activation = nn.Softmax(dim=1)
else:
# regression problem
self.final_activation = None
def forward(self, x, return_bottleneck_feat=False):
# encoder part
encoders_features = []
for encoder in self.encoders:
x = encoder(x)
# reverse the encoder outputs to be aligned with the decoder
encoders_features.insert(0, x)
# remove the last encoder's output from the list
# !!remember: it's the 1st in the list
bottleneck_feat = encoders_features[0]
if self.encoder_only:
return bottleneck_feat
else:
encoders_features = encoders_features[1:]
# decoder part
for decoder, encoder_features in zip(self.decoders, encoders_features):
# pass the output from the corresponding encoder and the output
# of the previous decoder
x = decoder(encoder_features, x)
x = self.final_conv(x)
# During training the network outputs logits
if self.final_activation is not None:
x = self.final_activation(x)
if return_bottleneck_feat:
return x, bottleneck_feat
else:
return x
class ResidualUNet3D(AbstractUNet):
"""
Residual 3DUnet model implementation based on https://arxiv.org/pdf/1706.00120.pdf.
Uses ResNetBlock as a basic building block, summation joining instead
of concatenation joining and transposed convolutions for upsampling (watch out for block artifacts).
Since the model effectively becomes a residual net, in theory it allows for deeper UNet.
"""
def __init__(self, in_channels, out_channels, final_sigmoid=True, f_maps=(8, 16, 64, 256, 1024), layer_order='gcr',
num_groups=8, num_levels=5, is_segmentation=True, conv_padding=1,
conv_upscale=2, upsample='default', dropout_prob=0.1, encoder_only=False, **kwargs):
super(ResidualUNet3D, self).__init__(in_channels=in_channels,
out_channels=out_channels,
final_sigmoid=final_sigmoid,
basic_module=ResNetBlock,
f_maps=f_maps,
layer_order=layer_order,
num_groups=num_groups,
num_levels=num_levels,
is_segmentation=is_segmentation,
conv_padding=conv_padding,
conv_upscale=conv_upscale,
upsample=upsample,
dropout_prob=dropout_prob,
encoder_only=encoder_only,
is3d=True)
@@ -0,0 +1,54 @@
import torch
import torch.nn as nn
class VanillaMLP(nn.Module):
def __init__(self, input_dim, output_dim, out_activation, n_hidden_layers=4, n_neurons=64, activation="ReLU"):
super().__init__()
self.n_neurons = n_neurons
self.n_hidden_layers = n_hidden_layers
self.activation = activation
self.out_activation = out_activation
layers = [
self.make_linear(input_dim, self.n_neurons, is_first=True, is_last=False),
self.make_activation(),
]
for i in range(self.n_hidden_layers - 1):
layers += [
self.make_linear(
self.n_neurons, self.n_neurons, is_first=False, is_last=False
),
self.make_activation(),
]
layers += [
self.make_linear(self.n_neurons, output_dim, is_first=False, is_last=True)
]
if self.out_activation == "sigmoid":
layers += [nn.Sigmoid()]
elif self.out_activation == "tanh":
layers += [nn.Tanh()]
elif self.out_activation == "hardtanh":
layers += [nn.Hardtanh()]
elif self.out_activation == "GELU":
layers += [nn.GELU()]
elif self.out_activation == "RELU":
layers += [nn.ReLU()]
else:
raise NotImplementedError
self.layers = nn.Sequential(*layers)
def forward(self, x, split_size=100000):
with torch.cuda.amp.autocast(enabled=False):
out = self.layers(x)
return out
def make_linear(self, dim_in, dim_out, is_first, is_last):
layer = nn.Linear(dim_in, dim_out, bias=False)
return layer
def make_activation(self):
if self.activation == "ReLU":
return nn.ReLU(inplace=True)
elif self.activation == "GELU":
return nn.GELU()
else:
raise NotImplementedError
@@ -0,0 +1,331 @@
#https://github.com/3DTopia/OpenLRM/blob/main/openlrm/models/modeling_lrm.py
# Copyright (c) 2023-2024, Zexin He
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# https://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
import torch
import torch.nn as nn
from functools import partial
def project_onto_planes(planes, coordinates):
"""
Does a projection of a 3D point onto a batch of 2D planes,
returning 2D plane coordinates.
Takes plane axes of shape n_planes, 3, 3
# Takes coordinates of shape N, M, 3
# returns projections of shape N*n_planes, M, 2
"""
N, M, C = coordinates.shape
n_planes, _, _ = planes.shape
coordinates = coordinates.unsqueeze(1).expand(-1, n_planes, -1, -1).reshape(N*n_planes, M, 3)
inv_planes = torch.linalg.inv(planes).unsqueeze(0).expand(N, -1, -1, -1).reshape(N*n_planes, 3, 3)
projections = torch.bmm(coordinates, inv_planes)
return projections[..., :2]
def sample_from_planes(plane_features, coordinates, mode='bilinear', padding_mode='zeros', box_warp=None):
plane_axes = torch.tensor([[[1, 0, 0],
[0, 1, 0],
[0, 0, 1]],
[[1, 0, 0],
[0, 0, 1],
[0, 1, 0]],
[[0, 0, 1],
[0, 1, 0],
[1, 0, 0]]], dtype=torch.float32).cuda()
assert padding_mode == 'zeros'
N, n_planes, C, H, W = plane_features.shape
_, M, _ = coordinates.shape
plane_features = plane_features.view(N*n_planes, C, H, W)
projected_coordinates = project_onto_planes(plane_axes, coordinates).unsqueeze(1)
output_features = torch.nn.functional.grid_sample(plane_features, projected_coordinates.float(), mode=mode, padding_mode=padding_mode, align_corners=False).permute(0, 3, 2, 1).reshape(N, n_planes, M, C)
return output_features
def get_grid_coord(grid_size = 256, align_corners=False):
if align_corners == False:
coords = torch.linspace(-1 + 1/(grid_size), 1 - 1/(grid_size), steps=grid_size)
else:
coords = torch.linspace(-1, 1, steps=grid_size)
i, j, k = torch.meshgrid(coords, coords, coords, indexing='ij')
coordinates = torch.stack((i, j, k), dim=-1).reshape(-1, 3)
return coordinates
class BasicBlock(nn.Module):
"""
Transformer block that is in its simplest form.
Designed for PF-LRM architecture.
"""
# Block contains a self-attention layer and an MLP
def __init__(self, inner_dim: int, num_heads: int, eps: float,
attn_drop: float = 0., attn_bias: bool = False,
mlp_ratio: float = 4., mlp_drop: float = 0.):
super().__init__()
self.norm1 = nn.LayerNorm(inner_dim, eps=eps)
self.self_attn = nn.MultiheadAttention(
embed_dim=inner_dim, num_heads=num_heads,
dropout=attn_drop, bias=attn_bias, batch_first=True)
self.norm2 = nn.LayerNorm(inner_dim, eps=eps)
self.mlp = nn.Sequential(
nn.Linear(inner_dim, int(inner_dim * mlp_ratio)),
nn.GELU(),
nn.Dropout(mlp_drop),
nn.Linear(int(inner_dim * mlp_ratio), inner_dim),
nn.Dropout(mlp_drop),
)
def forward(self, x):
# x: [N, L, D]
before_sa = self.norm1(x)
x = x + self.self_attn(before_sa, before_sa, before_sa, need_weights=False)[0]
x = x + self.mlp(self.norm2(x))
return x
class ConditionBlock(nn.Module):
"""
Transformer block that takes in a cross-attention condition.
Designed for SparseLRM architecture.
"""
# Block contains a cross-attention layer, a self-attention layer, and an MLP
def __init__(self, inner_dim: int, cond_dim: int, num_heads: int, eps: float,
attn_drop: float = 0., attn_bias: bool = False,
mlp_ratio: float = 4., mlp_drop: float = 0.):
super().__init__()
self.norm1 = nn.LayerNorm(inner_dim, eps=eps)
self.cross_attn = nn.MultiheadAttention(
embed_dim=inner_dim, num_heads=num_heads, kdim=cond_dim, vdim=cond_dim,
dropout=attn_drop, bias=attn_bias, batch_first=True)
self.norm2 = nn.LayerNorm(inner_dim, eps=eps)
self.self_attn = nn.MultiheadAttention(
embed_dim=inner_dim, num_heads=num_heads,
dropout=attn_drop, bias=attn_bias, batch_first=True)
self.norm3 = nn.LayerNorm(inner_dim, eps=eps)
self.mlp = nn.Sequential(
nn.Linear(inner_dim, int(inner_dim * mlp_ratio)),
nn.GELU(),
nn.Dropout(mlp_drop),
nn.Linear(int(inner_dim * mlp_ratio), inner_dim),
nn.Dropout(mlp_drop),
)
def forward(self, x, cond):
# x: [N, L, D]
# cond: [N, L_cond, D_cond]
x = x + self.cross_attn(self.norm1(x), cond, cond, need_weights=False)[0]
before_sa = self.norm2(x)
x = x + self.self_attn(before_sa, before_sa, before_sa, need_weights=False)[0]
x = x + self.mlp(self.norm3(x))
return x
class TransformerDecoder(nn.Module):
def __init__(self, block_type: str,
num_layers: int, num_heads: int,
inner_dim: int, cond_dim: int = None,
eps: float = 1e-6):
super().__init__()
self.block_type = block_type
self.layers = nn.ModuleList([
self._block_fn(inner_dim, cond_dim)(
num_heads=num_heads,
eps=eps,
)
for _ in range(num_layers)
])
self.norm = nn.LayerNorm(inner_dim, eps=eps)
@property
def block_type(self):
return self._block_type
@block_type.setter
def block_type(self, block_type):
assert block_type in ['cond', 'basic'], \
f"Unsupported block type: {block_type}"
self._block_type = block_type
def _block_fn(self, inner_dim, cond_dim):
assert inner_dim is not None, f"inner_dim must always be specified"
if self.block_type == 'basic':
return partial(BasicBlock, inner_dim=inner_dim)
elif self.block_type == 'cond':
assert cond_dim is not None, f"Condition dimension must be specified for ConditionBlock"
return partial(ConditionBlock, inner_dim=inner_dim, cond_dim=cond_dim)
else:
raise ValueError(f"Unsupported block type during runtime: {self.block_type}")
def forward_layer(self, layer: nn.Module, x: torch.Tensor, cond: torch.Tensor,):
if self.block_type == 'basic':
return layer(x)
elif self.block_type == 'cond':
return layer(x, cond)
else:
raise NotImplementedError
def forward(self, x: torch.Tensor, cond: torch.Tensor = None):
# x: [N, L, D]
# cond: [N, L_cond, D_cond] or None
for layer in self.layers:
x = self.forward_layer(layer, x, cond)
x = self.norm(x)
return x
class Voxel2Triplane(nn.Module):
"""
Full model of the basic single-view large reconstruction model.
"""
def __init__(self, transformer_dim: int, transformer_layers: int, transformer_heads: int,
triplane_low_res: int, triplane_high_res: int, triplane_dim: int, voxel_feat_dim: int, normalize_vox_feat=False, voxel_dim=16):
super().__init__()
# attributes
self.triplane_low_res = triplane_low_res
self.triplane_high_res = triplane_high_res
self.triplane_dim = triplane_dim
self.voxel_feat_dim = voxel_feat_dim
# initialize pos_embed with 1/sqrt(dim) * N(0, 1)
self.pos_embed = nn.Parameter(torch.randn(1, 3*triplane_low_res**2, transformer_dim) * (1. / transformer_dim) ** 0.5)
self.transformer = TransformerDecoder(
block_type='cond',
num_layers=transformer_layers, num_heads=transformer_heads,
inner_dim=transformer_dim, cond_dim=voxel_feat_dim
)
self.upsampler = nn.ConvTranspose2d(transformer_dim, triplane_dim, kernel_size=8, stride=8, padding=0)
self.normalize_vox_feat = normalize_vox_feat
if normalize_vox_feat:
self.vox_norm = nn.LayerNorm(voxel_feat_dim, eps=1e-6)
self.vox_pos_embed = nn.Parameter(torch.randn(1, voxel_dim * voxel_dim * voxel_dim, voxel_feat_dim) * (1. / voxel_feat_dim) ** 0.5)
def forward_transformer(self, voxel_feats):
N = voxel_feats.shape[0]
x = self.pos_embed.repeat(N, 1, 1) # [N, L, D]
if self.normalize_vox_feat:
vox_pos_embed = self.vox_pos_embed.repeat(N, 1, 1) # [N, L, D]
voxel_feats = self.vox_norm(voxel_feats + vox_pos_embed)
x = self.transformer(
x,
cond=voxel_feats
)
return x
def reshape_upsample(self, tokens):
N = tokens.shape[0]
H = W = self.triplane_low_res
x = tokens.view(N, 3, H, W, -1)
x = torch.einsum('nihwd->indhw', x) # [3, N, D, H, W]
x = x.contiguous().view(3*N, -1, H, W) # [3*N, D, H, W]
x = self.upsampler(x) # [3*N, D', H', W']
x = x.view(3, N, *x.shape[-3:]) # [3, N, D', H', W']
x = torch.einsum('indhw->nidhw', x) # [N, 3, D', H', W']
x = x.contiguous()
return x
def forward(self, voxel_feats):
N = voxel_feats.shape[0]
# encode image
assert voxel_feats.shape[-1] == self.voxel_feat_dim, \
f"Feature dimension mismatch: {voxel_feats.shape[-1]} vs {self.voxel_feat_dim}"
# transformer generating planes
tokens = self.forward_transformer(voxel_feats)
planes = self.reshape_upsample(tokens)
assert planes.shape[0] == N, "Batch size mismatch for planes"
assert planes.shape[1] == 3, "Planes should have 3 channels"
return planes
class TriplaneTransformer(nn.Module):
"""
Full model of the basic single-view large reconstruction model.
"""
def __init__(self, input_dim: int, transformer_dim: int, transformer_layers: int, transformer_heads: int,
triplane_low_res: int, triplane_high_res: int, triplane_dim: int):
super().__init__()
# attributes
self.triplane_low_res = triplane_low_res
self.triplane_high_res = triplane_high_res
self.triplane_dim = triplane_dim
# initialize pos_embed with 1/sqrt(dim) * N(0, 1)
self.pos_embed = nn.Parameter(torch.randn(1, 3*triplane_low_res**2, transformer_dim) * (1. / transformer_dim) ** 0.5)
self.transformer = TransformerDecoder(
block_type='basic',
num_layers=transformer_layers, num_heads=transformer_heads,
inner_dim=transformer_dim,
)
self.downsampler = nn.Sequential(
nn.Conv2d(input_dim, transformer_dim, kernel_size=3, stride=1, padding=1),
nn.ReLU(),
nn.MaxPool2d(kernel_size=2, stride=2), # Reduces size from 128x128 to 64x64
nn.Conv2d(transformer_dim, transformer_dim, kernel_size=3, stride=1, padding=1),
nn.ReLU(),
nn.MaxPool2d(kernel_size=2, stride=2), # Reduces size from 64x64 to 32x32
)
self.upsampler = nn.ConvTranspose2d(transformer_dim, triplane_dim, kernel_size=4, stride=4, padding=0)
self.mlp = nn.Sequential(
nn.Linear(input_dim, triplane_dim),
nn.ReLU(),
nn.Linear(triplane_dim, triplane_dim)
)
def forward_transformer(self, triplanes):
N = triplanes.shape[0]
tokens = torch.einsum('nidhw->nihwd', triplanes).reshape(N, self.pos_embed.shape[1], -1) # [N, L, D]
x = self.pos_embed.repeat(N, 1, 1) + tokens # [N, L, D]
x = self.transformer(x)
return x
def reshape_downsample(self, triplanes):
N = triplanes.shape[0]
H = W = self.triplane_high_res
x = triplanes.view(N, 3, -1, H, W)
x = torch.einsum('nidhw->indhw', x) # [3, N, D, H, W]
x = x.contiguous().view(3*N, -1, H, W) # [3*N, D, H, W]
x = self.downsampler(x) # [3*N, D', H', W']
x = x.view(3, N, *x.shape[-3:]) # [3, N, D', H', W']
x = torch.einsum('indhw->nidhw', x) # [N, 3, D', H', W']
x = x.contiguous()
return x
def reshape_upsample(self, tokens):
N = tokens.shape[0]
H = W = self.triplane_low_res
x = tokens.view(N, 3, H, W, -1)
x = torch.einsum('nihwd->indhw', x) # [3, N, D, H, W]
x = x.contiguous().view(3*N, -1, H, W) # [3*N, D, H, W]
x = self.upsampler(x) # [3*N, D', H', W']
x = x.view(3, N, *x.shape[-3:]) # [3, N, D', H', W']
x = torch.einsum('indhw->nidhw', x) # [N, 3, D', H', W']
x = x.contiguous()
return x
def forward(self, triplanes):
downsampled_triplanes = self.reshape_downsample(triplanes)
tokens = self.forward_transformer(downsampled_triplanes)
residual = self.reshape_upsample(tokens)
triplanes = triplanes.permute(0, 1, 3, 4, 2).contiguous()
triplanes = self.mlp(triplanes)
triplanes = triplanes.permute(0, 1, 4, 2, 3).contiguous()
planes = triplanes + residual
return planes
@@ -0,0 +1,283 @@
import torch
import lightning.pytorch as pl
from .dataloader import Demo_Dataset, Demo_Remesh_Dataset, Correspondence_Demo_Dataset
from torch.utils.data import DataLoader
from partfield.model.UNet.model import ResidualUNet3D
from partfield.model.triplane import TriplaneTransformer, get_grid_coord #, sample_from_planes, Voxel2Triplane
from partfield.model.model_utils import VanillaMLP
import torch.nn.functional as F
import torch.nn as nn
import os
import trimesh
import skimage
import numpy as np
import h5py
import torch.distributed as dist
from partfield.model.PVCNN.encoder_pc import TriPlanePC2Encoder, sample_triplane_feat
import json
import gc
import time
from plyfile import PlyData, PlyElement
class Model(pl.LightningModule):
def __init__(self, cfg):
super().__init__()
self.save_hyperparameters()
self.cfg = cfg
self.automatic_optimization = False
self.triplane_resolution = cfg.triplane_resolution
self.triplane_channels_low = cfg.triplane_channels_low
self.triplane_transformer = TriplaneTransformer(
input_dim=cfg.triplane_channels_low * 2,
transformer_dim=1024,
transformer_layers=6,
transformer_heads=8,
triplane_low_res=32,
triplane_high_res=128,
triplane_dim=cfg.triplane_channels_high,
)
self.sdf_decoder = VanillaMLP(input_dim=64,
output_dim=1,
out_activation="tanh",
n_neurons=64, #64
n_hidden_layers=6) #6
self.use_pvcnn = cfg.use_pvcnnonly
self.use_2d_feat = cfg.use_2d_feat
if self.use_pvcnn:
self.pvcnn = TriPlanePC2Encoder(
cfg.pvcnn,
device="cuda",
shape_min=-1,
shape_length=2,
use_2d_feat=self.use_2d_feat) #.cuda()
self.logit_scale = nn.Parameter(torch.tensor([1.0], requires_grad=True))
self.grid_coord = get_grid_coord(256)
self.mse_loss = torch.nn.MSELoss()
self.l1_loss = torch.nn.L1Loss(reduction='none')
if cfg.regress_2d_feat:
self.feat_decoder = VanillaMLP(input_dim=64,
output_dim=192,
out_activation="GELU",
n_neurons=64, #64
n_hidden_layers=6) #6
def predict_dataloader(self):
if self.cfg.remesh_demo:
dataset = Demo_Remesh_Dataset(self.cfg)
elif self.cfg.correspondence_demo:
dataset = Correspondence_Demo_Dataset(self.cfg)
else:
dataset = Demo_Dataset(self.cfg)
dataloader = DataLoader(dataset,
num_workers=self.cfg.dataset.val_num_workers,
batch_size=self.cfg.dataset.val_batch_size,
shuffle=False,
pin_memory=True,
drop_last=False)
return dataloader
@torch.no_grad()
def predict_step(self, batch, batch_idx):
save_dir = f"{self.cfg.result_name}"
os.makedirs(save_dir, exist_ok=True)
uid = batch['uid'][0]
view_id = 0
starttime = time.time()
if uid == "car" or uid == "complex_car":
# if uid == "complex_car":
print("Skipping this for now.")
print(uid)
return
### Skip if model already processed
if os.path.exists(f'{save_dir}/part_feat_{uid}_{view_id}.npy') or os.path.exists(f'{save_dir}/part_feat_{uid}_{view_id}_batch.npy'):
print("Already processed "+uid)
return
N = batch['pc'].shape[0]
assert N == 1
if self.use_2d_feat:
print("ERROR. Dataloader not implemented with input 2d feat.")
exit()
else:
pc_feat = self.pvcnn(batch['pc'], batch['pc'])
planes = pc_feat
planes = self.triplane_transformer(planes)
sdf_planes, part_planes = torch.split(planes, [64, planes.shape[2] - 64], dim=2)
if self.cfg.is_pc:
tensor_vertices = batch['pc'].reshape(1, -1, 3).cuda().to(torch.float16)
point_feat = sample_triplane_feat(part_planes, tensor_vertices) # N, M, C
point_feat = point_feat.cpu().detach().numpy().reshape(-1, 448)
np.save(f'{save_dir}/part_feat_{uid}_{view_id}.npy', point_feat)
print(f"Exported part_feat_{uid}_{view_id}.npy")
###########
from sklearn.decomposition import PCA
data_scaled = point_feat / np.linalg.norm(point_feat, axis=-1, keepdims=True)
pca = PCA(n_components=3)
data_reduced = pca.fit_transform(data_scaled)
data_reduced = (data_reduced - data_reduced.min()) / (data_reduced.max() - data_reduced.min())
colors_255 = (data_reduced * 255).astype(np.uint8)
points = batch['pc'].squeeze().detach().cpu().numpy()
if colors_255 is None:
colors_255 = np.full_like(points, 255) # Default to white color (255,255,255)
else:
assert colors_255.shape == points.shape, "Colors must have the same shape as points"
# Convert to structured array for PLY format
vertex_data = np.array(
[(*point, *color) for point, color in zip(points, colors_255)],
dtype=[("x", "f4"), ("y", "f4"), ("z", "f4"), ("red", "u1"), ("green", "u1"), ("blue", "u1")]
)
# Create PLY element
el = PlyElement.describe(vertex_data, "vertex")
# Write to file
filename = f'{save_dir}/feat_pca_{uid}_{view_id}.ply'
PlyData([el], text=True).write(filename)
print(f"Saved PLY file: {filename}")
############
else:
use_cuda_version = True
if use_cuda_version:
def sample_points(vertices, faces, n_point_per_face):
# Generate random barycentric coordinates
# borrowed from Kaolin https://github.com/NVIDIAGameWorks/kaolin/blob/master/kaolin/ops/mesh/trianglemesh.py#L43
n_f = faces.shape[0]
u = torch.sqrt(torch.rand((n_f, n_point_per_face, 1),
device=vertices.device,
dtype=vertices.dtype))
v = torch.rand((n_f, n_point_per_face, 1),
device=vertices.device,
dtype=vertices.dtype)
w0 = 1 - u
w1 = u * (1 - v)
w2 = u * v
face_v_0 = torch.index_select(vertices, 0, faces[:, 0].reshape(-1))
face_v_1 = torch.index_select(vertices, 0, faces[:, 1].reshape(-1))
face_v_2 = torch.index_select(vertices, 0, faces[:, 2].reshape(-1))
points = w0 * face_v_0.unsqueeze(dim=1) + w1 * face_v_1.unsqueeze(dim=1) + w2 * face_v_2.unsqueeze(dim=1)
return points
def sample_and_mean_memory_save_version(part_planes, tensor_vertices, n_point_per_face):
n_sample_each = self.cfg.n_sample_each # we iterate over this to avoid OOM
n_v = tensor_vertices.shape[1]
n_sample = n_v // n_sample_each + 1
all_sample = []
for i_sample in range(n_sample):
sampled_feature = sample_triplane_feat(part_planes, tensor_vertices[:, i_sample * n_sample_each: i_sample * n_sample_each + n_sample_each,])
assert sampled_feature.shape[1] % n_point_per_face == 0
sampled_feature = sampled_feature.reshape(1, -1, n_point_per_face, sampled_feature.shape[-1])
sampled_feature = torch.mean(sampled_feature, axis=-2)
all_sample.append(sampled_feature)
return torch.cat(all_sample, dim=1)
if self.cfg.vertex_feature:
tensor_vertices = batch['vertices'][0].reshape(1, -1, 3).to(torch.float32)
point_feat = sample_and_mean_memory_save_version(part_planes, tensor_vertices, 1)
else:
n_point_per_face = self.cfg.n_point_per_face
tensor_vertices = sample_points(batch['vertices'][0], batch['faces'][0], n_point_per_face)
tensor_vertices = tensor_vertices.reshape(1, -1, 3).to(torch.float32)
point_feat = sample_and_mean_memory_save_version(part_planes, tensor_vertices, n_point_per_face) # N, M, C
#### Take mean feature in the triangle
print("Time elapsed for feature prediction: " + str(time.time() - starttime))
point_feat = point_feat.reshape(-1, 448).cpu().numpy()
np.save(f'{save_dir}/part_feat_{uid}_{view_id}_batch.npy', point_feat)
print(f"Exported part_feat_{uid}_{view_id}.npy")
###########
from sklearn.decomposition import PCA
data_scaled = point_feat / np.linalg.norm(point_feat, axis=-1, keepdims=True)
pca = PCA(n_components=3)
data_reduced = pca.fit_transform(data_scaled)
data_reduced = (data_reduced - data_reduced.min()) / (data_reduced.max() - data_reduced.min())
colors_255 = (data_reduced * 255).astype(np.uint8)
V = batch['vertices'][0].cpu().numpy()
F = batch['faces'][0].cpu().numpy()
if self.cfg.vertex_feature:
colored_mesh = trimesh.Trimesh(vertices=V, faces=F, vertex_colors=colors_255, process=False)
else:
colored_mesh = trimesh.Trimesh(vertices=V, faces=F, face_colors=colors_255, process=False)
colored_mesh.export(f'{save_dir}/feat_pca_{uid}_{view_id}.ply')
############
torch.cuda.empty_cache()
else:
### Mesh input (obj file)
V = batch['vertices'][0].cpu().numpy()
F = batch['faces'][0].cpu().numpy()
##### Loop through faces #####
num_samples_per_face = self.cfg.n_point_per_face
all_point_feats = []
for face in F:
# Get the vertices of the current face
v0, v1, v2 = V[face]
# Generate random barycentric coordinates
u = np.random.rand(num_samples_per_face, 1)
v = np.random.rand(num_samples_per_face, 1)
is_prob = (u+v) >1
u[is_prob] = 1 - u[is_prob]
v[is_prob] = 1 - v[is_prob]
w = 1 - u - v
# Calculate points in Cartesian coordinates
points = u * v0 + v * v1 + w * v2
tensor_vertices = torch.from_numpy(points.copy()).reshape(1, -1, 3).cuda().to(torch.float32)
point_feat = sample_triplane_feat(part_planes, tensor_vertices) # N, M, C
#### Take mean feature in the triangle
point_feat = torch.mean(point_feat, axis=1).cpu().detach().numpy()
all_point_feats.append(point_feat)
##############################
all_point_feats = np.array(all_point_feats).reshape(-1, 448)
point_feat = all_point_feats
np.save(f'{save_dir}/part_feat_{uid}_{view_id}.npy', point_feat)
print(f"Exported part_feat_{uid}_{view_id}.npy")
###########
from sklearn.decomposition import PCA
data_scaled = point_feat / np.linalg.norm(point_feat, axis=-1, keepdims=True)
pca = PCA(n_components=3)
data_reduced = pca.fit_transform(data_scaled)
data_reduced = (data_reduced - data_reduced.min()) / (data_reduced.max() - data_reduced.min())
colors_255 = (data_reduced * 255).astype(np.uint8)
colored_mesh = trimesh.Trimesh(vertices=V, faces=F, face_colors=colors_255, process=False)
colored_mesh.export(f'{save_dir}/feat_pca_{uid}_{view_id}.ply')
############
print("Time elapsed: " + str(time.time()-starttime))
return
@@ -0,0 +1,103 @@
import torch
import lightning.pytorch as pl
# from .dataloader import Demo_Dataset, Demo_Remesh_Dataset, Correspondence_Demo_Dataset
from torch.utils.data import DataLoader
from .model.UNet.model import ResidualUNet3D
from .model.triplane import TriplaneTransformer, get_grid_coord #, sample_from_planes, Voxel2Triplane
from .model.model_utils import VanillaMLP
import torch.nn.functional as F
import torch.nn as nn
import os
import trimesh
import skimage
import numpy as np
import h5py
import torch.distributed as dist
from .model.PVCNN.encoder_pc import TriPlanePC2Encoder, sample_triplane_feat
import json
import gc
import time
from plyfile import PlyData, PlyElement
class Model(pl.LightningModule):
def __init__(self, cfg):
super().__init__()
self.save_hyperparameters()
self.cfg = cfg
self.automatic_optimization = False
self.triplane_resolution = cfg.triplane_resolution
self.triplane_channels_low = cfg.triplane_channels_low
self.triplane_transformer = TriplaneTransformer(
input_dim=cfg.triplane_channels_low * 2,
transformer_dim=1024,
transformer_layers=6,
transformer_heads=8,
triplane_low_res=32,
triplane_high_res=128,
triplane_dim=cfg.triplane_channels_high,
)
self.sdf_decoder = VanillaMLP(input_dim=64,
output_dim=1,
out_activation="tanh",
n_neurons=64, #64
n_hidden_layers=6) #6
self.use_pvcnn = cfg.use_pvcnnonly
self.use_2d_feat = cfg.use_2d_feat
if self.use_pvcnn:
self.pvcnn = TriPlanePC2Encoder(
cfg.pvcnn,
device="cuda",
shape_min=-1,
shape_length=2,
use_2d_feat=self.use_2d_feat) #.cuda()
self.logit_scale = nn.Parameter(torch.tensor([1.0], requires_grad=True))
self.grid_coord = get_grid_coord(256)
self.mse_loss = torch.nn.MSELoss()
self.l1_loss = torch.nn.L1Loss(reduction='none')
if cfg.regress_2d_feat:
self.feat_decoder = VanillaMLP(input_dim=64,
output_dim=192,
out_activation="GELU",
n_neurons=64, #64
n_hidden_layers=6) #6
# def predict_dataloader(self):
# if self.cfg.remesh_demo:
# dataset = Demo_Remesh_Dataset(self.cfg)
# elif self.cfg.correspondence_demo:
# dataset = Correspondence_Demo_Dataset(self.cfg)
# else:
# dataset = Demo_Dataset(self.cfg)
# dataloader = DataLoader(dataset,
# num_workers=self.cfg.dataset.val_num_workers,
# batch_size=self.cfg.dataset.val_batch_size,
# shuffle=False,
# pin_memory=True,
# drop_last=False)
# return dataloader
@torch.no_grad()
def encode(self, points):
N = points.shape[0]
# assert N == 1
pcd = points[..., :3]
pc_feat = self.pvcnn(pcd, pcd)
planes = pc_feat
planes = self.triplane_transformer(planes)
sdf_planes, part_planes = torch.split(planes, [64, planes.shape[2] - 64], dim=2)
tensor_vertices = pcd.reshape(N, -1, 3).cuda().to(pcd.dtype)
point_feat = sample_triplane_feat(part_planes, tensor_vertices) # N, M, C
# point_feat = point_feat.cpu().detach().numpy().reshape(-1, 448)
point_feat = point_feat.reshape(N, -1, 448)
return point_feat
@@ -0,0 +1,5 @@
import trimesh
def load_mesh_util(input_fname):
mesh = trimesh.load(input_fname, force='mesh', process=False)
return mesh
+57
View File
@@ -0,0 +1,57 @@
import os
from omegaconf import OmegaConf, DictConfig
from dataclasses import dataclass, field
from typing import Any, Dict, List, Optional, Union
from datetime import datetime
@dataclass
class ExperimentConfig:
name: str = "default"
tag: str = ""
use_timestamp: bool = False
timestamp: Optional[str] = None
exp_root_dir: str = "outputs"
### these shouldn't be set manually
exp_dir: str = "outputs/default"
trial_name: str = "exp"
trial_dir: str = "outputs/default/exp"
###
resume: Optional[str] = None
ckpt_path: Optional[str] = None
data: dict = field(default_factory=dict)
model_pl: dict = field(default_factory=dict)
trainer: dict = field(default_factory=dict)
checkpoint: dict = field(default_factory=dict)
checkpoint_epoch: Optional[dict] = None
wandb: dict = field(default_factory=dict)
def load_config(*yamls: str, cli_args: list = [], from_string=False, **kwargs) -> Any:
if from_string:
yaml_confs = [OmegaConf.create(s) for s in yamls]
else:
yaml_confs = [OmegaConf.load(f) for f in yamls]
cli_conf = OmegaConf.from_cli(cli_args)
cfg = OmegaConf.merge(*yaml_confs, cli_conf, kwargs)
OmegaConf.resolve(cfg)
assert isinstance(cfg, DictConfig)
scfg = parse_structured(ExperimentConfig, cfg)
return scfg
def config_to_primitive(config, resolve: bool = True) -> Any:
return OmegaConf.to_container(config, resolve=resolve)
def dump_config(path: str, config) -> None:
with open(path, "w") as fp:
OmegaConf.save(config=config, f=fp)
def parse_structured(fields: Any, cfg: Optional[Union[dict, DictConfig]] = None) -> Any:
scfg = OmegaConf.structured(fields(**cfg))
return scfg
@@ -0,0 +1,305 @@
from dataclasses import dataclass
import os
import sys
import torch
import trimesh
from torch import nn
from transformers import AutoModelForCausalLM
from transformers.generation.logits_process import LogitsProcessorList
from einops import rearrange
from .image_encoder import DINOv2ImageEncoder
from ..config import parse_structured
from .bboxopt import BBoxOPT, BBoxOPTConfig
from ..utils.bbox_tokenizer import BoundsTokenizerDiag
from .bbox_gen_models import GroupEmbedding, MultiModalProjector, MeshDecodeLogitsProcessor, SparseStructureEncoder
current_dir = os.path.dirname(os.path.abspath(__file__))
modules_dir = os.path.dirname(os.path.dirname(current_dir))
partfield_dir = os.path.join(modules_dir, 'PartField')
if partfield_dir not in sys.path:
sys.path.insert(0, partfield_dir)
import importlib.util
from ...PartField.partfield.config import default_argument_parser, setup
class BboxGen(nn.Module):
@dataclass
class Config:
# encoder config
encoder_dim_feat: int = 3
encoder_dim: int = 64
encoder_heads: int = 4
encoder_token_num: int = 256
encoder_qkv_bias: bool = False
encoder_use_ln_post: bool = True
encoder_use_checkpoint: bool = False
encoder_num_embed_freqs: int = 8
encoder_embed_include_pi: bool = False
encoder_init_scale: float = 0.25
encoder_random_fps: bool = True
encoder_learnable_query: bool = False
encoder_layers: int = 4
group_embedding_dim: int = 64
# decoder config
vocab_size: int = 518
decoder_hidden_size: int = 1536
decoder_num_hidden_layers: int = 24
decoder_ffn_dim: int = 6144
decoder_heads: int = 16
decoder_use_flash_attention: bool = True
decoder_gradient_checkpointing: bool = True
# data config
bins: int = 64
BOS_id: int = 64
EOS_id: int = 65
PAD_id: int = 66
max_length: int = 2187 # bos + 50x2x3 + 1374 + 512
voxel_token_length: int = 1886
voxel_token_placeholder: int = -1
# tokenizer config
max_group_size: int = 50
# voxel encoder
partfield_encoder_path: str = ""
cfg: Config
def __init__(self, cfg,model_name="facebook/dinov2-with-registers-large",ckpt=""):
super().__init__()
self.cfg = parse_structured(self.Config, cfg)
self.image_encoder = DINOv2ImageEncoder(
model_name=model_name,ckpt=ckpt
)
self.image_projector = MultiModalProjector(
in_features=(1024 + self.cfg.group_embedding_dim),
out_features=self.cfg.decoder_hidden_size,
)
self.group_embedding = GroupEmbedding(
max_group_size=self.cfg.max_group_size,
hidden_size=self.cfg.group_embedding_dim,
)
self.decoder_config = BBoxOPTConfig(
vocab_size=self.cfg.vocab_size,
hidden_size=self.cfg.decoder_hidden_size,
num_hidden_layers=self.cfg.decoder_num_hidden_layers,
ffn_dim=self.cfg.decoder_ffn_dim,
max_position_embeddings=self.cfg.max_length,
num_attention_heads=self.cfg.decoder_heads,
pad_token_id=self.cfg.PAD_id,
bos_token_id=self.cfg.BOS_id,
eos_token_id=self.cfg.EOS_id,
use_cache=True,
init_std=0.02,
)
if self.cfg.decoder_use_flash_attention:
self.decoder: BBoxOPT = AutoModelForCausalLM.from_config(
self.decoder_config,
torch_dtype=torch.bfloat16,
attn_implementation="flash_attention_2"
)
else:
self.decoder: BBoxOPT = AutoModelForCausalLM.from_config(
self.decoder_config,
)
if self.cfg.decoder_gradient_checkpointing:
self.decoder.gradient_checkpointing_enable()
self.logits_processor = LogitsProcessorList()
self.logits_processor.append(MeshDecodeLogitsProcessor(
bins=self.cfg.bins,
BOS_id=self.cfg.BOS_id,
EOS_id=self.cfg.EOS_id,
PAD_id=self.cfg.PAD_id,
vertices_num=2,
))
self.tokenizer = BoundsTokenizerDiag(
bins=self.cfg.bins,
BOS_id=self.cfg.BOS_id,
EOS_id=self.cfg.EOS_id,
PAD_id=self.cfg.PAD_id,
)
self._load_partfield_encoder()
self.partfield_voxel_encoder = SparseStructureEncoder(
in_channels=451,
channels=[448, 448, 448, 1024],
latent_channels=448,
num_res_blocks=1,
num_res_blocks_middle=1,
norm_type="layer",
)
def _load_partfield_encoder(self):
# Load PartField encoder
model_spec = importlib.util.spec_from_file_location(
"partfield.partfield_encoder",
os.path.join(partfield_dir, "partfield", "partfield_encoder.py")
)
model_module = importlib.util.module_from_spec(model_spec)
model_spec.loader.exec_module(model_module)
Model = model_module.Model
parser = default_argument_parser()
args = []
args.extend(["-c", os.path.join(partfield_dir, "configs/final/demo.yaml")])
args.append("--opts")
args.extend(["continue_ckpt", self.cfg.partfield_encoder_path])
parsed_args = parser.parse_args(args)
cfg = setup(parsed_args, freeze=False)
self.partfield_encoder = Model(cfg)
self.partfield_encoder.eval()
weights = torch.load(self.cfg.partfield_encoder_path,weights_only=False)["state_dict"]
self.partfield_encoder.load_state_dict(weights)
for param in self.partfield_encoder.parameters():
param.requires_grad = False
print("PartField encoder loaded")
def _prepare_lm_inputs(self, voxel_token, input_ids):
inputs_embeds = torch.zeros(input_ids.shape[0], input_ids.shape[1], self.cfg.decoder_hidden_size, device=input_ids.device, dtype=voxel_token.dtype)
voxel_token_mask = (input_ids == self.cfg.voxel_token_placeholder)
inputs_embeds[voxel_token_mask] = voxel_token.view(-1, self.cfg.decoder_hidden_size)
inputs_embeds[~voxel_token_mask] = self.decoder.get_input_embeddings()(input_ids[~voxel_token_mask]).to(dtype=inputs_embeds.dtype)
attention_mask = (input_ids != self.cfg.PAD_id)
return inputs_embeds, attention_mask.long()
def forward(self, batch):
image_latents = self.image_encoder(batch['images'])
masks = batch['masks']
masks_emb = self.group_embedding(masks)
masks_emb = rearrange(masks_emb, 'b c h w -> b (h w) c') # B x Q x C
group_emb = torch.zeros((image_latents.shape[0], image_latents.shape[1], masks_emb.shape[2]), device=image_latents.device, dtype=image_latents.dtype)
group_emb[:, :masks_emb.shape[1], :] = masks_emb
image_latents = torch.cat([image_latents, group_emb], dim=-1)
image_latents = self.image_projector(image_latents)
points = batch['points'][..., :3]
rot_matrix = torch.tensor([[1, 0, 0], [0, 0, -1], [0, 1, 0]], device=points.device, dtype=points.dtype)
rot_points = torch.matmul(points, rot_matrix)
rot_points = rot_points * (2 * 0.9) # from (-0.5, 0.5) to (-1, 1)
partfield_feat = self.partfield_encoder.encode(rot_points)
feat_volume = torch.zeros((points.shape[0], 448, 64, 64, 64), device=partfield_feat.device, dtype=partfield_feat.dtype)
whole_voxel_index = batch['whole_voxel_index'] # (b, m, 3)
batch_size, num_points = whole_voxel_index.shape[0], whole_voxel_index.shape[1]
batch_indices = torch.arange(batch_size, device=whole_voxel_index.device).unsqueeze(1).expand(-1, num_points) # (b, m)
batch_flat = batch_indices.flatten() # (b*m,)
x_flat = whole_voxel_index[..., 0].flatten() # (b*m,)
y_flat = whole_voxel_index[..., 1].flatten() # (b*m,)
z_flat = whole_voxel_index[..., 2].flatten() # (b*m,)
partfield_feat_flat = partfield_feat.reshape(-1, 448) # (b*m, 448)
feat_volume[batch_flat, :, x_flat, y_flat, z_flat] = partfield_feat_flat
xyz_volume = torch.zeros((points.shape[0], 3, 64, 64, 64), device=points.device, dtype=points.dtype)
xyz_volume[batch_flat, :, x_flat, y_flat, z_flat] = points.reshape(-1, 3)
feat_volume = torch.cat([feat_volume, xyz_volume], dim=1)
feat_volume = self.partfield_voxel_encoder(feat_volume)
feat_volume = rearrange(feat_volume, 'b c x y z -> b (x y z) c')
voxel_token = torch.cat([image_latents, feat_volume], dim=1) # B x N x D
input_ids = batch['input_ids']
inputs_embeds, attention_mask = self._prepare_lm_inputs(voxel_token, input_ids)
output = self.decoder(
attention_mask=attention_mask,
inputs_embeds=inputs_embeds,
return_dict=True,
)
return {
"logits": output.logits,
}
def gen_mesh_from_bounds(self, bounds, random_color):
bboxes = []
for j in range(bounds.shape[0]):
bbox = trimesh.primitives.Box(bounds=bounds[j])
color = random_color[j]
bbox.visual.vertex_colors = color
bboxes.append(bbox)
mesh = trimesh.Scene(bboxes)
return mesh
def generate(self, batch):
image_latents = self.image_encoder(batch['images'])
masks = batch['masks']
masks_emb = self.group_embedding(masks)
masks_emb = rearrange(masks_emb, 'b c h w -> b (h w) c') # B x Q x C
group_emb = torch.zeros((image_latents.shape[0], image_latents.shape[1], masks_emb.shape[2]), device=image_latents.device, dtype=image_latents.dtype)
group_emb[:, :masks_emb.shape[1], :] = masks_emb
image_latents = torch.cat([image_latents, group_emb], dim=-1)
image_latents = self.image_projector(image_latents)
points = batch['points'][..., :3]
rot_matrix = torch.tensor([[1, 0, 0], [0, 0, -1], [0, 1, 0]], device=points.device, dtype=points.dtype)
rot_points = torch.matmul(points, rot_matrix)
rot_points = rot_points * (2 * 0.9) # from (-0.5, 0.5) to (-1, 1)
partfield_feat = self.partfield_encoder.encode(rot_points)
feat_volume = torch.zeros((points.shape[0], 448, 64, 64, 64), device=partfield_feat.device, dtype=partfield_feat.dtype)
whole_voxel_index = batch['whole_voxel_index'] # (b, m, 3)
batch_size, num_points = whole_voxel_index.shape[0], whole_voxel_index.shape[1]
batch_indices = torch.arange(batch_size, device=whole_voxel_index.device).unsqueeze(1).expand(-1, num_points) # (b, m)
batch_flat = batch_indices.flatten() # (b*m,)
x_flat = whole_voxel_index[..., 0].flatten() # (b*m,)
y_flat = whole_voxel_index[..., 1].flatten() # (b*m,)
z_flat = whole_voxel_index[..., 2].flatten() # (b*m,)
partfield_feat_flat = partfield_feat.reshape(-1, 448) # (b*m, 448)
feat_volume[batch_flat, :, x_flat, y_flat, z_flat] = partfield_feat_flat
xyz_volume = torch.zeros((points.shape[0], 3, 64, 64, 64), device=points.device, dtype=points.dtype)
xyz_volume[batch_flat, :, x_flat, y_flat, z_flat] = points.reshape(-1, 3)
feat_volume = torch.cat([feat_volume, xyz_volume], dim=1)
feat_volume = self.partfield_voxel_encoder(feat_volume)
feat_volume = rearrange(feat_volume, 'b c x y z -> b (x y z) c')
voxel_token = torch.cat([image_latents, feat_volume], dim=1) # B x N x D
meshes = []
mesh_names = []
bboxes = []
output = self.decoder.generate(
inputs_embeds=voxel_token,
max_new_tokens=self.cfg.max_length - voxel_token.shape[1],
logits_processor=self.logits_processor,
do_sample=True,
top_k=5,
top_p=0.95,
temperature=0.5,
use_cache=True,
)
for i in range(output.shape[0]):
bounds = self.tokenizer.decode(output[i].detach().cpu().numpy(), coord_rg=(-0.5, 0.5))
# mesh = self.gen_mesh_from_bounds(bounds, batch['random_color'][i])
# meshes.append(mesh)
mesh_names.append("topk=5")
bboxes.append(bounds)
return {
# 'meshes': meshes,
'mesh_names': mesh_names,
'bboxes': bboxes,
}
@@ -0,0 +1,215 @@
import torch
from torch import nn
import torch.nn.functional as F
from diffusers.models.normalization import FP32LayerNorm
from diffusers.models.attention import FeedForward
from transformers.generation.logits_process import LogitsProcessor
from typing import List, Literal, Optional
from ..modules.norm import GroupNorm32, ChannelLayerNorm32
class GroupEmbedding(nn.Module):
def __init__(self, max_group_size, hidden_size=64):
super().__init__()
self.group_embedding = nn.Embedding(max_group_size + 1, hidden_size) # +1 for background
self.group_embedding.weight.data.normal_(mean=0.0, std=0.02)
def forward(self, masks):
batch_size, height, width = masks.shape
masks_flat = masks.reshape(batch_size, -1)
embeddings = self.group_embedding(masks_flat)
embeddings = embeddings.reshape(batch_size, height, width, -1)
embeddings = embeddings.permute(0, 3, 1, 2)
return embeddings
class MultiModalProjector(torch.nn.Module):
def __init__(self, in_features: int, out_features: int, pos_embed_seq_len=None):
super().__init__()
self.norm1 = FP32LayerNorm(in_features)
self.ff = FeedForward(in_features, out_features, mult=1, activation_fn="gelu")
self.norm2 = FP32LayerNorm(out_features)
if pos_embed_seq_len is not None:
self.pos_embed = nn.Parameter(torch.zeros(1, pos_embed_seq_len, in_features))
else:
self.pos_embed = None
def forward(self, encoder_hidden_states_image: torch.Tensor) -> torch.Tensor:
if self.pos_embed is not None:
batch_size, seq_len, embed_dim = encoder_hidden_states_image.shape
encoder_hidden_states_image = encoder_hidden_states_image.view(-1, 2 * seq_len, embed_dim)
encoder_hidden_states_image = encoder_hidden_states_image + self.pos_embed
hidden_states = self.norm1(encoder_hidden_states_image)
hidden_states = self.ff(hidden_states)
hidden_states = self.norm2(hidden_states)
return hidden_states
class MeshDecodeLogitsProcessor(LogitsProcessor):
def __init__(self, bins, BOS_id, EOS_id, PAD_id, vertices_num=8):
super().__init__()
self.bins = bins
self.BOS_id = BOS_id
self.EOS_id = EOS_id
self.PAD_id = PAD_id
self.filter_value = -float('inf')
self.vertices_num = vertices_num
def force_token(self, scores, token_id):
mask = torch.ones_like(scores, dtype=torch.bool)
mask[:, token_id] = False
scores[mask] = self.filter_value
def __call__(self, input_ids, scores):
# # all rules:
# # 1. first token: BOS
current_len = input_ids.shape[-1]
if current_len == 0:
# force bos
self.force_token(scores, self.BOS_id)
elif current_len <= self.vertices_num * 3 + 1:
scores[:, self.bins:] = self.filter_value
else:
scores[:, self.BOS_id] = self.filter_value
scores[:, self.PAD_id] = self.filter_value
effective_tokens = current_len - 1
complete_boxes = effective_tokens % (self.vertices_num * 3) == 0
# print(effective_tokens, complete_boxes)
if not complete_boxes:
scores[:, self.EOS_id] = self.filter_value
return scores
def norm_layer(norm_type: str, *args, **kwargs) -> nn.Module:
"""
Return a normalization layer.
"""
if norm_type == "group":
return GroupNorm32(32, *args, **kwargs)
elif norm_type == "layer":
return ChannelLayerNorm32(*args, **kwargs)
else:
raise ValueError(f"Invalid norm type {norm_type}")
class ResBlock3d(nn.Module):
def __init__(
self,
channels: int,
out_channels: Optional[int] = None,
norm_type: Literal["group", "layer"] = "layer",
):
super().__init__()
self.channels = channels
self.out_channels = out_channels or channels
self.norm1 = norm_layer(norm_type, channels)
self.norm2 = norm_layer(norm_type, self.out_channels)
self.conv1 = nn.Conv3d(channels, self.out_channels, 3, padding=1)
self.conv2 = zero_module(nn.Conv3d(self.out_channels, self.out_channels, 3, padding=1))
self.skip_connection = nn.Conv3d(channels, self.out_channels, 1) if channels != self.out_channels else nn.Identity()
def forward(self, x: torch.Tensor) -> torch.Tensor:
h = self.norm1(x)
h = F.silu(h)
h = self.conv1(h)
h = self.norm2(h)
h = F.silu(h)
h = self.conv2(h)
h = h + self.skip_connection(x)
return h
class DownsampleBlock3d(nn.Module):
def __init__(
self,
in_channels: int,
out_channels: int,
mode: Literal["conv", "avgpool"] = "conv",
):
assert mode in ["conv", "avgpool"], f"Invalid mode {mode}"
super().__init__()
self.in_channels = in_channels
self.out_channels = out_channels
if mode == "conv":
self.conv = nn.Conv3d(in_channels, out_channels, 2, stride=2)
elif mode == "avgpool":
assert in_channels == out_channels, "Pooling mode requires in_channels to be equal to out_channels"
def forward(self, x: torch.Tensor) -> torch.Tensor:
if hasattr(self, "conv"):
return self.conv(x)
else:
return F.avg_pool3d(x, 2)
def zero_module(module):
"""
Zero out the parameters of a module and return it.
"""
for p in module.parameters():
p.detach().zero_()
return module
class SparseStructureEncoder(nn.Module):
def __init__(
self,
in_channels: int,
latent_channels: int,
num_res_blocks: int,
channels: List[int],
num_res_blocks_middle: int = 2,
norm_type: Literal["group", "layer"] = "layer",
):
super().__init__()
self.in_channels = in_channels
self.latent_channels = latent_channels
self.num_res_blocks = num_res_blocks
self.channels = channels
self.num_res_blocks_middle = num_res_blocks_middle
self.norm_type = norm_type
self.dtype = torch.float16
self.input_layer = nn.Conv3d(in_channels, channels[0], 3, padding=1)
self.blocks = nn.ModuleList([])
for i, ch in enumerate(channels):
self.blocks.extend([
ResBlock3d(ch, ch)
for _ in range(num_res_blocks)
])
if i < len(channels) - 1:
self.blocks.append(
DownsampleBlock3d(ch, channels[i+1])
)
self.middle_block = nn.Sequential(*[
ResBlock3d(channels[-1], channels[-1])
for _ in range(num_res_blocks_middle)
])
@property
def device(self) -> torch.device:
"""
Return the device of the model.
"""
return next(self.parameters()).device
def forward(self, x: torch.Tensor):
h = self.input_layer(x)
h = h.type(self.dtype)
for block in self.blocks:
h = block(h)
h = self.middle_block(h)
h = h.type(x.dtype)
return h
+223
View File
@@ -0,0 +1,223 @@
import torch
import torch.utils.checkpoint
from torch import nn
from transformers import AutoModelForCausalLM, AutoConfig
from transformers.models.opt.modeling_opt import OPTForCausalLM, OPTModel, OPTDecoder, OPTConfig
from transformers.utils import logging
from typing import Optional, Union
from transformers import __version__ as transformers_version
from transformers.generation.logits_process import LogitsProcessorList
from transformers.generation.utils import GenerateNonBeamOutput, GenerateEncoderDecoderOutput, GenerateDecoderOnlyOutput
from transformers.generation.stopping_criteria import StoppingCriteriaList
from transformers.generation.configuration_utils import GenerationConfig
from transformers.generation.streamers import BaseStreamer
logger = logging.get_logger(__name__)
class BBoxOPTConfig(OPTConfig):
model_type = "mesh_opt"
class BBoxOPTDecoder(OPTDecoder):
config_class = BBoxOPTConfig
class BBoxOPTModel(OPTModel):
config_class = BBoxOPTConfig
def __init__(self, config: BBoxOPTConfig):
super(OPTModel, self).__init__(config)
self.decoder = BBoxOPTDecoder(config)
# Initialize weights and apply final processing
self.post_init()
class BBoxOPT(OPTForCausalLM):
config_class = BBoxOPTConfig
def __init__(self, config: BBoxOPTConfig):
super(OPTForCausalLM, self).__init__(config)
self.model = BBoxOPTModel(config)
# the lm_head weight is automatically tied to the embed tokens weight
self.lm_head = nn.Linear(config.word_embed_proj_dim, config.vocab_size, bias=False)
# Initialize weights and apply final processing
self.post_init()
def _sample(
self,
input_ids: torch.LongTensor,
logits_processor: LogitsProcessorList,
stopping_criteria: StoppingCriteriaList,
generation_config: GenerationConfig,
synced_gpus: bool,
streamer: Optional["BaseStreamer"],
**model_kwargs,
) -> Union[GenerateNonBeamOutput, torch.LongTensor]:
r"""
Generates sequences of token ids for models with a language modeling head using **multinomial sampling** and
can be used for text-decoder, text-to-text, speech-to-text, and vision-to-text models.
Parameters:
input_ids (`torch.LongTensor` of shape `(batch_size, sequence_length)`):
The sequence used as a prompt for the generation.
logits_processor (`LogitsProcessorList`):
An instance of [`LogitsProcessorList`]. List of instances of class derived from [`LogitsProcessor`]
used to modify the prediction scores of the language modeling head applied at each generation step.
stopping_criteria (`StoppingCriteriaList`):
An instance of [`StoppingCriteriaList`]. List of instances of class derived from [`StoppingCriteria`]
used to tell if the generation loop should stop.
generation_config ([`~generation.GenerationConfig`]):
The generation configuration to be used as parametrization of the decoding method.
synced_gpus (`bool`):
Whether to continue running the while loop until max_length (needed for ZeRO stage 3)
streamer (`BaseStreamer`, *optional*):
Streamer object that will be used to stream the generated sequences. Generated tokens are passed
through `streamer.put(token_ids)` and the streamer is responsible for any further processing.
model_kwargs:
Additional model specific kwargs will be forwarded to the `forward` function of the model. If model is
an encoder-decoder model the kwargs should include `encoder_outputs`.
Return:
[`~generation.GenerateDecoderOnlyOutput`], [`~generation.GenerateEncoderDecoderOutput`] or `torch.LongTensor`:
A `torch.LongTensor` containing the generated tokens (default behaviour) or a
[`~generation.GenerateDecoderOnlyOutput`] if `model.config.is_encoder_decoder=False` and
`return_dict_in_generate=True` or a [`~generation.GenerateEncoderDecoderOutput`] if
`model.config.is_encoder_decoder=True`.
"""
# init values
pad_token_id = generation_config._pad_token_tensor
output_attentions = generation_config.output_attentions
output_hidden_states = generation_config.output_hidden_states
output_scores = generation_config.output_scores
output_logits = generation_config.output_logits
return_dict_in_generate = generation_config.return_dict_in_generate
max_length = generation_config.max_length
has_eos_stopping_criteria = any(hasattr(criteria, "eos_token_id") for criteria in stopping_criteria)
do_sample = generation_config.do_sample
# init attention / hidden states / scores tuples
scores = () if (return_dict_in_generate and output_scores) else None
raw_logits = () if (return_dict_in_generate and output_logits) else None
decoder_attentions = () if (return_dict_in_generate and output_attentions) else None
cross_attentions = () if (return_dict_in_generate and output_attentions) else None
decoder_hidden_states = () if (return_dict_in_generate and output_hidden_states) else None
# if model is an encoder-decoder, retrieve encoder attention weights and hidden states
if return_dict_in_generate and self.config.is_encoder_decoder:
encoder_attentions = model_kwargs["encoder_outputs"].get("attentions") if output_attentions else None
encoder_hidden_states = (
model_kwargs["encoder_outputs"].get("hidden_states") if output_hidden_states else None
)
# keep track of which sequences are already finished
batch_size, cur_len = input_ids.shape[:2]
this_peer_finished = False
unfinished_sequences = torch.ones(batch_size, dtype=torch.long, device=input_ids.device)
if transformers_version=="4.50.3":
model_kwargs = self._get_initial_cache_position(input_ids, model_kwargs)
else:
model_kwargs = self._get_initial_cache_position(cur_len,input_ids.device, model_kwargs)
while self._has_unfinished_sequences(
this_peer_finished, synced_gpus, device=input_ids.device
) and cur_len < max_length:
# prepare model inputs
model_inputs = self.prepare_inputs_for_generation(input_ids, **model_kwargs)
# prepare variable output controls (note: some models won't accept all output controls)
model_inputs.update({"output_attentions": output_attentions} if output_attentions else {})
model_inputs.update({"output_hidden_states": output_hidden_states} if output_hidden_states else {})
# forward pass to get next token
outputs = self(**model_inputs, return_dict=True)
if synced_gpus and this_peer_finished:
continue # don't waste resources running the code we don't need
# Clone is needed to avoid keeping a hanging ref to outputs.logits which may be very large for first iteration
# (the clone itself is always small)
next_token_logits = outputs.logits.clone()[:, -1, :].float()
# pre-process distribution
next_token_scores = logits_processor(input_ids, next_token_logits)
# Store scores, attentions and hidden_states when required
if return_dict_in_generate:
if output_scores:
scores += (next_token_scores,)
if output_logits:
raw_logits += (next_token_logits,)
if output_attentions:
decoder_attentions += (
(outputs.decoder_attentions,) if self.config.is_encoder_decoder else (outputs.attentions,)
)
if self.config.is_encoder_decoder:
cross_attentions += (outputs.cross_attentions,)
if output_hidden_states:
decoder_hidden_states += (
(outputs.decoder_hidden_states,)
if self.config.is_encoder_decoder
else (outputs.hidden_states,)
)
# token selection
if do_sample:
probs = nn.functional.softmax(next_token_scores, dim=-1)
# TODO (joao): this OP throws "skipping cudagraphs due to ['incompatible ops']", find solution
next_tokens = torch.multinomial(probs, num_samples=1).squeeze(1)
else:
next_tokens = torch.argmax(next_token_scores, dim=-1)
# finished sentences should have their next token be a padding token
if has_eos_stopping_criteria:
next_tokens = next_tokens * unfinished_sequences + pad_token_id * (1 - unfinished_sequences)
# update generated ids, model inputs, and length for next step
input_ids = torch.cat([input_ids, next_tokens[:, None]], dim=-1)
if streamer is not None:
streamer.put(next_tokens.cpu())
model_kwargs = self._update_model_kwargs_for_generation(
outputs,
model_kwargs,
is_encoder_decoder=self.config.is_encoder_decoder,
)
unfinished_sequences = unfinished_sequences & ~stopping_criteria(input_ids, scores)
this_peer_finished = unfinished_sequences.max() == 0
cur_len += 1
# This is needed to properly delete outputs.logits which may be very large for first iteration
# Otherwise a reference to outputs is kept which keeps the logits alive in the next iteration
del outputs
if streamer is not None:
streamer.end()
if return_dict_in_generate:
if self.config.is_encoder_decoder:
return GenerateEncoderDecoderOutput(
sequences=input_ids,
scores=scores,
logits=raw_logits,
encoder_attentions=encoder_attentions,
encoder_hidden_states=encoder_hidden_states,
decoder_attentions=decoder_attentions,
cross_attentions=cross_attentions,
decoder_hidden_states=decoder_hidden_states,
past_key_values=model_kwargs.get("past_key_values"),
)
else:
return GenerateDecoderOnlyOutput(
sequences=input_ids,
scores=scores,
logits=raw_logits,
attentions=decoder_attentions,
hidden_states=decoder_hidden_states,
past_key_values=model_kwargs.get("past_key_values"),
)
else:
return input_ids
AutoConfig.register("mesh_opt", BBoxOPTConfig)
AutoModelForCausalLM.register(BBoxOPTConfig, BBoxOPT)
@@ -0,0 +1,44 @@
from typing import Literal
import numpy as np
import torch
import torch.nn as nn
import torch.nn.functional as F
from PIL import Image
from safetensors.torch import load_file
from transformers import AutoModel,AutoConfig
class DINOv2ImageEncoder(nn.Module):
def __init__(self, model_name: Literal[
"facebook/dinov2-with-registers-large",
"facebook/dinov2-large"
],ckpt=""):
super().__init__()
config = AutoConfig.from_pretrained(model_name)
self.model=AutoModel.from_config(config).to(dtype=torch.float16)
self.model.load_state_dict(torch.load(ckpt,map_location="cpu",weights_only=False) if not ckpt.endswith('.safetensors') else load_file(ckpt,device="cpu"))
#self.model = AutoModel.from_pretrained(model_name, torch_dtype=torch.bfloat16)
self.model.requires_grad_(False)
self.model.eval()
DINOv2_INPUT_MEAN = torch.as_tensor([0.485, 0.456, 0.406], dtype=torch.float32)[
None, :, None, None
]
DINOv2_INPUT_STD = torch.as_tensor([0.229, 0.224, 0.225], dtype=torch.float32)[
None, :, None, None
]
self.register_buffer("DINOv2_INPUT_MEAN", DINOv2_INPUT_MEAN, persistent=False)
self.register_buffer("DINOv2_INPUT_STD", DINOv2_INPUT_STD, persistent=False)
self.max_size = 518
self.hidden_size = self.model.config.hidden_size
def preprocess(self, image: torch.Tensor):
B, C, H, W = image.shape
assert C == 3 and H <= self.max_size and W <= self.max_size
image = (image - self.DINOv2_INPUT_MEAN.to(image)) / self.DINOv2_INPUT_STD.to(image)
return image
def forward(self, image: torch.Tensor):
image = self.preprocess(image)
features = self.model(image).last_hidden_state
return features
+34
View File
@@ -0,0 +1,34 @@
import torch
import torch.nn as nn
import torch.nn.functional as F
class LayerNorm32(nn.LayerNorm):
def forward(self, inputs: torch.Tensor) -> torch.Tensor:
origin_dtype = inputs.dtype
return F.layer_norm(
inputs.float(),
self.normalized_shape,
self.weight.float() if self.weight is not None else None,
self.bias.float() if self.bias is not None else None,
self.eps,
).to(origin_dtype)
class GroupNorm32(nn.GroupNorm):
"""
A GroupNorm layer that converts to float32 before the forward pass.
"""
def forward(self, x: torch.Tensor) -> torch.Tensor:
return super().forward(x.float()).type(x.dtype)
class ChannelLayerNorm32(LayerNorm32):
def forward(self, x: torch.Tensor) -> torch.Tensor:
# print(x.dtype)
DIM = x.dim()
x = x.permute(0, *range(2, DIM), 1).contiguous()
x = super().forward(x)
x = x.permute(0, DIM-1, *range(1, DIM-1)).contiguous()
return x
@@ -0,0 +1,64 @@
import numpy as np
from .mesh import change_pcd_range
class BoundsTokenizerDiag:
def __init__(self, bins, BOS_id, EOS_id, PAD_id):
self.bins = bins
self.BOS_id = BOS_id
self.EOS_id = EOS_id
self.PAD_id = PAD_id
def encode(self, data_dict, coord_rg=(-1,1)):
"""
Encode bounding boxes to token sequence
Args:
data_dict: dictionary containing bounding boxes
coord_rg: range of coordinate values
Returns:
token sequence
"""
bounds = data_dict["bounds"] # (s, 2, 3)
all_vertices = bounds.reshape(-1, 6)
all_vertices = change_pcd_range(all_vertices, from_rg=coord_rg, to_rg=(0.5/self.bins, 1-0.5/self.bins))
quantized_vertices = (all_vertices * self.bins).astype(np.int32)
tokens = []
tokens.append(self.BOS_id)
tokens.extend(quantized_vertices.flatten().tolist())
tokens.append(self.EOS_id)
tokens = np.array(tokens)
return tokens
def decode(self, tokens, coord_rg=(-1,1)):
"""
Decode token sequence back to bounding boxes
Args:
tokens: token sequence
Returns:
bounding box array [N, 2, 3]
"""
# Remove special tokens
valid_tokens = []
for t in tokens:
if t != self.BOS_id and t != self.EOS_id and t != self.PAD_id:
valid_tokens.append(t)
# Ensure correct number of tokens (2 vertices per box, 3 coordinates per vertex)
if len(valid_tokens) % (2 * 3) != 0:
raise ValueError(f"Invalid token count: {len(valid_tokens)}")
# Reshape to vertex coordinates
points = np.array(valid_tokens).reshape(-1, 2, 3)
# Convert quantized coordinates back to continuous values
points = points / self.bins
points = change_pcd_range(points, from_rg=(0.5/self.bins, 1-0.5/self.bins), to_rg=coord_rg)
return points
+42
View File
@@ -0,0 +1,42 @@
import numpy as np
import trimesh
import torch
def normalize_scene(scene, rg=(-0.5, 0.5)):
# put to [-0.5, 0.5]
whole_center = scene.bounding_box.centroid
scene.apply_translation(-whole_center)
whole_scale = max(scene.bounding_box.extents)
scene.apply_scale((rg[1]-rg[0]) / whole_scale)
return scene
def normalize_mesh(mesh, rg=(-1,1)):
# put to [-1, 1]
vmin = mesh.vertices.min(axis=0)
vmax = mesh.vertices.max(axis=0)
center = (vmin + vmax) / 2
scale = (vmax - vmin).max()
mesh.vertices = (mesh.vertices - center) / scale * (rg[1] - rg[0]) + (rg[0] + rg[1]) / 2
def change_mesh_range(mesh, from_rg=(-1,1), to_rg=(-1,1)):
mesh.vertices = (mesh.vertices - (from_rg[0] + from_rg[1]) / 2) / (from_rg[1] - from_rg[0]) * (to_rg[1] - to_rg[0]) + (to_rg[0] + to_rg[1]) / 2
return mesh
def change_pcd_range(pcd, from_rg=(-1,1), to_rg=(-1,1)):
pcd = (pcd - (from_rg[0] + from_rg[1]) / 2) / (from_rg[1] - from_rg[0]) * (to_rg[1] - to_rg[0]) + (to_rg[0] + to_rg[1]) / 2
return pcd
def quantize_vertices(v, bins):
return (v * bins).astype(np.int32)
def sample_points(mesh, n):
points, face_index = trimesh.sample.sample_surface(mesh, n)
normals = mesh.face_normals[face_index]
return points, normals
def clear_mesh(mesh):
mesh.update_faces(mesh.nondegenerate_faces(height=1.e-8))
mesh.remove_unreferenced_vertices()
mesh.merge_vertices(digits_vertex=0)
return mesh
+408
View File
@@ -0,0 +1,408 @@
import os
os.environ['OPENCV_IO_ENABLE_OPENEXR'] = '1'
import numpy as np
from typing import Optional
from PIL import Image, ImageDraw
import torchvision.transforms.functional as TF
import cv2
import torch
import trimesh
import glob
from tqdm import tqdm
def tensor_to_mask_array(mask_tensor):
if isinstance(mask_tensor, torch.Tensor):
mask_array = mask_tensor.cpu().numpy()
if mask_array.max() <= 1.0:
mask_array = (mask_array * 255).astype(np.uint8)
else:
mask_array = mask_array.astype(np.uint8)
return mask_array
def load_img_mask(img_path, mask_path, size=(518, 518)): #()
image = Image.open(img_path)
alpha = np.array(image.getchannel(3))
bbox = np.array(alpha).nonzero()
bbox = [bbox[1].min(), bbox[0].min(), bbox[1].max(), bbox[0].max()]
center = [(bbox[0] + bbox[2]) / 2, (bbox[1] + bbox[3]) / 2]
hsize = max(bbox[2] - bbox[0], bbox[3] - bbox[1]) / 2
aug_size_ratio = 1.2
aug_hsize = hsize * aug_size_ratio
aug_center_offset = [0, 0]
aug_center = [center[0] + aug_center_offset[0], center[1] + aug_center_offset[1]]
aug_bbox = [int(aug_center[0] - aug_hsize), int(aug_center[1] - aug_hsize), int(aug_center[0] + aug_hsize), int(aug_center[1] + aug_hsize)]
img_height, img_width = alpha.shape
if mask_path.endswith('.npy'):
mask = np.load(mask_path)
elif mask_path.endswith('.exr'):
mask = cv2.imread(mask_path, cv2.IMREAD_UNCHANGED)
else:
mask = cv2.imread(mask_path, cv2.IMREAD_UNCHANGED)
if mask.dtype == np.uint8 or mask.dtype == np.uint16:
mask = mask.astype(np.float32)
#mask = cv2.imread(mask_path, cv2.IMREAD_UNCHANGED)
pad_left = max(0, -aug_bbox[0])
pad_top = max(0, -aug_bbox[1])
pad_right = max(0, aug_bbox[2] - img_width)
pad_bottom = max(0, aug_bbox[3] - img_height)
if pad_left > 0 or pad_top > 0 or pad_right > 0 or pad_bottom > 0:
img_array = np.array(image)
padded_img_array = np.pad(
img_array,
((pad_top, pad_bottom), (pad_left, pad_right), (0, 0)),
mode='constant',
constant_values=0
)
padded_mask_array = np.pad(mask, ((pad_top, pad_bottom), (pad_left, pad_right), (0, 0)), mode='constant', constant_values=0)
image = Image.fromarray(padded_img_array.astype('uint8'))
aug_bbox[0] += pad_left
aug_bbox[1] += pad_top
aug_bbox[2] += pad_left
aug_bbox[3] += pad_top
mask = padded_mask_array
image = image.crop(aug_bbox)
mask = mask[aug_bbox[1]:aug_bbox[3], aug_bbox[0]:aug_bbox[2]]
ordered_mask_input, mask_vis = load_bottom_up_mask(mask)
image_white_bg = np.array(image)
image_black_bg = np.array(image)
if image_white_bg.shape[-1] == 4:
mask_img = image_white_bg[..., 3] == 0
image_white_bg[mask_img] = [255, 255, 255, 255]
image_black_bg[mask_img] = [0, 0, 0, 255]
image_white_bg = image_white_bg[..., :3]
image_black_bg = image_black_bg[..., :3]
img_white_bg = Image.fromarray(image_white_bg.astype('uint8'))
img_black_bg = Image.fromarray(image_black_bg.astype('uint8'))
img_white_bg = img_white_bg.resize(size, resample=Image.Resampling.LANCZOS)
img_black_bg = img_black_bg.resize(size, resample=Image.Resampling.LANCZOS)
img_mask_vis = vis_mask_on_img(img_white_bg, mask_vis)
img_white_bg = TF.to_tensor(img_white_bg)
img_black_bg = TF.to_tensor(img_black_bg)
return img_white_bg, img_black_bg, ordered_mask_input, img_mask_vis
def load_bottom_up_mask(mask, size=(518, 518)):
mask_input = smart_downsample_mask(mask, (37, 37))
mask_vis = cv2.resize(mask_input, (518, 518), interpolation=cv2.INTER_NEAREST)
mask_input = np.array(mask_input, dtype=np.int32)
unique_indices = np.unique(mask_input)
unique_indices = unique_indices[unique_indices > 0]
part_positions = {}
for idx in unique_indices:
y_coords, _ = np.where(mask_input == idx)
if len(y_coords) > 0:
part_positions[idx] = np.max(y_coords)
sorted_parts = sorted(part_positions.items(), key=lambda x: -x[1]) # Sort by y-coordinate in descending order
# Create mapping from old indices to new indices (ordered by position)
index_map = {}
for new_idx, (old_idx, _) in enumerate(sorted_parts, 1): # Start from 1 (0 is background)
index_map[old_idx] = new_idx
# Apply the mapping to create position-ordered mask
ordered_mask_input = np.zeros_like(mask_input)
for old_idx, new_idx in index_map.items():
ordered_mask_input[mask_input == old_idx] = new_idx
mask_vis = np.array(mask_vis, dtype=np.int32)
ordered_mask_input = torch.from_numpy(ordered_mask_input).long()
return ordered_mask_input, mask_vis
def smart_downsample_mask(mask, target_size):
h, w = mask.shape[:2]
target_h, target_w = target_size
h_ratio = h / target_h
w_ratio = w / target_w
downsampled = np.zeros((target_h, target_w), dtype=mask.dtype)
for i in range(target_h):
for j in range(target_w):
y_start = int(i * h_ratio)
y_end = min(int((i + 1) * h_ratio), h)
x_start = int(j * w_ratio)
x_end = min(int((j + 1) * w_ratio), w)
region = mask[y_start:y_end, x_start:x_end]
if region.size == 0:
continue
unique_values, counts = np.unique(region.flatten(), return_counts=True)
non_zero_mask = unique_values > 0
if np.any(non_zero_mask):
non_zero_values = unique_values[non_zero_mask]
non_zero_counts = counts[non_zero_mask]
max_idx = np.argmax(non_zero_counts)
downsampled[i, j] = non_zero_values[max_idx]
else:
max_idx = np.argmax(counts)
downsampled[i, j] = unique_values[max_idx]
return downsampled
def vis_mask_on_img(img, mask):
H, W = mask.shape
mask_vis = np.zeros((H, W, 3), dtype=np.uint8) + 255
for part_id in range(1, int(mask.max()) + 1):
part_mask = (mask == part_id)
if part_mask.sum() > 0:
color = get_random_color((part_id - 1), use_float=False)[:3]
mask_vis[part_mask, 0:3] = color
mask_img = Image.fromarray(mask_vis)
combined_width = W * 2
combined_height = H
combined_img = Image.new('RGB', (combined_width, combined_height), (255, 255, 255))
combined_img.paste(img, (0, 0))
combined_img.paste(mask_img, (W, 0))
draw = ImageDraw.Draw(combined_img)
draw.line([(W, 0), (W, H)], fill=(0, 0, 0), width=2)
return combined_img
def get_random_color(index: Optional[int] = None, use_float: bool = False):
# some pleasing colors
# matplotlib.colormaps['Set3'].colors + matplotlib.colormaps['Set2'].colors + matplotlib.colormaps['Set1'].colors
palette = np.array(
[
[141, 211, 199, 255],
[255, 255, 179, 255],
[190, 186, 218, 255],
[251, 128, 114, 255],
[128, 177, 211, 255],
[253, 180, 98, 255],
[179, 222, 105, 255],
[252, 205, 229, 255],
[217, 217, 217, 255],
[188, 128, 189, 255],
[204, 235, 197, 255],
[255, 237, 111, 255],
[102, 194, 165, 255],
[252, 141, 98, 255],
[141, 160, 203, 255],
[231, 138, 195, 255],
[166, 216, 84, 255],
[255, 217, 47, 255],
[229, 196, 148, 255],
[179, 179, 179, 255],
[228, 26, 28, 255],
[55, 126, 184, 255],
[77, 175, 74, 255],
[152, 78, 163, 255],
[255, 127, 0, 255],
[255, 255, 51, 255],
[166, 86, 40, 255],
[247, 129, 191, 255],
[153, 153, 153, 255],
],
dtype=np.uint8,
)
if index is None:
index = np.random.randint(0, len(palette))
if index >= len(palette):
index = index % len(palette)
if use_float:
return palette[index].astype(np.float32) / 255
else:
return palette[index]
def change_pcd_range(pcd, from_rg=(-1,1), to_rg=(-1,1)):
pcd = (pcd - (from_rg[0] + from_rg[1]) / 2) / (from_rg[1] - from_rg[0]) * (to_rg[1] - to_rg[0]) + (to_rg[0] + to_rg[1]) / 2
return pcd
def prepare_bbox_gen_input(voxel_coords_path, img_white_bg, ordered_mask_input, bins=64, device="cuda"):
whole_voxel = np.load(voxel_coords_path)
whole_voxel = whole_voxel[:, 1:]
whole_voxel = (whole_voxel + 0.5) / bins - 0.5
whole_voxel_index = change_pcd_range(whole_voxel, from_rg=(-0.5, 0.5), to_rg=(0.5/bins, 1-0.5/bins))
whole_voxel_index = (whole_voxel_index * bins).astype(np.int32)
points = torch.from_numpy(whole_voxel).to(torch.float16).unsqueeze(0).to(device)
whole_voxel_index = torch.from_numpy(whole_voxel_index).long().unsqueeze(0).to(device)
images = img_white_bg.unsqueeze(0).to(device)
masks = ordered_mask_input.unsqueeze(0).to(device)
return {
"points": points,
"whole_voxel_index": whole_voxel_index,
"images": images,
"masks": masks,
}
def vis_voxel_coords(voxel_coords, bins=64):
voxel_coords = voxel_coords[:, 1:]
voxel_coords = (voxel_coords + 0.5) / bins - 0.5
voxel_coords_ply = trimesh.PointCloud(voxel_coords)
rot_matrix = np.array([[1, 0, 0, 0], [0, 0, 1, 0], [0, -1, 0, 0], [0, 0, 0, 1]])
voxel_coords_ply.apply_transform(rot_matrix)
return voxel_coords_ply
def gen_mesh_from_bounds(bounds):
bboxes = []
rot_matrix = np.array([[1, 0, 0, 0], [0, 0, 1, 0], [0, -1, 0, 0], [0, 0, 0, 1]])
for j in range(bounds.shape[0]):
bbox = trimesh.primitives.Box(bounds=bounds[j])
color = get_random_color(j, use_float=True)
bbox.visual.vertex_colors = color
bboxes.append(bbox)
mesh = trimesh.Scene(bboxes)
mesh.apply_transform(rot_matrix)
return mesh
def prepare_part_synthesis_input(voxel_coords_path, bbox_depth_path, ordered_mask_input, padding_size=2, bins=64, device="cuda"):
overall_coords = np.load(voxel_coords_path)
overall_coords = overall_coords[:, 1:] # Remove first column
bbox_scene = np.load(bbox_depth_path)
all_coords_wnoise = []
part_layouts = []
start_idx = 0
part_layouts.append(slice(start_idx, start_idx + overall_coords.shape[0]))
start_idx += overall_coords.shape[0]
assigned_points = np.zeros(overall_coords.shape[0], dtype=bool)
bbox_coords_list = []
bbox_masks = []
for bbox in bbox_scene:
points = change_pcd_range(bbox, from_rg=(-0.5, 0.5), to_rg=(0.5/bins, 1-0.5/bins))
bbox_min = np.floor(points[0] * bins).astype(np.int32)
bbox_max = np.ceil(points[1] * bins).astype(np.int32)
bbox_min = np.clip(bbox_min - padding_size, 0, bins - 1)
bbox_max = np.clip(bbox_max + padding_size, 0, bins - 1)
bbox_mask = np.all((overall_coords >= bbox_min) & (overall_coords <= bbox_max), axis=1)
bbox_masks.append(bbox_mask)
if np.sum(bbox_mask) == 0:
continue
assigned_points = assigned_points | bbox_mask
bbox_coords = overall_coords[bbox_mask]
bbox_coords_list.append(bbox_coords)
part_layouts.append(slice(start_idx, start_idx + bbox_coords.shape[0]))
start_idx += bbox_coords.shape[0]
bbox_coords = torch.from_numpy(bbox_coords)
all_coords_wnoise.append(bbox_coords)
unassigned_mask = ~assigned_points
unassigned_coords = overall_coords[unassigned_mask]
if np.sum(unassigned_mask) > 0 and len(bbox_scene) > 0:
print(f"Assigning {np.sum(unassigned_mask)} unassigned points to nearest bboxes")
nearest_bbox_indices = []
for point_idx, point in enumerate(unassigned_coords):
min_dist = float('inf')
nearest_idx = -1
for bbox_idx, bbox in enumerate(bbox_scene):
points = change_pcd_range(bbox, from_rg=(-0.5, 0.5), to_rg=(0.5/bins, 1-0.5/bins))
bbox_min = np.floor(points[0] * bins).astype(np.int32)
bbox_max = np.ceil(points[1] * bins).astype(np.int32)
dx = min(abs(point[0] - bbox_min[0]), abs(point[0] - bbox_max[0]))
dy = min(abs(point[1] - bbox_min[1]), abs(point[1] - bbox_max[1]))
dz = min(abs(point[2] - bbox_min[2]), abs(point[2] - bbox_max[2]))
# dist = dx + dy + dz
dist = min(dx, dy, dz)
if dist < min_dist:
min_dist = dist;
nearest_idx = bbox_idx
nearest_bbox_indices.append(nearest_idx)
for bbox_idx in range(len(bbox_scene)):
points_for_this_bbox = np.array([i for i, idx in enumerate(nearest_bbox_indices) if idx == bbox_idx])
if len(points_for_this_bbox) > 0:
additional_coords = unassigned_coords[points_for_this_bbox]
if bbox_idx < len(bbox_coords_list):
combined_coords = np.vstack([bbox_coords_list[bbox_idx], additional_coords])
old_slice = part_layouts[bbox_idx + 1] # +1 because first slice is whole model
new_slice = slice(old_slice.start, old_slice.start + combined_coords.shape[0])
part_layouts[bbox_idx + 1] = new_slice
additional_points = additional_coords.shape[0]
for i in range(bbox_idx + 2, len(part_layouts)):
old_slice = part_layouts[i]
new_slice = slice(old_slice.start + additional_points, old_slice.stop + additional_points)
part_layouts[i] = new_slice
all_coords_wnoise[bbox_idx] = torch.from_numpy(combined_coords)
start_idx += additional_points
else:
part_layouts.append(slice(start_idx, start_idx + additional_coords.shape[0]))
start_idx += additional_coords.shape[0]
all_coords_wnoise.append(torch.from_numpy(additional_coords))
overall_coords = torch.from_numpy(overall_coords)
all_coords_wnoise.insert(0, overall_coords)
combined_coords = torch.cat(all_coords_wnoise, dim=0).int()
coords = torch.cat(
[torch.full((combined_coords.shape[0], 1), 0, dtype=torch.int32), combined_coords],
dim=-1
).to(device)
masks = ordered_mask_input.unsqueeze(0).to(device)
return {
'coords': coords,
'part_layouts': part_layouts,
'masks': masks,
}
def merge_parts(save_dir,prefix):
scene_list = []
scene_list_texture = []
part_list = glob.glob(os.path.join(save_dir, "*.glb"))
part_list = [p for p in part_list if "part" in p and "parts" not in p and "part0" not in p] # part 0 is the overall model
part_list.sort()
for i, part_path in enumerate(tqdm(part_list, desc="Merging parts")):
part_mesh = trimesh.load(part_path, force='mesh')
scene_list_texture.append(part_mesh)
random_color = get_random_color(i, use_float=True)
part_mesh_color = part_mesh.copy()
part_mesh_color.visual = trimesh.visual.ColorVisuals(
mesh=part_mesh_color,
vertex_colors=random_color
)
scene_list.append(part_mesh_color)
os.remove(part_path)
scene_texture = trimesh.Scene(scene_list_texture)
mesh_texture_path=os.path.join(save_dir, f"{prefix}_mesh_textured.glb")
scene_texture.export(mesh_texture_path)
scene = trimesh.Scene(scene_list)
mesh_segment_path=os.path.join(save_dir, f"{prefix}_mesh_segment.glb")
scene.export(mesh_segment_path)
return mesh_texture_path,mesh_segment_path
@@ -0,0 +1,575 @@
"""
Image Part Segmentation and Labeling Tool
This script segments images into meaningful parts using the Segment Anything Model (SAM)
and optionally removes backgrounds using BriaRMBG. It identifies, visualizes, and merges
different parts of objects in images.
Key features:
- Background removal with alpha channel preservation
- Automatic part segmentation with SAM
- Intelligent part merging for logical grouping
- Detection of parts that SAM might miss
- Splitting of disconnected parts into separate components
- Edge cleaning and smoothing of segmentations
- Visualization of segmented parts with clear labeling
"""
import os
import argparse
import numpy as np
import cv2
import torch
from PIL import Image
from torchvision.transforms import functional as F
from torchvision import transforms
import torch.nn.functional as F_nn
from segment_anything import SamAutomaticMaskGenerator, build_sam
from .visualizer import Visualizer
# Minimum size threshold for considering a segment (in pixels)
size_th = 2000
def get_mask(group_ids, image, ids=None, img_name=None, save_dir=None):
"""
Creates and saves a colored visualization of mask segments.
Args:
group_ids: Array of segment IDs for each pixel
image: Input image
ids: Identifier to append to output filename
img_name: Base name of the image for saving
Returns:
Array of segment IDs (unchanged, just for convenience)
"""
colored_mask = np.zeros((image.shape[0], image.shape[1], 3), dtype=np.uint8)
colored_mask[group_ids == -1] = [255, 255, 255]
unique_ids = np.unique(group_ids)
unique_ids = unique_ids[unique_ids >= 0]
for i, unique_id in enumerate(unique_ids):
color_r = (i * 50 + 80) % 256
color_g = (i * 120 + 40) % 256
color_b = (i * 180 + 20) % 256
mask = (group_ids == unique_id)
colored_mask[mask] = [color_r, color_g, color_b]
mask_path = os.path.join(save_dir, f"{img_name}_mask_segments_{ids}.png")
cv2.imwrite(mask_path, cv2.cvtColor(colored_mask, cv2.COLOR_RGB2BGR))
print(f"Saved mask segments visualization to {mask_path}")
return group_ids
def clean_segment_edges(group_ids):
"""
Clean up segment edges by applying morphological operations to each segment.
Args:
group_ids: Array of segment IDs for each pixel
Returns:
Cleaned array of segment IDs with smoother boundaries
"""
# Get unique segment IDs (excluding background -1)
unique_ids = np.unique(group_ids)
unique_ids = unique_ids[unique_ids >= 0]
# Create a clean group_ids array
cleaned_group_ids = np.full_like(group_ids, -1) # Start with all background
# Define kernel for morphological operations
kernel = np.ones((3, 3), np.uint8)
# Process each segment individually
for segment_id in unique_ids:
# Extract the mask for this segment
segment_mask = (group_ids == segment_id).astype(np.uint8)
# Apply morphological closing to smooth edges
smoothed_mask = cv2.morphologyEx(segment_mask, cv2.MORPH_CLOSE, kernel, iterations=1)
# Apply morphological opening to remove small isolated pixels
smoothed_mask = cv2.morphologyEx(smoothed_mask, cv2.MORPH_OPEN, kernel, iterations=1)
# Add this segment back to the cleaned result
cleaned_group_ids[smoothed_mask > 0] = segment_id
print(f"Cleaned edges for {len(unique_ids)} segments")
return cleaned_group_ids
def prepare_image(image, bg_color=None, rmbg_net=None):
image_size = (1024, 1024)
transform_image = transforms.Compose([
transforms.Resize(image_size),
transforms.ToTensor(),
transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225])
])
input_images = transform_image(image).unsqueeze(0).to('cuda')
# Prediction
with torch.no_grad():
preds = rmbg_net(input_images)[-1].sigmoid().cpu()
pred = preds[0].squeeze()
pred_pil = transforms.ToPILImage()(pred)
mask = pred_pil.resize(image.size)
image.putalpha(mask)
return image
def resize_and_pad_to_square(image, target_size=518):
"""
Resize image to have longest side equal to target_size and pad shorter side
to create a square image.
Args:
image: PIL image or numpy array
target_size: Target square size, defaults to 518
Returns:
PIL Image resized and padded to square (target_size x target_size)
"""
# Ensure image is a PIL Image object
if isinstance(image, np.ndarray):
image = Image.fromarray(cv2.cvtColor(image, cv2.COLOR_BGR2RGB) if len(image.shape) == 3 and image.shape[2] == 3 else image)
# Get original dimensions
width, height = image.size
# Determine which dimension is longer
if width > height:
# Width is longer
new_width = target_size
new_height = int(height * (target_size / width))
else:
# Height is longer
new_height = target_size
new_width = int(width * (target_size / height))
# Resize image while maintaining aspect ratio
resized_image = image.resize((new_width, new_height), Image.LANCZOS)
# Create new square image with proper mode (with or without alpha channel)
mode = "RGBA" if image.mode == "RGBA" else "RGB"
background_color = (255, 255, 255, 0) if mode == "RGBA" else (255, 255, 255)
square_image = Image.new(mode, (target_size, target_size), background_color)
# Calculate position to paste resized image (centered)
paste_x = (target_size - new_width) // 2
paste_y = (target_size - new_height) // 2
# Paste resized image onto square background
if mode == "RGBA":
square_image.paste(resized_image, (paste_x, paste_y), resized_image)
else:
square_image.paste(resized_image, (paste_x, paste_y))
return square_image
def split_disconnected_parts(group_ids, size_threshold=None):
"""
Split each part into separate parts if they contain disconnected regions.
Args:
group_ids: Array of segment IDs for each pixel
size_threshold: Minimum size threshold for considering a segment (in pixels).
If None, uses the global size_th variable.
Returns:
Updated array with each connected component having a unique ID
"""
# Use provided threshold or fall back to global variable
if size_threshold is None:
size_threshold = size_th
# Create a copy to hold the result
new_group_ids = np.full_like(group_ids, -1) # Start with all background
# Get unique part IDs (excluding background -1)
unique_ids = np.unique(group_ids)
unique_ids = unique_ids[unique_ids >= 0]
# Track the next available ID
next_id = 0
total_split_regions = 0
# For each existing part ID
for part_id in unique_ids:
# Extract the mask for this part
part_mask = (group_ids == part_id).astype(np.uint8)
# Find connected components within this part
num_labels, labels = cv2.connectedComponents(part_mask, connectivity=8)
if num_labels == 1: # Just background (0), no regions found
continue
if num_labels == 2: # One connected component (background + 1 region)
# Assign the original part's area to the next available ID
new_group_ids[labels == 1] = next_id
next_id += 1
else: # Multiple disconnected components
split_count = 0
print(f"Part {part_id} has {num_labels-1} disconnected regions, splitting...")
# For each connected component (skipping background label 0)
for label in range(1, num_labels):
region_mask = labels == label
region_size = np.sum(region_mask)
# Only include regions that are large enough
if region_size >= size_threshold / 5: # Using size threshold to avoid tiny fragments
new_group_ids[region_mask] = next_id
split_count += 1
next_id += 1
else:
print(f" Skipping small disconnected region ({region_size} pixels)")
total_split_regions += split_count
if total_split_regions > 0:
print(f"Split disconnected parts: original {len(unique_ids)} parts -> {next_id} connected parts")
else:
print("No parts needed splitting - all parts are already connected")
return new_group_ids
# -------------------------------------------------------
# MAIN SEGMENTATION FUNCTION
# -------------------------------------------------------
def get_sam_mask(image, mask_generator, visual, merge_groups=None, existing_group_ids=None,
check_undetected=True, rgba_image=None, img_name=None, skip_split=False, save_dir=None, size_threshold=None):
"""
Generate and process SAM masks for the image, with optional merging and undetected region detection.
Args:
size_threshold: Minimum size threshold for considering a segment (in pixels).
If None, uses the global size_th variable.
"""
# Use provided threshold or fall back to global variable
if size_threshold is None:
size_threshold = size_th
label_mode = '1'
anno_mode = ['Mask', 'Mark']
exist_group = False
# Use existing group IDs if provided, otherwise generate new ones with SAM
if existing_group_ids is not None:
group_ids = existing_group_ids.copy()
group_counter = np.max(group_ids) + 1
exist_group = True
else:
# Generate masks using SAM
masks = mask_generator.generate(image)
group_ids = np.full((image.shape[0], image.shape[1]), -1, dtype=int)
num_masks = len(masks)
group_counter = 0
# Sort masks by area (largest first)
area_sorted_masks = sorted(masks, key=lambda x: x["area"], reverse=True)
# Create background mask if we have RGBA image
background_mask = None
if rgba_image is not None:
rgba_array = np.array(rgba_image)
if rgba_array.shape[2] == 4:
# Use alpha channel to create foreground/background mask
background_mask = rgba_array[:, :, 3] <= 10 # Areas with very low alpha are background
# First pass: assign original group IDs
for i in range(0, num_masks):
if area_sorted_masks[i]["area"] < size_threshold:
print(f"Skipping mask {i}, area too small: {area_sorted_masks[i]['area']} < {size_threshold}")
continue
mask = area_sorted_masks[i]["segmentation"]
# Check proportion of background pixels in this mask
if background_mask is not None:
# Calculate how many pixels in this mask are background
background_pixels_in_mask = np.sum(mask & background_mask)
mask_area = np.sum(mask)
background_ratio = background_pixels_in_mask / mask_area
# Skip mask if background proportion is too high (>10%)
if background_ratio > 0.1:
print(f" Skipping mask {i}, background ratio: {background_ratio:.2f}")
continue
# Assign group ID to this mask's pixels
group_ids[mask] = group_counter
print(f"Assigned mask {i} with area {area_sorted_masks[i]['area']} to group {group_counter}")
group_counter += 1
# Split disconnected parts immediately after SAM segmentation
print("Splitting disconnected parts in initial segmentation...")
group_ids = split_disconnected_parts(group_ids, size_threshold)
# Update group counter after splitting
if np.max(group_ids) >= 0:
group_counter = np.max(group_ids) + 1
print(f"After early splitting, now have {len(np.unique(group_ids))-1} regions (excluding background)")
# Check for undetected parts using RGBA information
if check_undetected and rgba_image is not None:
print("Checking for undetected parts using RGBA image...")
# Create a foreground mask from the alpha channel
rgba_array = np.array(rgba_image)
# Check if the image has an alpha channel
if rgba_array.shape[2] == 4:
print(f"Image has alpha channel, checking for undetected parts...")
# Use alpha channel to identify non-transparent pixels (foreground)
alpha_mask = rgba_array[:, :, 3] > 0
# Create existing parts mask and dilate it
existing_parts_mask = (group_ids != -1)
kernel = np.ones((4, 4), np.uint8)
# Use larger kernel for faster dilation
large_kernel = np.ones((4, 4), np.uint8)
dilated_parts = cv2.dilate(existing_parts_mask.astype(np.uint8), large_kernel)
# Find undetected areas (foreground but not detected by SAM)
undetected_mask = alpha_mask & (~dilated_parts.astype(bool))
# Process only if there are enough undetected pixels
if np.sum(undetected_mask) > size_threshold:
print(f"Found undetected parts with {np.sum(undetected_mask)} pixels")
# Find connected components in undetected regions
num_labels, labels = cv2.connectedComponents(
undetected_mask.astype(np.uint8),
connectivity=8
)
print(f" Found {num_labels-1} initial regions")
# Use Union-Find data structure for efficient region merging
parent = list(range(num_labels))
# Find with path compression
def find(x):
"""Find with path compression for Union-Find"""
if parent[x] != x:
parent[x] = find(parent[x])
return parent[x]
# Union by rank/size
def union(x, y):
"""Union operation for Union-Find"""
root_x = find(x)
root_y = find(y)
if root_x != root_y:
# Use smaller ID as parent
if root_x < root_y:
parent[root_y] = root_x
else:
parent[root_x] = root_y
# Calculate areas for all regions at once
areas = np.bincount(labels.flatten())[1:] if num_labels > 1 else []
# Filter regions by minimum size
valid_regions = np.where(areas >= size_threshold/5)[0] + 1
# Barrier mask for connectivity checks
barrier_mask = existing_parts_mask
# Pre-compute dilated regions for all valid regions
dilated_regions = {}
for i in valid_regions:
region_mask = (labels == i).astype(np.uint8)
dilated_regions[i] = cv2.dilate(region_mask, kernel, iterations=2)
# Check for region merges based on proximity and overlap
for idx, i in enumerate(valid_regions[:-1]):
for j in valid_regions[idx+1:]:
# Check overlap between dilated regions
overlap = dilated_regions[i] & dilated_regions[j]
overlap_size = np.sum(overlap)
# Merge if significant overlap and not separated by existing parts
if overlap_size > 40 and not np.any(overlap & barrier_mask):
# Calculate overlap ratios
overlap_ratio_i = overlap_size / areas[i-1]
overlap_ratio_j = overlap_size / areas[j-1]
if max(overlap_ratio_i, overlap_ratio_j) > 0.03:
union(i, j)
print(f" Merging regions {i} and {j} (overlap: {overlap_size} px)")
# Apply the merging results to create merged labels
merged_labels = np.zeros_like(labels)
for label in range(1, num_labels):
merged_labels[labels == label] = find(label)
# Get unique merged regions
unique_merged_regions = np.unique(merged_labels[merged_labels > 0])
print(f" After merging: {len(unique_merged_regions)} connected regions")
# Add regions to group_ids if they're large enough
group_counter_start = group_counter
for label in unique_merged_regions:
region_mask = merged_labels == label
region_size = np.sum(region_mask)
if region_size > size_threshold:
print(f" Adding region with ID {label} ({region_size} pixels) as group {group_counter}")
group_ids[region_mask] = group_counter
group_counter += 1
else:
print(f" Skipping small region with ID {label} ({region_size} pixels < {size_threshold})")
print(f" Added {group_counter - group_counter_start} regions that weren't detected by SAM")
# Process edges for all new parts at once
if group_counter > group_counter_start:
print("Processing edges for newly detected parts...")
# Create combined mask for all new parts
new_parts_mask = np.zeros_like(group_ids, dtype=bool)
for part_id in range(group_counter_start, group_counter):
new_parts_mask |= (group_ids == part_id)
# Compute edges for all new parts at once
all_new_dilated = cv2.dilate(new_parts_mask.astype(np.uint8), kernel, iterations=1)
all_new_eroded = cv2.erode(new_parts_mask.astype(np.uint8), kernel, iterations=1)
all_new_edges = all_new_dilated.astype(bool) & (~all_new_eroded.astype(bool))
print(f"Edge processing completed for {group_counter - group_counter_start} new parts")
# Save debug visualization of initial segmentation
if not exist_group:
get_mask(group_ids, image, ids=2, img_name=img_name, save_dir=save_dir)
# Merge groups if specified
if merge_groups is not None:
# Start with current group_ids
merged_group_ids = group_ids
# Preserve background regions
merged_group_ids[group_ids == -1] = -1
# For each merge group, assign all pixels to the first ID in that group
for new_id, group in enumerate(merge_groups):
# Create a mask to include all original IDs in this group
group_mask = np.zeros_like(group_ids, dtype=bool)
orig_ids_first = group[0]
# Process each original ID
for orig_id in group:
# Get mask for this original ID
mask = (group_ids == orig_id)
pixels = np.sum(mask)
if pixels > 0:
print(f" Including original ID {orig_id} ({pixels} pixels)")
group_mask = group_mask | mask
else:
print(f" Warning: Original ID {orig_id} does not exist")
# Set all pixels in this group to the first ID in the group
if np.any(group_mask):
print(f" Merging {np.sum(group_mask)} pixels to ID {orig_ids_first}")
merged_group_ids[group_mask] = orig_ids_first
# Reassign IDs to be continuous from 0
unique_ids = np.unique(merged_group_ids)
unique_ids = unique_ids[unique_ids != -1] # Exclude background
id_reassignment = {old_id: new_id for new_id, old_id in enumerate(unique_ids)}
# Create new array with reassigned IDs
new_group_ids = np.full_like(merged_group_ids, -1) # Start with all background
for old_id, new_id in id_reassignment.items():
new_group_ids[merged_group_ids == old_id] = new_id
# Update merged_group_ids with continuous IDs
merged_group_ids = new_group_ids
print(f"ID reassignment complete: {len(id_reassignment)} groups now have sequential IDs from 0 to {len(id_reassignment)-1}")
# Replace original group IDs with merged result
group_ids = merged_group_ids
print(f"Merging complete, now have {len(np.unique(group_ids))-1} regions (excluding background)")
# Skip splitting disconnected parts if requested
if not skip_split:
# Split disconnected parts into separate parts
group_ids = split_disconnected_parts(group_ids, size_threshold)
print(f"After splitting disconnected parts, now have {len(np.unique(group_ids))-1} regions (excluding background)")
else:
# Always split disconnected parts for initial segmentation
group_ids = split_disconnected_parts(group_ids, size_threshold)
print(f"After splitting disconnected parts, now have {len(np.unique(group_ids))-1} regions (excluding background)")
# Create visualization with clear labeling
vis_mask = visual
# First draw background areas (ID -1)
background_mask = (group_ids == -1)
if np.any(background_mask):
vis_mask = visual.draw_binary_mask(background_mask, color=[1.0, 1.0, 1.0], alpha=0.0)
# Then draw each segment with unique colors and labels
for unique_id in np.unique(group_ids):
if unique_id == -1: # Skip background
continue
mask = (group_ids == unique_id)
# Calculate center point and area of this region
y_indices, x_indices = np.where(mask)
if len(y_indices) > 0 and len(x_indices) > 0:
area = len(y_indices) # Calculate region area
print(f"Labeling region {unique_id}, area: {area} pixels")
if area < 30: # Skip very small regions
continue
# Use different colors for different IDs to enhance visual distinction
color_r = (unique_id * 50 + 80) % 200 / 255.0 + 0.2
color_g = (unique_id * 120 + 40) % 200 / 255.0 + 0.2
color_b = (unique_id * 180 + 20) % 200 / 255.0 + 0.2
color = [color_r, color_g, color_b]
# Adjust transparency based on area size
adaptive_alpha = min(0.3, max(0.1, 0.1 + area / 100000))
# Extract edges of this region
kernel = np.ones((3, 3), np.uint8)
dilated = cv2.dilate(mask.astype(np.uint8), kernel, iterations=1)
eroded = cv2.erode(mask.astype(np.uint8), kernel, iterations=1)
edge = dilated.astype(bool) & (~eroded.astype(bool))
# Build label text
label = f"{unique_id}"
# First draw the main body of the region
vis_mask = visual.draw_binary_mask_with_number(
mask,
text=label,
label_mode=label_mode,
alpha=adaptive_alpha,
anno_mode=anno_mode,
color=color,
font_size=20
)
# Enhance edges (add border effect for all parts)
edge_color = [min(c*1.3, 1.0) for c in color] # Slightly brighter edge color
vis_mask = visual.draw_binary_mask(
edge,
alpha=0.8, # Lower transparency for edges to make them more visible
color=edge_color
)
im = vis_mask.get_image()
return group_ids, im
File diff suppressed because it is too large Load Diff
@@ -0,0 +1,6 @@
from . import models
from . import modules
from . import pipelines
from . import renderers
from . import representations
from . import utils
@@ -0,0 +1,135 @@
import importlib
__attributes = {
'SparseStructureEncoder': 'sparse_structure_vae',
'SparseStructureDecoder': 'sparse_structure_vae',
'SparseStructureFlowModel': 'sparse_structure_flow',
'SLatEncoder': 'structured_latent_vae',
'SLatGaussianDecoder': 'structured_latent_vae',
'SLatRadianceFieldDecoder': 'structured_latent_vae',
'SLatMeshDecoder': 'structured_latent_vae',
'ElasticSLatEncoder': 'structured_latent_vae',
'ElasticSLatGaussianDecoder': 'structured_latent_vae',
'ElasticSLatRadianceFieldDecoder': 'structured_latent_vae',
'ElasticSLatMeshDecoder': 'structured_latent_vae',
'SLatFlowModel': 'structured_latent_flow',
'ElasticSLatFlowModel': 'structured_latent_flow',
}
__submodules = []
__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]
def from_pretrained(path: str, **kwargs):
"""
Load a model from a pretrained checkpoint.
Args:
path: The path to the checkpoint. Can be either local path or a Hugging Face model name.
NOTE: config file and model file should take the name f'{path}.json' and f'{path}.safetensors' respectively.
**kwargs: Additional arguments for the model constructor.
"""
import os
import json
from safetensors.torch import load_file
is_local = os.path.exists(f"{path}.json") and os.path.exists(f"{path}.safetensors")
# print(f"is local: {is_local}, path: {path} because {os.path.exists(f'{path}.json')} and {os.path.exists(f'{path}.safetensors')}")
if is_local:
config_file = f"{path}.json"
model_file = f"{path}.safetensors"
else:
from huggingface_hub import hf_hub_download
path_parts = path.split('/')
repo_id = f'{path_parts[0]}/{path_parts[1]}'
model_name = '/'.join(path_parts[2:])
config_file = hf_hub_download(repo_id, f"{model_name}.json")
model_file = hf_hub_download(repo_id, f"{model_name}.safetensors")
with open(config_file, 'r') as f:
config = json.load(f)
# print(f"Config loaded successfully: {config.get('name', 'Name not found in config')}")
if 'name' not in config:
raise ValueError(f"Config file missing required 'name' field")
model_class = config['name']
if model_class.lower() in [k.lower() for k in __attributes.keys()]:
# Try to find case-insensitive match
for k in __attributes.keys():
if k.lower() == model_class.lower():
model_class = k
break
# print(f"Using model class: {model_class}")
try:
model_constructor = __getattr__(model_class)
except AttributeError as e:
print(f"Model lookup failed: {e}")
raise ValueError(f"Model class '{model_class}' not found in available models: {list(__attributes.keys())}")
# print(f"Initializing model with args: {config.get('args', {})}")
model = model_constructor(**config.get('args', {}), **kwargs)
# Load state dict
state_dict = load_file(model_file)
# print(f"State dict loaded successfully from {model_file}")
# Check key compatibility
model_keys = set(model.state_dict().keys())
loaded_keys = set(state_dict.keys())
missing_keys = model_keys - loaded_keys
unexpected_keys = loaded_keys - model_keys
if missing_keys:
print(f"Missing keys in state dict: {missing_keys}")
if unexpected_keys:
print(f"Unexpected keys in state dict: {unexpected_keys}")
# Load state dict with strict=False to allow missing keys
model.load_state_dict(state_dict, strict=False)
return model
# For Pylance
if __name__ == '__main__':
from .sparse_structure_vae import (
SparseStructureEncoder,
SparseStructureDecoder,
)
from .sparse_structure_flow import SparseStructureFlowModel
from .structured_latent_vae import (
SLatEncoder,
SLatGaussianDecoder,
SLatRadianceFieldDecoder,
SLatMeshDecoder,
ElasticSLatEncoder,
ElasticSLatGaussianDecoder,
ElasticSLatRadianceFieldDecoder,
ElasticSLatMeshDecoder,
)
from .structured_latent_flow import (
SLatFlowModel,
ElasticSLatFlowModel,
)
@@ -0,0 +1,67 @@
"""
This file defines a mixin class for sparse transformers that enables elastic memory management.
It provides functionality to dynamically adjust memory usage by controlling gradient checkpointing
across transformer blocks, allowing for trading computation for memory efficiency.
"""
from contextlib import contextmanager
from typing import *
import math
from ..modules import sparse as sp
from ..utils.elastic_utils import ElasticModuleMixin
class SparseTransformerElasticMixin(ElasticModuleMixin):
"""
A mixin class for sparse transformers that provides elastic memory management capabilities.
Extends the base ElasticModuleMixin with sparse tensor-specific functionality.
"""
def _get_input_size(self, x: sp.SparseTensor, *args, **kwargs):
"""
Determines the input size from a sparse tensor.
Args:
x: A SparseTensor input
*args, **kwargs: Additional arguments (unused)
Returns:
The size of the feature dimension of the sparse tensor
"""
return x.feats.shape[0]
@contextmanager
def with_mem_ratio(self, mem_ratio=1.0):
"""
Context manager that temporarily adjusts memory usage by enabling gradient checkpointing
for a portion of the transformer blocks based on the specified memory ratio.
Args:
mem_ratio: A value between 0 and 1 indicating the desired memory ratio.
1.0 means use all available memory (no checkpointing).
Lower values enable more checkpointing to reduce memory usage.
Yields:
The exact memory ratio that could be achieved with the block granularity.
"""
if mem_ratio == 1.0:
# No memory optimization needed if ratio is 1.0
yield 1.0
return
# Calculate how many blocks should use checkpointing
num_blocks = len(self.blocks)
num_checkpoint_blocks = min(math.ceil((1 - mem_ratio) * num_blocks) + 1, num_blocks)
# Calculate the actual memory ratio based on the number of checkpointed blocks
exact_mem_ratio = 1 - (num_checkpoint_blocks - 1) / num_blocks
# Enable checkpointing for the calculated number of blocks
for i in range(num_blocks):
self.blocks[i].use_checkpoint = i < num_checkpoint_blocks
yield exact_mem_ratio
# Restore all blocks to not use checkpointing after context exit
for i in range(num_blocks):
self.blocks[i].use_checkpoint = False
@@ -0,0 +1,299 @@
"""
This file implements a Sparse Structure Flow model for 3D data generation or transformation.
It contains a transformer-based architecture that processes 3D volumes by:
1. Embedding timesteps for diffusion/flow-based modeling
2. Patchifying 3D inputs for efficient processing
3. Using cross-attention mechanisms to condition the generation on external features
4. Supporting various positional encoding schemes for 3D data
The model is designed for high-dimensional structure generation with conditional inputs
and follows a transformer-based architecture similar to DiT (Diffusion Transformers).
"""
from typing import *
import torch
import torch.nn as nn
import torch.nn.functional as F
import numpy as np
from ..modules.utils import convert_module_to_f16, convert_module_to_f32
from ..modules.transformer import AbsolutePositionEmbedder, ModulatedTransformerCrossBlock
from ..modules.spatial import patchify, unpatchify
class TimestepEmbedder(nn.Module):
"""
Embeds scalar timesteps into vector representations.
This is crucial for diffusion models where the model needs to know
which noise level (timestep) it's currently operating at.
"""
def __init__(self, hidden_size, frequency_embedding_size=256):
"""
Initialize the timestep embedder.
Args:
hidden_size: Dimension of the output embeddings
frequency_embedding_size: Dimension of the intermediate frequency embeddings
"""
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 similar to positional encodings in transformers.
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.
"""
# Implementation based on OpenAI's GLIDE repository
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):
"""
Embed timesteps into vectors.
Args:
t: Timesteps to embed [batch_size]
Returns:
Embedded timesteps [batch_size, hidden_size]
"""
t_freq = self.timestep_embedding(t, self.frequency_embedding_size)
t_emb = self.mlp(t_freq)
return t_emb
class SparseStructureFlowModel(nn.Module):
"""
A transformer-based model for processing 3D data with conditional inputs.
The model patchifies 3D volumes, processes them with transformer blocks,
and then reconstructs the 3D volume at the output.
"""
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,
patch_size: int = 2,
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,
):
"""
Initialize the Sparse Structure Flow model.
Args:
resolution: Input resolution (assumes cubic inputs of shape [resolution, resolution, resolution])
in_channels: Number of input channels
model_channels: Number of model's internal channels
cond_channels: Number of channels in conditional input
out_channels: Number of output channels
num_blocks: Number of transformer blocks
num_heads: Number of attention heads (defaults to model_channels // num_head_channels)
num_head_channels: Number of channels per attention head
mlp_ratio: Ratio for MLP hidden dimension relative to model_channels
patch_size: Size of patches for patchifying the input
pe_mode: Type of positional encoding ("ape" for absolute, "rope" for rotary)
use_fp16: Whether to use FP16 precision for most operations
use_checkpoint: Whether to use gradient checkpointing to save memory
share_mod: Whether to share modulation layers across blocks
qk_rms_norm: Whether to use RMS normalization for query and key in self-attention
qk_rms_norm_cross: Whether to use RMS normalization for query and key in cross-attention
"""
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.patch_size = patch_size
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.dtype = torch.float16 if use_fp16 else torch.float32
# Timestep embedding network
self.t_embedder = TimestepEmbedder(model_channels)
# Optional shared modulation for all blocks
if share_mod:
self.adaLN_modulation = nn.Sequential(
nn.SiLU(),
nn.Linear(model_channels, 6 * model_channels, bias=True)
)
# Set up positional encoding
if pe_mode == "ape":
pos_embedder = AbsolutePositionEmbedder(model_channels, 3)
# Create a grid of 3D coordinates for each patch position
coords = torch.meshgrid(*[torch.arange(res, device=self.device) for res in [resolution // patch_size] * 3], indexing='ij')
coords = torch.stack(coords, dim=-1).reshape(-1, 3)
pos_emb = pos_embedder(coords)
self.register_buffer("pos_emb", pos_emb)
# Input projection layer
self.input_layer = nn.Linear(in_channels * patch_size**3, model_channels)
# Transformer blocks with cross-attention for conditioning
self.blocks = nn.ModuleList([
ModulatedTransformerCrossBlock(
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=share_mod,
qk_rms_norm=self.qk_rms_norm,
qk_rms_norm_cross=self.qk_rms_norm_cross,
)
for _ in range(num_blocks)
])
# Output projection layer
self.out_layer = nn.Linear(model_channels, out_channels * patch_size**3)
# Initialize model weights
self.initialize_weights()
if use_fp16:
self.convert_to_fp16()
@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 transformer blocks of the model to float16 for improved efficiency.
"""
self.blocks.apply(convert_module_to_f16)
def convert_to_fp32(self) -> None:
"""
Convert the transformer blocks of the model back to float32 (e.g., for inference).
"""
self.blocks.apply(convert_module_to_f32)
def initialize_weights(self) -> None:
"""
Initialize the weights of the model using carefully chosen initialization schemes.
"""
# Initialize transformer layers with Xavier uniform initialization
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)
# Initialize timestep embedding MLP with normal distribution
nn.init.normal_(self.t_embedder.mlp[0].weight, std=0.02)
nn.init.normal_(self.t_embedder.mlp[2].weight, std=0.02)
# Zero-out adaLN modulation layers to ensure stable training initially
if self.share_mod:
nn.init.constant_(self.adaLN_modulation[-1].weight, 0)
nn.init.constant_(self.adaLN_modulation[-1].bias, 0)
else:
for block in self.blocks:
nn.init.constant_(block.adaLN_modulation[-1].weight, 0)
nn.init.constant_(block.adaLN_modulation[-1].bias, 0)
# Zero-out output layers to ensure initial predictions are near zero
nn.init.constant_(self.out_layer.weight, 0)
nn.init.constant_(self.out_layer.bias, 0)
def forward(self, x: torch.Tensor, t: torch.Tensor, cond: torch.Tensor) -> torch.Tensor:
"""
Forward pass of the model.
Args:
x: Input tensor of shape [batch_size, in_channels, resolution, resolution, resolution]
t: Timestep tensor of shape [batch_size]
cond: Conditional input tensor
Returns:
Output tensor of shape [batch_size, out_channels, resolution, resolution, resolution]
"""
# Validate input shape
assert [*x.shape] == [x.shape[0], self.in_channels, *[self.resolution] * 3], \
f"Input shape mismatch, got {x.shape}, expected {[x.shape[0], self.in_channels, *[self.resolution] * 3]}"
# Patchify the input volume and reshape for transformer processing
h = patchify(x, self.patch_size)
h = h.view(*h.shape[:2], -1).permute(0, 2, 1).contiguous() # [B, num_patches, patch_dim]
# Project to model dimension
h = self.input_layer(h)
# Add positional embeddings
h = h + self.pos_emb[None]
# Get timestep embeddings
t_emb = self.t_embedder(t)
if self.share_mod:
t_emb = self.adaLN_modulation(t_emb)
# Convert to appropriate dtype for computation
t_emb = t_emb.type(self.dtype)
h = h.type(self.dtype)
cond = cond.type(self.dtype)
# print("transfer cond")
# print("*" * 20)
# print(cond.shape) # torch.Size([4, 4122, 1024])
# Process through transformer blocks
for block in self.blocks:
h = block(h, t_emb, cond)
# print("transferred ")
# Convert back to original dtype
h = h.type(x.dtype)
# Final normalization and projection
h = F.layer_norm(h, h.shape[-1:])
h = self.out_layer(h)
# Reshape and unpatchify to get final 3D output
h = h.permute(0, 2, 1).view(h.shape[0], h.shape[2], *[self.resolution // self.patch_size] * 3)
h = unpatchify(h, self.patch_size).contiguous()
return h
@@ -0,0 +1,450 @@
"""
sparse_structure_vae.py
This file implements a Variational Autoencoder (VAE) for 3D sparse structural representations.
It's part of the TRELLIS framework and contains components for encoding volumetric data
into a latent space and decoding it back to volumetric representation.
The implementation includes:
- 3D normalization layers
- 3D residual blocks for feature extraction
- 3D downsampling and upsampling blocks for resolution changes
- Encoder (SparseStructureEncoder) that maps input volumes to a latent distribution
- Decoder (SparseStructureDecoder) that reconstructs volumes from latent codes
This VAE architecture is specifically designed for capturing structural information
in a compressed latent representation that can be sampled probabilistically.
"""
from typing import *
import torch
import torch.nn as nn
import torch.nn.functional as F
from ..modules.norm import GroupNorm32, ChannelLayerNorm32
from ..modules.spatial import pixel_shuffle_3d
from ..modules.utils import zero_module, convert_module_to_f16, convert_module_to_f32
def norm_layer(norm_type: str, *args, **kwargs) -> nn.Module:
"""
Return a normalization layer based on the specified type.
Args:
norm_type: Either "group" for GroupNorm or "layer" for LayerNorm
*args, **kwargs: Arguments passed to the normalization layer
Returns:
An instance of the requested normalization layer
"""
if norm_type == "group":
return GroupNorm32(32, *args, **kwargs)
elif norm_type == "layer":
return ChannelLayerNorm32(*args, **kwargs)
else:
raise ValueError(f"Invalid norm type {norm_type}")
class ResBlock3d(nn.Module):
"""
3D Residual Block with two convolutions and a skip connection.
The block applies normalization, activation, and convolution twice,
with a skip connection from the input to the output.
"""
def __init__(
self,
channels: int,
out_channels: Optional[int] = None,
norm_type: Literal["group", "layer"] = "layer",
):
"""
Initialize a 3D ResBlock.
Args:
channels: Number of input channels
out_channels: Number of output channels (defaults to input channels)
norm_type: Type of normalization to use
"""
super().__init__()
self.channels = channels
self.out_channels = out_channels or channels
# First normalization and convolution
self.norm1 = norm_layer(norm_type, channels)
self.norm2 = norm_layer(norm_type, self.out_channels)
self.conv1 = nn.Conv3d(channels, self.out_channels, 3, padding=1)
# Second convolution is initialized with zeros for stable training
self.conv2 = zero_module(nn.Conv3d(self.out_channels, self.out_channels, 3, padding=1))
# Skip connection: identity if channels match, otherwise 1x1 conv
self.skip_connection = nn.Conv3d(channels, self.out_channels, 1) if channels != self.out_channels else nn.Identity()
def forward(self, x: torch.Tensor) -> torch.Tensor:
"""
Forward pass for the ResBlock.
Args:
x: Input tensor of shape [B, C, D, H, W]
Returns:
Output tensor after residual computation
"""
h = self.norm1(x)
h = F.silu(h)
h = self.conv1(h)
h = self.norm2(h)
h = F.silu(h)
h = self.conv2(h)
h = h + self.skip_connection(x) # Residual connection
return h
class DownsampleBlock3d(nn.Module):
"""
3D downsampling block to reduce spatial dimensions by a factor of 2.
Supports downsampling via strided convolution or average pooling.
"""
def __init__(
self,
in_channels: int,
out_channels: int,
mode: Literal["conv", "avgpool"] = "conv",
):
"""
Initialize a 3D downsampling block.
Args:
in_channels: Number of input channels
out_channels: Number of output channels
mode: Downsampling method ("conv" or "avgpool")
"""
assert mode in ["conv", "avgpool"], f"Invalid mode {mode}"
super().__init__()
self.in_channels = in_channels
self.out_channels = out_channels
if mode == "conv":
self.conv = nn.Conv3d(in_channels, out_channels, 2, stride=2)
elif mode == "avgpool":
assert in_channels == out_channels, "Pooling mode requires in_channels to be equal to out_channels"
def forward(self, x: torch.Tensor) -> torch.Tensor:
"""
Forward pass for the downsampling block.
Args:
x: Input tensor of shape [B, C, D, H, W]
Returns:
Downsampled tensor
"""
if hasattr(self, "conv"):
return self.conv(x)
else:
return F.avg_pool3d(x, 2)
class UpsampleBlock3d(nn.Module):
"""
3D upsampling block to increase spatial dimensions by a factor of 2.
Supports upsampling via transposed convolution or nearest-neighbor interpolation.
"""
def __init__(
self,
in_channels: int,
out_channels: int,
mode: Literal["conv", "nearest"] = "conv",
):
"""
Initialize a 3D upsampling block.
Args:
in_channels: Number of input channels
out_channels: Number of output channels
mode: Upsampling method ("conv" or "nearest")
"""
assert mode in ["conv", "nearest"], f"Invalid mode {mode}"
super().__init__()
self.in_channels = in_channels
self.out_channels = out_channels
if mode == "conv":
# For pixel shuffle upsampling, we need 8x channels (2³ = 8)
self.conv = nn.Conv3d(in_channels, out_channels*8, 3, padding=1)
elif mode == "nearest":
assert in_channels == out_channels, "Nearest mode requires in_channels to be equal to out_channels"
def forward(self, x: torch.Tensor) -> torch.Tensor:
"""
Forward pass for the upsampling block.
Args:
x: Input tensor of shape [B, C, D, H, W]
Returns:
Upsampled tensor
"""
if hasattr(self, "conv"):
x = self.conv(x)
return pixel_shuffle_3d(x, 2) # 3D pixel shuffle for upsampling
else:
return F.interpolate(x, scale_factor=2, mode="nearest")
class SparseStructureEncoder(nn.Module):
"""
Encoder for Sparse Structure (\mathcal{E}_S in the paper Sec. 3.3).
Takes a 3D volume as input and encodes it into a latent distribution (mean and logvar).
Can sample from this distribution to get a latent representation.
Args:
in_channels (int): Channels of the input.
latent_channels (int): Channels of the latent representation.
num_res_blocks (int): Number of residual blocks at each resolution.
channels (List[int]): Channels of the encoder blocks.
num_res_blocks_middle (int): Number of residual blocks in the middle.
norm_type (Literal["group", "layer"]): Type of normalization layer.
use_fp16 (bool): Whether to use FP16.
"""
def __init__(
self,
in_channels: int,
latent_channels: int,
num_res_blocks: int,
channels: List[int],
num_res_blocks_middle: int = 2,
norm_type: Literal["group", "layer"] = "layer",
use_fp16: bool = False,
):
"""
Initialize the encoder for sparse structure.
"""
super().__init__()
self.in_channels = in_channels
self.latent_channels = latent_channels
self.num_res_blocks = num_res_blocks
self.channels = channels
self.num_res_blocks_middle = num_res_blocks_middle
self.norm_type = norm_type
self.use_fp16 = use_fp16
self.dtype = torch.float16 if use_fp16 else torch.float32
# Initial projection from input to feature space
self.input_layer = nn.Conv3d(in_channels, channels[0], 3, padding=1)
# Encoder blocks with progressive downsampling
self.blocks = nn.ModuleList([])
for i, ch in enumerate(channels):
# Add residual blocks at the current resolution
self.blocks.extend([
ResBlock3d(ch, ch)
for _ in range(num_res_blocks)
])
# Add downsampling block if not at the final resolution
if i < len(channels) - 1:
self.blocks.append(
DownsampleBlock3d(ch, channels[i+1])
)
# Middle blocks at the lowest resolution
self.middle_block = nn.Sequential(*[
ResBlock3d(channels[-1], channels[-1])
for _ in range(num_res_blocks_middle)
])
# Output layer produces both mean and logvar for the latent distribution
self.out_layer = nn.Sequential(
norm_layer(norm_type, channels[-1]),
nn.SiLU(),
nn.Conv3d(channels[-1], latent_channels*2, 3, padding=1)
)
if use_fp16:
self.convert_to_fp16()
@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.use_fp16 = True
self.dtype = torch.float16
self.blocks.apply(convert_module_to_f16)
self.middle_block.apply(convert_module_to_f16)
def convert_to_fp32(self) -> None:
"""
Convert the torso of the model to float32.
"""
self.use_fp16 = False
self.dtype = torch.float32
self.blocks.apply(convert_module_to_f32)
self.middle_block.apply(convert_module_to_f32)
def forward(self, x: torch.Tensor, sample_posterior: bool = False, return_raw: bool = False) -> torch.Tensor:
"""
Forward pass through the encoder.
Args:
x: Input tensor of shape [B, C, D, H, W]
sample_posterior: Whether to sample from the posterior distribution or just return mean
return_raw: Whether to return the raw outputs (z, mean, logvar) instead of just z
Returns:
Either the latent representation or a tuple of (z, mean, logvar) if return_raw=True
"""
h = self.input_layer(x)
h = h.type(self.dtype) # Convert to FP16 if needed
# Process through encoder blocks
for block in self.blocks:
h = block(h)
h = self.middle_block(h)
h = h.type(x.dtype) # Convert back to input dtype
h = self.out_layer(h)
# Split output into mean and log variance
mean, logvar = h.chunk(2, dim=1)
# Sample from the posterior if requested
if sample_posterior:
std = torch.exp(0.5 * logvar)
z = mean + std * torch.randn_like(std) # Reparameterization trick
else:
z = mean
if return_raw:
return z, mean, logvar
return z
class SparseStructureDecoder(nn.Module):
"""
Decoder for Sparse Structure (\mathcal{D}_S in the paper Sec. 3.3).
Takes a latent representation and decodes it back to a 3D volume.
Uses a symmetric architecture to the encoder with upsampling instead of downsampling.
Args:
out_channels (int): Channels of the output.
latent_channels (int): Channels of the latent representation.
num_res_blocks (int): Number of residual blocks at each resolution.
channels (List[int]): Channels of the decoder blocks.
num_res_blocks_middle (int): Number of residual blocks in the middle.
norm_type (Literal["group", "layer"]): Type of normalization layer.
use_fp16 (bool): Whether to use FP16.
"""
def __init__(
self,
out_channels: int,
latent_channels: int,
num_res_blocks: int,
channels: List[int],
num_res_blocks_middle: int = 2,
norm_type: Literal["group", "layer"] = "layer",
use_fp16: bool = False,
):
"""
Initialize the decoder for sparse structure.
"""
super().__init__()
self.out_channels = out_channels
self.latent_channels = latent_channels
self.num_res_blocks = num_res_blocks
self.channels = channels
self.num_res_blocks_middle = num_res_blocks_middle
self.norm_type = norm_type
self.use_fp16 = use_fp16
self.dtype = torch.float16 if use_fp16 else torch.float32
# Initial projection from latent space to feature space
self.input_layer = nn.Conv3d(latent_channels, channels[0], 3, padding=1)
# Middle blocks at the lowest resolution
self.middle_block = nn.Sequential(*[
ResBlock3d(channels[0], channels[0])
for _ in range(num_res_blocks_middle)
])
# Decoder blocks with progressive upsampling
self.blocks = nn.ModuleList([])
for i, ch in enumerate(channels):
# Add residual blocks at the current resolution
self.blocks.extend([
ResBlock3d(ch, ch)
for _ in range(num_res_blocks)
])
# Add upsampling block if not at the final resolution
if i < len(channels) - 1:
self.blocks.append(
UpsampleBlock3d(ch, channels[i+1])
)
# Final output layer to generate the desired output channels
self.out_layer = nn.Sequential(
norm_layer(norm_type, channels[-1]),
nn.SiLU(),
nn.Conv3d(channels[-1], out_channels, 3, padding=1)
)
if use_fp16:
self.convert_to_fp16()
@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.use_fp16 = True
self.dtype = torch.float16
self.blocks.apply(convert_module_to_f16)
self.middle_block.apply(convert_module_to_f16)
def convert_to_fp32(self) -> None:
"""
Convert the torso of the model to float32.
"""
self.use_fp16 = False
self.dtype = torch.float32
self.blocks.apply(convert_module_to_f32)
self.middle_block.apply(convert_module_to_f32)
def forward(self, x: torch.Tensor) -> torch.Tensor:
"""
Forward pass through the decoder.
Args:
x: Latent representation tensor of shape [B, C, D, H, W]
Returns:
Reconstructed output tensor
"""
h = self.input_layer(x)
h = h.type(self.dtype) # Convert to FP16 if needed
h = self.middle_block(h)
# Process through decoder blocks
for block in self.blocks:
h = block(h)
h = h.type(x.dtype) # Convert back to input dtype
h = self.out_layer(h)
return h
@@ -0,0 +1,470 @@
from typing import *
from einops import rearrange
import torch
import torch.nn as nn
import torch.nn.functional as F
import numpy as np
from ..modules.utils import zero_module, convert_module_to_f16, convert_module_to_f32
from ..modules.transformer import AbsolutePositionEmbedder
from ..modules.norm import LayerNorm32
from ..modules import sparse as sp
from ..modules.sparse.transformer import ModulatedSparseTransformerCrossBlock
from .sparse_structure_flow import TimestepEmbedder
from .sparse_elastic_mixin import SparseTransformerElasticMixin
class SparseResBlock3d(nn.Module):
"""
3D Sparse Residual Block with time embedding conditioning.
This block performs normalization, convolution operations on sparse tensors,
and incorporates time embeddings via adaptive layer normalization.
Supports optional up/downsampling.
"""
def __init__(
self,
channels: int,
emb_channels: int,
out_channels: Optional[int] = None,
downsample: bool = False,
upsample: bool = False,
):
super().__init__()
self.channels = channels
self.emb_channels = emb_channels
self.out_channels = out_channels or channels
self.downsample = downsample
self.upsample = upsample
assert not (downsample and upsample), "Cannot downsample and upsample at the same time"
# First normalization and convolution
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)
# Second convolution initialized to zero for stable training
self.conv2 = zero_module(sp.SparseConv3d(self.out_channels, self.out_channels, 3))
# Time embedding projection for adaptive layer norm
self.emb_layers = nn.Sequential(
nn.SiLU(),
nn.Linear(emb_channels, 2 * self.out_channels, bias=True),
)
# Skip connection with linear projection if channel dimensions change
self.skip_connection = sp.SparseLinear(channels, self.out_channels) if channels != self.out_channels else nn.Identity()
# Optional up/downsampling
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:
"""Apply up/downsampling if configured"""
if self.updown is not None:
x = self.updown(x)
return x
def forward(self, x: sp.SparseTensor, emb: torch.Tensor) -> sp.SparseTensor:
"""
Forward pass of the residual block.
Args:
x: Input sparse tensor
emb: Time embedding tensor
Returns:
Processed sparse tensor
"""
# print(f"number of points in the input: {x.coords.shape[0]}")
# Project embedding to scale and shift factors
emb_out = self.emb_layers(emb).type(x.dtype)
scale, shift = torch.chunk(emb_out, 2, dim=1)
# Apply up/downsampling if needed
x = self._updown(x)
# Main processing path
h = x.replace(self.norm1(x.feats))
h = h.replace(F.silu(h.feats))
h = self.conv1(h)
# Apply adaptive layer norm using scale and shift from time embedding
h = h.replace(self.norm2(h.feats)) * (1 + scale) + shift
h = h.replace(F.silu(h.feats))
h = self.conv2(h)
# Residual connection
h = h + self.skip_connection(x)
return h
class SLatFlowModel(nn.Module):
"""
Structured Latent Flow Model for 3D generative modeling.
This model combines sparse convolutions with transformer blocks and
supports conditional generation. It uses a U-Net-like architecture with
skip connections and has optional mixed precision support.
"""
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,
patch_size: int = 2,
num_io_res_blocks: int = 2,
io_block_channels: List[int] = None,
pe_mode: Literal["ape", "rope"] = "ape",
use_fp16: bool = False,
use_checkpoint: bool = False,
use_skip_connection: bool = True,
share_mod: bool = False,
qk_rms_norm: bool = False,
qk_rms_norm_cross: bool = False,
):
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.patch_size = patch_size
self.num_io_res_blocks = num_io_res_blocks
self.io_block_channels = io_block_channels
self.pe_mode = pe_mode
self.use_fp16 = use_fp16
self.use_checkpoint = use_checkpoint
self.use_skip_connection = use_skip_connection
self.share_mod = share_mod
self.qk_rms_norm = qk_rms_norm
self.qk_rms_norm_cross = qk_rms_norm_cross
self.dtype = torch.float16 if use_fp16 else torch.float32
# Validate configurations
if self.io_block_channels is not None:
assert int(np.log2(patch_size)) == np.log2(patch_size), "Patch size must be a power of 2"
assert np.log2(patch_size) == len(io_block_channels), "Number of IO ResBlocks must match the number of stages"
# Time step embedder
self.t_embedder = TimestepEmbedder(model_channels)
# Shared modulation for all transformer blocks if enabled
if share_mod:
self.adaLN_modulation = nn.Sequential(
nn.SiLU(),
nn.Linear(model_channels, 6 * model_channels, bias=True)
)
self.part_max_size = 50
# Positional embedding for transformer blocks
if pe_mode == "ape":
self.pos_embedder = AbsolutePositionEmbedder(model_channels)
self.part_pe = nn.Embedding(self.part_max_size + 1, model_channels) # +1 for overall object
self.part_pe_proj = nn.Linear(model_channels, model_channels)
# Mask embedding
self.dinov2_hidden_size = 1024
self.mask_group_emb_dim = 128
self.mask_group_emb = nn.Embedding(self.part_max_size + 1, self.mask_group_emb_dim) # +1 for background
self.mask_group_emb_proj = nn.Linear(self.mask_group_emb_dim, self.dinov2_hidden_size)
# Input projection layer
self.input_layer = sp.SparseLinear(in_channels, model_channels if io_block_channels is None else io_block_channels[0])
# Input processing blocks (downsampling path)
self.input_blocks = nn.ModuleList([])
# print(f"io_block_channels: {io_block_channels}") # io_block_channels: [128]
# print(f"model_channels: {model_channels}") # model_channels: 1024
if io_block_channels is not None:
for chs, next_chs in zip(io_block_channels, io_block_channels[1:] + [model_channels]):
# Add regular residual blocks at current resolution
self.input_blocks.extend([
SparseResBlock3d(
chs,
model_channels,
out_channels=chs,
)
for _ in range(num_io_res_blocks-1)
])
# Add downsampling block at the end of each resolution level
self.input_blocks.append(
SparseResBlock3d(
chs,
model_channels,
out_channels=next_chs,
downsample=True,
)
)
# Core transformer blocks
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(num_blocks)
])
# Output processing blocks (upsampling path)
self.out_blocks = nn.ModuleList([])
if io_block_channels is not None:
for chs, prev_chs in zip(reversed(io_block_channels), [model_channels] + list(reversed(io_block_channels[1:]))):
# Add upsampling block at the beginning of each resolution level
self.out_blocks.append(
SparseResBlock3d(
prev_chs * 2 if self.use_skip_connection else prev_chs,
model_channels,
out_channels=chs,
upsample=True,
)
)
# Add regular residual blocks at current resolution
self.out_blocks.extend([
SparseResBlock3d(
chs * 2 if self.use_skip_connection else chs,
model_channels,
out_channels=chs,
)
for _ in range(num_io_res_blocks-1)
])
# Final output projection
self.out_layer = sp.SparseLinear(model_channels if io_block_channels is None else io_block_channels[0], out_channels)
# Initialize model weights
self.initialize_weights()
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 for mixed precision training.
"""
self.input_blocks.apply(convert_module_to_f16)
self.blocks.apply(convert_module_to_f16)
self.out_blocks.apply(convert_module_to_f16)
def convert_to_fp32(self) -> None:
"""
Convert the torso of the model back to float32.
"""
self.input_blocks.apply(convert_module_to_f32)
self.blocks.apply(convert_module_to_f32)
self.out_blocks.apply(convert_module_to_f32)
def initialize_weights(self) -> None:
"""
Initialize model weights with specialized initialization for different components.
"""
# Initialize transformer layers with Xavier uniform initialization
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)
# Initialize timestep embedding MLP with normal distribution
nn.init.normal_(self.t_embedder.mlp[0].weight, std=0.02)
nn.init.normal_(self.t_embedder.mlp[2].weight, std=0.02)
# Zero-out adaLN modulation layers for stable training
if self.share_mod:
nn.init.constant_(self.adaLN_modulation[-1].weight, 0)
nn.init.constant_(self.adaLN_modulation[-1].bias, 0)
else:
for block in self.blocks:
nn.init.constant_(block.adaLN_modulation[-1].weight, 0)
nn.init.constant_(block.adaLN_modulation[-1].bias, 0)
# Zero-out output layers for stable training
nn.init.constant_(self.out_layer.weight, 0)
nn.init.constant_(self.out_layer.bias, 0)
# part embedding initialization
nn.init.zeros_(self.part_pe_proj.weight)
nn.init.zeros_(self.part_pe_proj.bias)
# Initialize layer positional embeddings
self.part_pe.weight.data.normal_(mean=0.0,std=0.02)
# Initialize group embedding
nn.init.zeros_(self.mask_group_emb_proj.weight)
nn.init.zeros_(self.mask_group_emb_proj.bias)
self.mask_group_emb.weight.data.normal_(mean=0.0, std=0.02)
def forward(self, x: sp.SparseTensor, t: torch.Tensor, cond: torch.Tensor, **kwargs) -> sp.SparseTensor:
"""
Forward pass of the Structured Latent Flow model.
Args:
x: Input sparse tensor
t: Timestep embedding inputs
cond: Conditional input for cross-attention
**kwargs: Additional arguments, including part_layouts if available
Returns:
Output sparse tensor
"""
# x = x.type(self.dtype)
# t = t.type(self.dtype)
# cond = cond.type(self.dtype)
input_dtype = x.dtype
masks = kwargs['masks'] # [b, h, w]
# Ensure masks are always long type regardless of source
masks = masks.long() # Explicitly convert to long type for embedding
masks = rearrange(masks, 'b h w -> b (h w)') # [b, h*w]
masks_emb = self.mask_group_emb(masks) # [b, h*w, 128]
masks_emb = self.mask_group_emb_proj(masks_emb) # [b, h*w, 1024]
group_emb = torch.zeros((cond.shape[0], cond.shape[1], masks_emb.shape[2]), device=cond.device, dtype=cond.dtype)
group_emb[:, :masks_emb.shape[1], :] = masks_emb
cond = cond + group_emb
cond = cond.type(self.dtype)
# Store original batch IDs for later restoration
original_batch_ids = x.coords[:, 0].clone()
# Create new batch IDs to represent individual parts (instead of batches)
new_batch_ids = torch.zeros_like(original_batch_ids)
# Assign unique IDs to each part across all batches
part_layouts = kwargs['part_layouts']
part_id = 0
len_before = 0
batch_last_partid = []
for batch_idx, part_layout in enumerate(part_layouts):
for layout_idx, layout in enumerate(part_layout):
adjusted_layout = slice(layout.start + len_before, layout.stop + len_before, layout.step)
new_batch_ids[adjusted_layout] = part_id
part_id += 1
batch_last_partid.append(part_id)
len_before += part_layout[-1].stop
# Project input to model dimensions and convert to target dtype
x = self.input_layer(x).type(self.dtype)
x = sp.SparseTensor(
feats = x.feats,
coords = torch.cat([new_batch_ids.view(-1, 1), x.coords[:, 1:]], dim=1),)
# Process timestep embedding and condition input
t_emb = self.t_embedder(t)
if self.share_mod:
t_emb = self.adaLN_modulation(t_emb)
t_emb = t_emb.type(self.dtype)
t_emb_updown = []
for batch_idx, part_layout in enumerate(part_layouts):
t_emb_updown_batch = t_emb[batch_idx:batch_idx+1].repeat(len(part_layout), 1)
t_emb_updown.append(t_emb_updown_batch)
t_emb_updown = torch.cat(t_emb_updown, dim=0).type(self.dtype)
# Store features for skip connections
skips = []
# Downsampling path through input blocks
for block in self.input_blocks:
x = block(x, t_emb_updown)
skips.append(x.feats)
# Store part-wise batch IDs before transformer processing
part_wise_batch_ids = x.coords[:, 0].clone()
# Convert to batch-wise IDs for transformer blocks
new_transformer_batch_ids = torch.zeros_like(part_wise_batch_ids)
part_ids_in_each_object = torch.zeros_like(part_wise_batch_ids)
start_reform = 0
last_part_id = 0
for part_id in batch_last_partid:
mask = (part_wise_batch_ids >= last_part_id) & (part_wise_batch_ids < part_id)
new_transformer_batch_ids[mask] = start_reform
part_ids_in_each_object[mask] = part_wise_batch_ids[mask] - last_part_id
last_part_id = part_id
start_reform += 1
# Update coordinates with batch-wise IDs for transformer processing
h = sp.SparseTensor(
feats = x.feats,
coords = torch.cat([new_transformer_batch_ids.view(-1, 1), x.coords[:, 1:]], dim=1))
# Add positional embeddings for transformer blocks
if self.pe_mode == "ape":
# Add absolute positional embeddings to spatial coordinates
h = h + self.pos_embedder(h.coords[:, 1:]).type(self.dtype)
# Part-with PE; overall is 0
part_pe = self.part_pe(part_ids_in_each_object)
part_pe = self.part_pe_proj(part_pe)
h = h + part_pe.type(self.dtype)
else:
raise NotImplementedError
# Process with transformer blocks
for block in self.blocks:
h = block(h, t_emb, cond)
h = x.replace(feats=h.feats, coords=torch.cat([part_wise_batch_ids.view(-1, 1), h.coords[:, 1:]], dim=1))
# Upsampling path with output blocks and skip connections
for block, skip in zip(self.out_blocks, reversed(skips)):
if self.use_skip_connection:
h = block(h.replace(torch.cat([h.feats, skip], dim=1)), t_emb_updown)
else:
h = block(h, t_emb_updown)
h = h.replace(F.layer_norm(h.feats, h.feats.shape[-1:]))
h = self.out_layer(h.type(input_dtype))
h = sp.SparseTensor(
feats = h.feats,
coords = torch.cat([original_batch_ids.view(-1, 1), h.coords[:, 1:]], dim=1))
return h
class ElasticSLatFlowModel(SparseTransformerElasticMixin, SLatFlowModel):
"""
Structured Latent Flow Model with elastic memory management.
This class extends SLatFlowModel with memory-efficient operations,
allowing training with limited VRAM by dynamically managing memory
allocation for sparse tensors.
"""
pass
@@ -0,0 +1,4 @@
from .encoder import SLatEncoder, ElasticSLatEncoder
from .decoder_gs import SLatGaussianDecoder, ElasticSLatGaussianDecoder
from .decoder_rf import SLatRadianceFieldDecoder, ElasticSLatRadianceFieldDecoder
from .decoder_mesh import SLatMeshDecoder, ElasticSLatMeshDecoder
@@ -0,0 +1,185 @@
"""
Base Sparse Transformer Implementation for TRELLIS Framework
This file implements the base architecture for sparse transformers used in structured latent variable models.
It provides a configurable foundation with multiple attention mechanisms (full, windowed, shifted window)
and supports different positional encoding strategies. The sparse implementation allows for efficient
processing of data with varying density patterns.
The main class SparseTransformerBase serves as the foundation for encoder and decoder implementations
in the structured latent VAE models.
"""
from typing import *
import torch
import torch.nn as nn
from ...modules.utils import convert_module_to_f16, convert_module_to_f32
from ...modules import sparse as sp
from ...modules.transformer import AbsolutePositionEmbedder
from ...modules.sparse.transformer import SparseTransformerBlock
def block_attn_config(self):
"""
Return the attention configuration for each transformer block.
Generates configurations for each block based on the specified attention mode:
- shift_window: Uses serialized attention with shifting window patterns
- shift_sequence: Uses serialized attention with sequence shifts
- shift_order: Uses serialized attention with different serialization orders
- full: Uses standard full attention (non-sparse)
- swin: Uses Swin Transformer-style windowed attention
Yields:
Tuple containing attention mode and its parameters
"""
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.
Implements a transformer architecture that can work with sparse data structures,
supporting various attention mechanisms and positional encodings.
"""
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,
):
"""
Initialize the sparse transformer base model.
Args:
in_channels: Number of input channels
model_channels: Hidden dimension size
num_blocks: Number of transformer blocks
num_heads: Number of attention heads (calculated from head_channels if None)
num_head_channels: Number of channels per attention head
mlp_ratio: Ratio for MLP hidden dimension
attn_mode: Attention mechanism type
window_size: Size of attention window for windowed modes
pe_mode: Positional encoding mode (absolute or rotary)
use_fp16: Whether to use half precision
use_checkpoint: Whether to use gradient checkpointing
qk_rms_norm: Whether to use RMS normalization for query and key
"""
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
# Create positional embedder if using absolute positional encoding
if pe_mode == "ape":
self.pos_embedder = AbsolutePositionEmbedder(model_channels)
# Input projection layer
self.input_layer = sp.SparseLinear(in_channels, model_channels)
# Build transformer blocks with configurations from block_attn_config
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 precision.
Used for mixed precision training.
"""
self.blocks.apply(convert_module_to_f16)
def convert_to_fp32(self) -> None:
"""
Convert the torso of the model back to float32 precision.
Used after mixed precision training or inference.
"""
self.blocks.apply(convert_module_to_f32)
def initialize_weights(self) -> None:
"""
Initialize the weights of the model using Xavier uniform initialization.
This helps with training stability and convergence.
"""
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:
"""
Forward pass through the sparse transformer.
Args:
x: Input sparse tensor
Returns:
Processed sparse tensor after passing through all transformer blocks
"""
# Project input to model dimension
h = self.input_layer(x)
# Add positional embeddings if using absolute positional encoding
if self.pe_mode == "ape":
h = h + self.pos_embedder(x.coords[:, 1:])
# Convert to target precision
h = h.type(self.dtype)
# Pass through transformer blocks sequentially
for block in self.blocks:
h = block(h)
return h
@@ -0,0 +1,180 @@
"""
decoder_gs.py: Structured Latent Gaussian Decoder for 3D Representation Learning
This file contains decoder implementations that transform latent codes into 3D Gaussian
representations. The decoders use sparse transformer architectures for efficient processing
and flexible attention mechanisms. The main components are:
- SLatGaussianDecoder: Core decoder that maps latent codes to 3D Gaussians
- ElasticSLatGaussianDecoder: Memory-efficient variant with elastic memory management
"""
from typing import *
import torch
import torch.nn as nn
import torch.nn.functional as F
from ...modules import sparse as sp
from ...utils.random_utils import hammersley_sequence
from .base import SparseTransformerBase
from ...representations import Gaussian
from ..sparse_elastic_mixin import SparseTransformerElasticMixin
class SLatGaussianDecoder(SparseTransformerBase):
"""
Sparse Transformer-based decoder that converts latent codes to 3D Gaussian representations.
This decoder processes sparse tensors and outputs parameters for Gaussian primitives
that can be rendered in 3D space, including positions, features, scaling, rotation,
and opacity.
"""
def __init__(
self,
resolution: int, # The resolution of the 3D grid
model_channels: int, # Number of channels in the transformer layers
latent_channels: int, # Number of channels in the input latent code
num_blocks: int, # Number of transformer blocks
num_heads: Optional[int] = None, # Number of attention heads
num_head_channels: Optional[int] = 64, # Channels per attention head
mlp_ratio: float = 4, # Ratio for MLP size in transformer blocks
attn_mode: Literal["full", "shift_window", "shift_sequence", "shift_order", "swin"] = "swin", # Attention mechanism
window_size: int = 8, # Size of attention windows for windowed attention
pe_mode: Literal["ape", "rope"] = "ape", # Positional encoding mode
use_fp16: bool = False, # Whether to use half-precision
use_checkpoint: bool = False, # Whether to use gradient checkpointing
qk_rms_norm: bool = False, # Whether to use RMS normalization for attention
representation_config: dict = None, # Configuration for the Gaussian representation
):
super().__init__(
in_channels=latent_channels,
model_channels=model_channels,
num_blocks=num_blocks,
num_heads=num_heads,
num_head_channels=num_head_channels,
mlp_ratio=mlp_ratio,
attn_mode=attn_mode,
window_size=window_size,
pe_mode=pe_mode,
use_fp16=use_fp16,
use_checkpoint=use_checkpoint,
qk_rms_norm=qk_rms_norm,
)
self.resolution = resolution
self.rep_config = representation_config
self._calc_layout() # Calculate output tensor layout
self.out_layer = sp.SparseLinear(model_channels, self.out_channels) # Final projection layer
self._build_perturbation() # Build position perturbation for better initialization
self.initialize_weights()
if use_fp16:
self.convert_to_fp16()
def initialize_weights(self) -> None:
"""
Initialize model weights, with special handling for output layers.
Zero-initializes the output layer for stability.
"""
super().initialize_weights()
# Zero-out output layers:
nn.init.constant_(self.out_layer.weight, 0)
nn.init.constant_(self.out_layer.bias, 0)
def _build_perturbation(self) -> None:
"""
Build position perturbation for Gaussian means.
Uses Hammersley sequence for quasi-random uniform distribution of points,
then transforms to match the desired Gaussian spatial distribution.
"""
perturbation = [hammersley_sequence(3, i, self.rep_config['num_gaussians']) for i in range(self.rep_config['num_gaussians'])]
perturbation = torch.tensor(perturbation).float() * 2 - 1 # Scale to [-1, 1]
perturbation = perturbation / self.rep_config['voxel_size'] # Scale by voxel size
perturbation = torch.atanh(perturbation).to(self.device) # Apply inverse tanh for better gradient flow
self.register_buffer('offset_perturbation', perturbation) # Register as buffer (not a parameter)
def _calc_layout(self) -> None:
"""
Calculate the layout of the output tensor.
Defines the shape and size of each Gaussian parameter group (position, features, scaling, rotation, opacity)
and their positions in the output tensor.
"""
self.layout = {
'_xyz' : {'shape': (self.rep_config['num_gaussians'], 3), 'size': self.rep_config['num_gaussians'] * 3},
'_features_dc' : {'shape': (self.rep_config['num_gaussians'], 1, 3), 'size': self.rep_config['num_gaussians'] * 3},
'_scaling' : {'shape': (self.rep_config['num_gaussians'], 3), 'size': self.rep_config['num_gaussians'] * 3},
'_rotation' : {'shape': (self.rep_config['num_gaussians'], 4), 'size': self.rep_config['num_gaussians'] * 4},
'_opacity' : {'shape': (self.rep_config['num_gaussians'], 1), 'size': self.rep_config['num_gaussians']},
}
# Calculate ranges for each parameter group in the flattened output tensor
start = 0
for k, v in self.layout.items():
v['range'] = (start, start + v['size'])
start += v['size']
self.out_channels = start # Total number of output channels
def to_representation(self, x: sp.SparseTensor) -> List[Gaussian]:
"""
Convert a batch of network outputs to 3D Gaussian representations.
Args:
x: The [N x * x C] sparse tensor output by the network.
Returns:
list of Gaussian representations, one per batch item
"""
ret = []
for i in range(x.shape[0]):
# Create a new Gaussian representation object with proper configuration
representation = Gaussian(
sh_degree=0, # No spherical harmonics, just using DC term
aabb=[-0.5, -0.5, -0.5, 1.0, 1.0, 1.0], # Axis-aligned bounding box
mininum_kernel_size = self.rep_config['3d_filter_kernel_size'],
scaling_bias = self.rep_config['scaling_bias'],
opacity_bias = self.rep_config['opacity_bias'],
scaling_activation = self.rep_config['scaling_activation']
)
# Get base positions from sparse tensor coordinates
xyz = (x.coords[x.layout[i]][:, 1:].float() + 0.5) / self.resolution
# Process each parameter group
for k, v in self.layout.items():
if k == '_xyz':
# Handle positions with special perturbation logic
offset = x.feats[x.layout[i]][:, v['range'][0]:v['range'][1]].reshape(-1, *v['shape'])
offset = offset * self.rep_config['lr'][k] # Apply learning rate scale
if self.rep_config['perturb_offset']:
offset = offset + self.offset_perturbation # Add perturbation
# Transform offsets through tanh and scale appropriately
offset = torch.tanh(offset) / self.resolution * 0.5 * self.rep_config['voxel_size']
_xyz = xyz.unsqueeze(1) + offset
setattr(representation, k, _xyz.flatten(0, 1))
else:
# Handle other parameters (features, scaling, rotation, opacity)
feats = x.feats[x.layout[i]][:, v['range'][0]:v['range'][1]].reshape(-1, *v['shape']).flatten(0, 1)
feats = feats * self.rep_config['lr'][k] # Apply parameter-specific learning rate
setattr(representation, k, feats)
ret.append(representation)
return ret
def forward(self, x: sp.SparseTensor) -> List[Gaussian]:
"""
Forward pass through the decoder.
Args:
x: Input sparse tensor containing latent codes
Returns:
List of Gaussian representations ready for rendering
"""
h = super().forward(x) # Process through transformer blocks
h = h.type(x.dtype) # Ensure consistent dtype
h = h.replace(F.layer_norm(h.feats, h.feats.shape[-1:])) # Apply layer normalization
h = self.out_layer(h) # Project to final output dimensions
return self.to_representation(h) # Convert to Gaussian representations
class ElasticSLatGaussianDecoder(SparseTransformerElasticMixin, SparseTransformerBase):
"""
Slat VAE Gaussian decoder with elastic memory management.
Used for training with low VRAM by dynamically managing memory allocations
and using efficient sparse operations.
"""
pass
@@ -0,0 +1,222 @@
"""
Mesh Decoder Module for Structured Latent VAE
This file implements a mesh-based decoder for the structured latent variational autoencoder (SLAT VAE).
It contains specialized sparse neural network components that transform latent representations into
3D mesh structures through a series of sparse convolutions and subdivisions.
The module implements:
1. SparseSubdivideBlock3d - A block that subdivides sparse tensors to increase resolution
2. SLatMeshDecoder - Main decoder that transforms latent codes into 3D meshes
3. ElasticSLatMeshDecoder - Memory-efficient version for low VRAM environments
"""
from typing import *
import torch
import torch.nn as nn
import torch.nn.functional as F
import numpy as np
from ...modules.utils import zero_module, convert_module_to_f16, convert_module_to_f32
from ...modules import sparse as sp
from .base import SparseTransformerBase
from ...representations import MeshExtractResult
from ...representations.mesh import SparseFeatures2Mesh
from ..sparse_elastic_mixin import SparseTransformerElasticMixin
class SparseSubdivideBlock3d(nn.Module):
"""
A 3D subdivide block that can subdivide the sparse tensor.
This block increases the resolution of sparse tensors by a factor of 2,
and optionally changes the number of channels.
Args:
channels: channels in the inputs and outputs.
resolution: the current resolution of the sparse tensor.
out_channels: if specified, the number of output channels.
num_groups: the number of groups for the group norm.
"""
def __init__(
self,
channels: int,
resolution: int,
out_channels: Optional[int] = None,
num_groups: int = 32
):
super().__init__()
self.channels = channels
self.resolution = resolution
self.out_resolution = resolution * 2
self.out_channels = out_channels or channels
# Normalization and activation before subdivision
self.act_layers = nn.Sequential(
sp.SparseGroupNorm32(num_groups, channels),
sp.SparseSiLU()
)
# Subdivision operator that doubles the resolution
self.sub = sp.SparseSubdivide()
# Post-subdivision processing with residual connection
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}")),
)
# Skip connection that handles potential channel dimension changes
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}")
def forward(self, x: sp.SparseTensor) -> sp.SparseTensor:
"""
Apply the block to a Tensor, conditioned on a timestep embedding.
Args:
x: an [N x C x ...] Tensor of features.
Returns:
an [N x C x ...] Tensor of outputs with doubled resolution.
"""
h = self.act_layers(x)
h = self.sub(h) # Double the resolution
x = self.sub(x) # Also subdivide the input for skip connection
h = self.out_layers(h)
h = h + self.skip_connection(x) # Add skip connection
return h
class SLatMeshDecoder(SparseTransformerBase):
"""
Structured Latent Mesh Decoder that transforms latent codes into 3D meshes.
Uses sparse transformers followed by upsampling blocks to generate high-resolution
features that are then converted to meshes.
"""
def __init__(
self,
resolution: int,
model_channels: int,
latent_channels: int,
num_blocks: int,
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,
representation_config: dict = None,
):
# Initialize the transformer backbone
super().__init__(
in_channels=latent_channels,
model_channels=model_channels,
num_blocks=num_blocks,
num_heads=num_heads,
num_head_channels=num_head_channels,
mlp_ratio=mlp_ratio,
attn_mode=attn_mode,
window_size=window_size,
pe_mode=pe_mode,
use_fp16=use_fp16,
use_checkpoint=use_checkpoint,
qk_rms_norm=qk_rms_norm,
)
self.resolution = resolution
self.rep_config = representation_config
# Mesh extractor to convert features to mesh representation
self.mesh_extractor = SparseFeatures2Mesh(res=self.resolution*4, use_color=self.rep_config.get('use_color', False))
self.out_channels = self.mesh_extractor.feats_channels
# Upsampling blocks that progressively increase resolution
self.upsample = nn.ModuleList([
SparseSubdivideBlock3d(
channels=model_channels,
resolution=resolution,
out_channels=model_channels // 4
),
SparseSubdivideBlock3d(
channels=model_channels // 4,
resolution=resolution * 2,
out_channels=model_channels // 8
)
])
# Final layer to map features to mesh attributes
self.out_layer = sp.SparseLinear(model_channels // 8, self.out_channels)
self.initialize_weights()
if use_fp16:
self.convert_to_fp16()
def initialize_weights(self) -> None:
"""Initialize model weights, with special handling for output layers."""
super().initialize_weights()
# Zero-out output layers for stable training
nn.init.constant_(self.out_layer.weight, 0)
nn.init.constant_(self.out_layer.bias, 0)
def convert_to_fp16(self) -> None:
"""
Convert the torso of the model to float16 for memory efficiency.
"""
super().convert_to_fp16()
self.upsample.apply(convert_module_to_f16)
def convert_to_fp32(self) -> None:
"""
Convert the torso of the model back to float32 for precision.
"""
super().convert_to_fp32()
self.upsample.apply(convert_module_to_f32)
def to_representation(self, x: sp.SparseTensor) -> List[MeshExtractResult]:
"""
Convert a batch of network outputs to 3D mesh representations.
Args:
x: The [N x * x C] sparse tensor output by the network.
Returns:
list of mesh representation results, one per batch item
"""
ret = []
for i in range(x.shape[0]):
mesh = self.mesh_extractor(x[i], training=self.training)
ret.append(mesh)
return ret
def forward(self, x: sp.SparseTensor) -> List[MeshExtractResult]:
"""
Process latent codes through the decoder and extract meshes.
Args:
x: Input sparse tensor of latent codes
Returns:
List of extracted mesh representations
"""
h = super().forward(x) # Process through transformer blocks
for block in self.upsample:
h = block(h) # Progressively increase resolution
h = h.type(x.dtype)
h = self.out_layer(h) # Final projection to mesh features
return self.to_representation(h) # Convert features to meshes
class ElasticSLatMeshDecoder(SparseTransformerElasticMixin, SLatMeshDecoder):
"""
Structured Latent Mesh Decoder with elastic memory management.
This variant uses elastic sparse tensor operations to reduce memory usage
during training, making it suitable for environments with limited VRAM.
"""
pass
@@ -0,0 +1,156 @@
"""
This file implements radiance field decoders for Structured Latent VAE models.
The main class SLatRadianceFieldDecoder is a sparse transformer-based decoder that
transforms latent codes into sparse representations of 3D scenes (Strivec representation).
It also includes an elastic memory version (ElasticSLatRadianceFieldDecoder) for low VRAM training.
"""
from typing import *
import torch
import torch.nn as nn
import torch.nn.functional as F
import numpy as np
from ...modules import sparse as sp
from .base import SparseTransformerBase
from ...representations import Strivec
from ..sparse_elastic_mixin import SparseTransformerElasticMixin
class SLatRadianceFieldDecoder(SparseTransformerBase):
"""
A sparse transformer-based decoder for converting latent codes to radiance field representations.
This decoder processes sparse tensors through transformer blocks and outputs parameters for Strivec representation.
"""
def __init__(
self,
resolution: int, # Resolution of the output 3D grid
model_channels: int, # Number of channels in the model's hidden layers
latent_channels: int, # Number of channels in the latent code
num_blocks: int, # Number of transformer blocks
num_heads: Optional[int] = None, # Number of attention heads
num_head_channels: Optional[int] = 64, # Channels per attention head
mlp_ratio: float = 4, # Ratio for MLP hidden dimension
attn_mode: Literal["full", "shift_window", "shift_sequence", "shift_order", "swin"] = "swin", # Attention mode
window_size: int = 8, # Size of local attention window
pe_mode: Literal["ape", "rope"] = "ape", # Positional encoding mode
use_fp16: bool = False, # Whether to use half precision
use_checkpoint: bool = False, # Whether to use gradient checkpointing
qk_rms_norm: bool = False, # Whether to normalize query and key
representation_config: dict = None, # Configuration for output representation
):
# Initialize the base sparse transformer
super().__init__(
in_channels=latent_channels,
model_channels=model_channels,
num_blocks=num_blocks,
num_heads=num_heads,
num_head_channels=num_head_channels,
mlp_ratio=mlp_ratio,
attn_mode=attn_mode,
window_size=window_size,
pe_mode=pe_mode,
use_fp16=use_fp16,
use_checkpoint=use_checkpoint,
qk_rms_norm=qk_rms_norm,
)
self.resolution = resolution
self.rep_config = representation_config
self._calc_layout() # Calculate the output layout
# Final layer to project features to the output representation
self.out_layer = sp.SparseLinear(model_channels, self.out_channels)
self.initialize_weights()
if use_fp16:
self.convert_to_fp16()
def initialize_weights(self) -> None:
"""
Initialize the weights of the model.
Zero-initializes the output layer for better training stability.
"""
super().initialize_weights()
# Zero-out output layers for better training stability
nn.init.constant_(self.out_layer.weight, 0)
nn.init.constant_(self.out_layer.bias, 0)
def _calc_layout(self) -> None:
"""
Calculate the output tensor layout for the Strivec representation.
Defines the shapes and sizes of different components and their positions in the output tensor.
"""
self.layout = {
'trivec': {'shape': (self.rep_config['rank'], 3, self.rep_config['dim']), 'size': self.rep_config['rank'] * 3 * self.rep_config['dim']},
'density': {'shape': (self.rep_config['rank'],), 'size': self.rep_config['rank']},
'features_dc': {'shape': (self.rep_config['rank'], 1, 3), 'size': self.rep_config['rank'] * 3},
}
# Calculate the range (start, end) indices for each component in the output tensor
start = 0
for k, v in self.layout.items():
v['range'] = (start, start + v['size'])
start += v['size']
self.out_channels = start
def to_representation(self, x: sp.SparseTensor) -> List[Strivec]:
"""
Convert a batch of network outputs to 3D representations.
Args:
x: The [N x * x C] sparse tensor output by the network.
Returns:
list of Strivec representations, one per batch item
"""
ret = []
for i in range(x.shape[0]):
# Create a new Strivec representation
representation = Strivec(
sh_degree=0,
resolution=self.resolution,
aabb=[-0.5, -0.5, -0.5, 1, 1, 1], # Axis-aligned bounding box
rank=self.rep_config['rank'],
dim=self.rep_config['dim'],
device='cuda',
)
representation.density_shift = 0.0
# Set position from sparse coordinates (normalized to [0,1])
representation.position = (x.coords[x.layout[i]][:, 1:].float() + 0.5) / self.resolution
# Set depth (octree level) based on resolution
representation.depth = torch.full((representation.position.shape[0], 1), int(np.log2(self.resolution)), dtype=torch.uint8, device='cuda')
# Extract each component from the output features according to the layout
for k, v in self.layout.items():
setattr(representation, k, x.feats[x.layout[i]][:, v['range'][0]:v['range'][1]].reshape(-1, *v['shape']))
# Add 1 to trivec for stability (prevent zero vectors)
representation.trivec = representation.trivec + 1
ret.append(representation)
return ret
def forward(self, x: sp.SparseTensor) -> List[Strivec]:
"""
Forward pass through the decoder.
Args:
x: Input sparse tensor containing latent codes
Returns:
List of Strivec representations
"""
# Pass through transformer backbone
h = super().forward(x)
h = h.type(x.dtype)
# Layer normalization on feature dimension
h = h.replace(F.layer_norm(h.feats, h.feats.shape[-1:]))
# Final projection to output features
h = self.out_layer(h)
# Convert network output to Strivec representations
return self.to_representation(h)
class ElasticSLatRadianceFieldDecoder(SparseTransformerElasticMixin, SLatRadianceFieldDecoder):
"""
Slat VAE Radiance Field Decoder with elastic memory management.
Used for training with low VRAM by dynamically managing memory allocation
and performing operations in chunks when needed.
"""
pass
@@ -0,0 +1,138 @@
"""
Structured Latent Variable Encoder Module
----------------------------------------
This file defines encoder classes for the Structured Latent Variable Autoencoder (SLatVAE).
It contains implementations for the sparse transformer-based encoder that maps input
features to a latent distribution, as well as a memory-efficient elastic version.
The encoder follows a variational approach, outputting means and log variances for
the latent space representation.
"""
from typing import *
import torch
import torch.nn as nn
import torch.nn.functional as F
from ...modules import sparse as sp
from .base import SparseTransformerBase
from ..sparse_elastic_mixin import SparseTransformerElasticMixin
class SLatEncoder(SparseTransformerBase):
"""
Sparse Latent Variable Encoder that uses transformer architecture to encode
sparse data into a latent distribution.
"""
def __init__(
self,
resolution: int,
in_channels: int,
model_channels: int,
latent_channels: int,
num_blocks: int,
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,
):
"""
Initialize the Sparse Latent Encoder.
Args:
resolution: Input data resolution
in_channels: Number of input feature channels
model_channels: Number of internal model feature channels
latent_channels: Dimension of the latent space
num_blocks: Number of transformer blocks
num_heads: Number of attention heads (optional)
num_head_channels: Channels per attention head if num_heads is None
mlp_ratio: Expansion ratio for MLP in transformer blocks
attn_mode: Type of attention mechanism to use
window_size: Size of attention windows if using windowed attention
pe_mode: Positional encoding mode (absolute or relative)
use_fp16: Whether to use half-precision floating point
use_checkpoint: Whether to use gradient checkpointing
qk_rms_norm: Whether to apply RMS normalization to query and key
"""
super().__init__(
in_channels=in_channels,
model_channels=model_channels,
num_blocks=num_blocks,
num_heads=num_heads,
num_head_channels=num_head_channels,
mlp_ratio=mlp_ratio,
attn_mode=attn_mode,
window_size=window_size,
pe_mode=pe_mode,
use_fp16=use_fp16,
use_checkpoint=use_checkpoint,
qk_rms_norm=qk_rms_norm,
)
self.resolution = resolution
# Output layer projects to twice the latent dimension (for mean and logvar)
self.out_layer = sp.SparseLinear(model_channels, 2 * latent_channels)
self.initialize_weights()
if use_fp16:
self.convert_to_fp16()
def initialize_weights(self) -> None:
"""
Initialize model weights with special handling for output layer.
The output layer weights are initialized to zero to stabilize training.
"""
super().initialize_weights()
# Zero-out output layers for better training stability
nn.init.constant_(self.out_layer.weight, 0)
nn.init.constant_(self.out_layer.bias, 0)
def forward(self, x: sp.SparseTensor, sample_posterior=True, return_raw=False):
"""
Forward pass through the encoder.
Args:
x: Input sparse tensor
sample_posterior: Whether to sample from posterior or return mean
return_raw: Whether to return mean and logvar in addition to samples
Returns:
If return_raw is True:
- sampled latent variables, mean, and logvar
Otherwise:
- sampled latent variables only
"""
# Process through transformer blocks
h = super().forward(x)
h = h.type(x.dtype)
# Apply layer normalization to features
h = h.replace(F.layer_norm(h.feats, h.feats.shape[-1:]))
h = self.out_layer(h)
# Split output into mean and logvar components
mean, logvar = h.feats.chunk(2, dim=-1)
if sample_posterior:
# Reparameterization trick: z = mean + std * epsilon
std = torch.exp(0.5 * logvar)
z = mean + std * torch.randn_like(std)
else:
# Use mean directly without sampling
z = mean
z = h.replace(z)
if return_raw:
return z, mean, logvar
else:
return z
class ElasticSLatEncoder(SparseTransformerElasticMixin, SLatEncoder):
"""
SLat VAE encoder with elastic memory management.
Used for training with low VRAM by dynamically managing memory allocation
and performing operations with reduced memory footprint.
"""
pass
@@ -0,0 +1,36 @@
from typing import *
BACKEND = 'flash_attn'
DEBUG = False
def __from_env():
import os
global BACKEND
global DEBUG
env_attn_backend = os.environ.get('ATTN_BACKEND')
env_sttn_debug = os.environ.get('ATTN_DEBUG')
if env_attn_backend is not None and env_attn_backend in ['xformers', 'flash_attn', 'sdpa', 'naive']:
BACKEND = env_attn_backend
if env_sttn_debug is not None:
DEBUG = env_sttn_debug == '1'
print(f"[ATTENTION] Using backend: {BACKEND}")
__from_env()
def set_backend(backend: Literal['xformers', 'flash_attn']):
global BACKEND
BACKEND = backend
def set_debug(debug: bool):
global DEBUG
DEBUG = debug
from .full_attn import *
from .modules import *
@@ -0,0 +1,199 @@
"""
Full Attention Module
This file implements different versions of the Scaled Dot-Product Attention mechanism used in transformer models.
It provides a unified interface that supports multiple backend implementations (xformers, flash_attn,
PyTorch's native SDPA, or a naive implementation) while maintaining consistent input/output formats.
The module allows for flexible calling patterns with different tensor arrangements for queries, keys, and values.
"""
from typing import *
import torch
import math
from . import DEBUG, BACKEND # Import configuration variables
# Select the appropriate attention backend based on configuration
if BACKEND == 'xformers':
import xformers.ops as xops
elif BACKEND == 'flash_attn':
import flash_attn
elif BACKEND == 'sdpa':
from torch.nn.functional import scaled_dot_product_attention as sdpa
elif BACKEND == 'naive':
pass # Will use the naive implementation defined below
else:
raise ValueError(f"Unknown attention backend: {BACKEND}")
__all__ = [
'scaled_dot_product_attention', # Only expose this main function
]
def _naive_sdpa(q, k, v):
"""
Naive implementation of scaled dot product attention.
Args:
q (torch.Tensor): Query tensor
k (torch.Tensor): Key tensor
v (torch.Tensor): Value tensor
Returns:
torch.Tensor: Output attention tensor
Note:
This implementation follows the standard attention formula:
Attention(Q,K,V) = softmax(QK^T/sqrt(d_k))V
"""
q = q.permute(0, 2, 1, 3) # [N, H, L, C] - Reshape for batched matrix multiplication
k = k.permute(0, 2, 1, 3) # [N, H, L, C]
v = v.permute(0, 2, 1, 3) # [N, H, L, C]
scale_factor = 1 / math.sqrt(q.size(-1)) # Scale factor to prevent softmax saturation
attn_weight = q @ k.transpose(-2, -1) * scale_factor # Compute scaled dot product
attn_weight = torch.softmax(attn_weight, dim=-1) # Apply softmax to get attention weights
out = attn_weight @ v # Apply attention weights to values
out = out.permute(0, 2, 1, 3) # [N, L, H, C] - Restore original dimension order
return out
@overload
def scaled_dot_product_attention(qkv: torch.Tensor) -> torch.Tensor:
"""
Apply scaled dot product attention.
Args:
qkv (torch.Tensor): A [N, L, 3, H, C] tensor containing Qs, Ks, and Vs.
"""
...
@overload
def scaled_dot_product_attention(q: torch.Tensor, kv: torch.Tensor) -> torch.Tensor:
"""
Apply scaled dot product attention.
Args:
q (torch.Tensor): A [N, L, H, C] tensor containing Qs.
kv (torch.Tensor): A [N, L, 2, H, C] tensor containing Ks and Vs.
"""
...
@overload
def scaled_dot_product_attention(q: torch.Tensor, k: torch.Tensor, v: torch.Tensor) -> torch.Tensor:
"""
Apply scaled dot product attention.
Args:
q (torch.Tensor): A [N, L, H, Ci] tensor containing Qs.
k (torch.Tensor): A [N, L, H, Ci] tensor containing Ks.
v (torch.Tensor): A [N, L, H, Co] tensor containing Vs.
Note:
k and v are assumed to have the same coordinate map.
"""
...
def scaled_dot_product_attention(*args, **kwargs):
"""
Unified interface for scaled dot product attention with multiple calling patterns.
Supports three calling patterns:
1. Single combined QKV tensor: scaled_dot_product_attention(qkv)
2. Separate Q and combined KV: scaled_dot_product_attention(q, kv)
3. Separate Q, K, V tensors: scaled_dot_product_attention(q, k, v)
The function automatically selects the appropriate backend implementation
based on the BACKEND configuration.
"""
# Define expected argument names for each calling pattern
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}"
# Handle case 1: Single combined QKV tensor
if num_all_args == 1:
qkv = args[0] if len(args) > 0 else kwargs['qkv']
assert len(qkv.shape) == 5 and qkv.shape[2] == 3, f"Invalid shape for qkv, got {qkv.shape}, expected [N, L, 3, H, C]"
device = qkv.device
# Handle case 2: Separate Q and combined KV tensors
elif num_all_args == 2:
# print("handle case 2")
q = args[0] if len(args) > 0 else kwargs['q']
kv = args[1] if len(args) > 1 else kwargs['kv']
assert q.shape[0] == kv.shape[0], f"Batch size mismatch, got {q.shape[0]} and {kv.shape[0]}"
assert len(q.shape) == 4, f"Invalid shape for q, got {q.shape}, expected [N, L, H, C]"
assert len(kv.shape) == 5, f"Invalid shape for kv, got {kv.shape}, expected [N, L, 2, H, C]"
device = q.device
# Handle case 3: Separate Q, K, V tensors
elif num_all_args == 3:
# print("handle case 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 q.shape[0] == k.shape[0] == v.shape[0], f"Batch size mismatch, got {q.shape[0]}, {k.shape[0]}, and {v.shape[0]}"
assert len(q.shape) == 4, f"Invalid shape for q, got {q.shape}, expected [N, L, H, Ci]"
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]"
device = q.device
# print("no problem")
# Use xformers backend
if BACKEND == 'xformers':
if num_all_args == 1:
q, k, v = qkv.unbind(dim=2) # Split combined tensor into separate Q, K, V
elif num_all_args == 2:
k, v = kv.unbind(dim=2) # Split combined KV tensor
out = xops.memory_efficient_attention(q, k, v)
# Use Flash Attention backend
elif BACKEND == 'flash_attn':
# print("flash_attn")
if num_all_args == 1:
# print("case 1")
out = flash_attn.flash_attn_qkvpacked_func(qkv) # Use packed QKV format
elif num_all_args == 2:
# print("case 2")
out = flash_attn.flash_attn_kvpacked_func(q, kv) # Use packed KV format with separate Q
elif num_all_args == 3:
# print("case 3")
out = flash_attn.flash_attn_func(q, k, v) # Use fully separate Q, K, V
# Use PyTorch's native scaled dot product attention
elif BACKEND == 'sdpa':
# print("sdpa")
if num_all_args == 1:
# print("case 1")
q, k, v = qkv.unbind(dim=2) # Split combined tensor
elif num_all_args == 2:
# print("case 2")
k, v = kv.unbind(dim=2) # Split combined KV tensor
# PyTorch's SDPA expects tensors in format [N, H, L, C]
q = q.permute(0, 2, 1, 3) # [N, H, L, C]
k = k.permute(0, 2, 1, 3) # [N, H, L, C]
v = v.permute(0, 2, 1, 3) # [N, H, L, C]
out = sdpa(q, k, v) # [N, H, L, C]
out = out.permute(0, 2, 1, 3) # Convert back to [N, L, H, C]
# Use naive implementation
elif BACKEND == 'naive':
# print("naive")
if num_all_args == 1:
# print("case 1")
q, k, v = qkv.unbind(dim=2) # Split combined tensor
elif num_all_args == 2:
# print("case 2")
k, v = kv.unbind(dim=2) # Split combined KV tensor
out = _naive_sdpa(q, k, v) # Call the naive implementation
else:
raise ValueError(f"Unknown attention module: {BACKEND}")
# print("no problem")
return out
@@ -0,0 +1,256 @@
"""
This file contains attention mechanism implementations for the TRELLIS framework.
It provides various components needed for building transformer-based architectures,
including custom normalization, rotary position embeddings, and attention modules
with different configurations and optimizations.
"""
from typing import *
import torch
import torch.nn as nn
import torch.nn.functional as F
from .full_attn import scaled_dot_product_attention
class MultiHeadRMSNorm(nn.Module):
"""
Multi-head RMS normalization layer that applies per-head normalization.
This helps stabilize attention computations by normalizing query and key vectors.
Args:
dim (int): The dimensionality of each head
heads (int): Number of attention heads
"""
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: torch.Tensor) -> torch.Tensor:
"""
Apply RMS normalization along the last dimension.
Args:
x (torch.Tensor): Input tensor of shape [..., dim]
Returns:
torch.Tensor: Normalized tensor with the same shape
"""
return (F.normalize(x.float(), dim = -1) * self.gamma * self.scale).to(x.dtype)
class RotaryPositionEmbedder(nn.Module):
"""
Implements Rotary Position Embedding (RoPE), which encodes position information
into the query and key tensors through a rotation-based approach.
Args:
hidden_size (int): Size of the hidden dimension
in_channels (int): Number of input channels, defaults to 3
"""
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
# Calculate frequency bands on a log scale
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:
"""
Compute phase shifts based on position indices.
Args:
indices (torch.Tensor): Position indices
Returns:
torch.Tensor: Complex tensor containing phase information
"""
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:
"""
Apply rotary embeddings to the input tensor.
Args:
x (torch.Tensor): Input tensor
phases (torch.Tensor): Phase tensor from _get_phases
Returns:
torch.Tensor: Tensor with rotary embeddings applied
"""
x_complex = torch.view_as_complex(x.float().reshape(*x.shape[:-1], -1, 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]:
"""
Apply rotary position embeddings to query and key tensors.
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 None,
sequential indices will be used.
Returns:
Tuple[torch.Tensor, torch.Tensor]: Position-encoded query and key tensors
"""
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 MultiHeadAttention(nn.Module):
"""
Flexible multi-head attention implementation supporting both self-attention
and cross-attention with various optimizations.
Args:
channels (int): Number of input/output channels
num_heads (int): Number of attention heads
ctx_channels (Optional[int]): Number of context channels for cross-attention
type (str): Type of attention, either "self" or "cross"
attn_mode (str): Attention computation mode, either "full" or "windowed"
window_size (Optional[int]): Size of attention window if windowed mode is used
shift_window (Optional[Tuple[int, int, int]]): Shift amount for windowed attention
qkv_bias (bool): Whether to include bias in QKV projections
use_rope (bool): Whether to use rotary position embeddings
qk_rms_norm (bool): Whether to apply RMS normalization to Q and K
"""
def __init__(
self,
channels: int,
num_heads: int,
ctx_channels: Optional[int]=None,
type: Literal["self", "cross"] = "self",
attn_mode: Literal["full", "windowed"] = "full",
window_size: Optional[int] = None,
shift_window: Optional[Tuple[int, int, int]] = 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", "windowed"], f"Invalid attention mode: {attn_mode}"
assert type == "self" or attn_mode == "full", "Cross-attention only supports full attention"
if attn_mode == "windowed":
raise NotImplementedError("Windowed attention is not yet implemented")
self.channels = channels
self.head_dim = channels // num_heads
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_window = shift_window
self.use_rope = use_rope
self.qk_rms_norm = qk_rms_norm
# Initialize projection layers based on attention type
if self._type == "self":
# For self-attention, create a single QKV projection
self.to_qkv = nn.Linear(channels, channels * 3, bias=qkv_bias)
else:
# For cross-attention, create separate projections for query and key-value
self.to_q = nn.Linear(channels, channels, bias=qkv_bias)
self.to_kv = nn.Linear(self.ctx_channels, channels * 2, bias=qkv_bias)
# Optional RMS normalization for stabilizing attention
if self.qk_rms_norm:
self.q_rms_norm = MultiHeadRMSNorm(self.head_dim, num_heads)
self.k_rms_norm = MultiHeadRMSNorm(self.head_dim, num_heads)
# Output projection
self.to_out = nn.Linear(channels, channels)
# Optional rotary position embeddings
if use_rope:
self.rope = RotaryPositionEmbedder(channels)
def forward(self, x: torch.Tensor, context: Optional[torch.Tensor] = None, indices: Optional[torch.Tensor] = None) -> torch.Tensor:
"""
Apply multi-head attention to the input tensor.
Args:
x (torch.Tensor): Input tensor of shape [B, L, C]
context (Optional[torch.Tensor]): Context tensor for cross-attention
indices (Optional[torch.Tensor]): Position indices for rotary embeddings
Returns:
torch.Tensor: Output tensor of shape [B, L, C]
"""
B, L, C = x.shape
if self._type == "self":
# Self-attention path
qkv = self.to_qkv(x)
qkv = qkv.reshape(B, L, 3, self.num_heads, -1)
if self.use_rope:
# Apply rotary position embeddings if enabled
q, k, v = qkv.unbind(dim=2)
q, k = self.rope(q, k, indices)
qkv = torch.stack([q, k, v], dim=2)
if self.attn_mode == "full":
if self.qk_rms_norm:
# Apply RMS normalization to queries and keys if enabled
q, k, v = qkv.unbind(dim=2)
q = self.q_rms_norm(q)
k = self.k_rms_norm(k)
h = scaled_dot_product_attention(q, k, v)
else:
# Standard attention with combined QKV tensor
h = scaled_dot_product_attention(qkv)
elif self.attn_mode == "windowed":
raise NotImplementedError("Windowed attention is not yet implemented")
else:
# Cross-attention path
Lkv = context.shape[1]
q = self.to_q(x)
# print(f"context shape: {context.shape}")
kv = self.to_kv(context)
# print("reshape kv")
q = q.reshape(B, L, self.num_heads, -1)
kv = kv.reshape(B, Lkv, 2, self.num_heads, -1)
# print("unbind kv")
if self.qk_rms_norm:
# print("qk_rms_norm")
# Apply RMS normalization to queries and keys if enabled
q = self.q_rms_norm(q)
k, v = kv.unbind(dim=2)
# print("unbind kv2")
k = self.k_rms_norm(k)
# print("unbind kv3")
h = scaled_dot_product_attention(q, k, v)
# print("unbind kv4")
else:
# Standard cross-attention
# print("unbind kv2")
# print(kv.shape)
h = scaled_dot_product_attention(q, kv)
# print("unbind kv3")
# Reshape and project back to the original dimension
h = h.reshape(B, L, -1)
h = self.to_out(h)
return h
@@ -0,0 +1,25 @@
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 GroupNorm32(nn.GroupNorm):
"""
A GroupNorm layer that converts to float32 before the forward pass.
"""
def forward(self, x: torch.Tensor) -> torch.Tensor:
return super().forward(x.float()).type(x.dtype)
class ChannelLayerNorm32(LayerNorm32):
def forward(self, x: torch.Tensor) -> torch.Tensor:
DIM = x.dim()
x = x.permute(0, *range(2, DIM), 1).contiguous()
x = super().forward(x)
x = x.permute(0, DIM-1, *range(1, DIM-1)).contiguous()
return x
@@ -0,0 +1,102 @@
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',
'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'
}
__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
@@ -0,0 +1,4 @@
from .full_attn import *
from .serialized_attn import *
from .windowed_attn import *
from .modules import *
@@ -0,0 +1,215 @@
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)
@@ -0,0 +1,139 @@
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
from ...attention import RotaryPositionEmbedder
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)
@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":
qkv = self._linear(self.to_qkv, x)
qkv = self._fused_pre(qkv, num_fused=3)
if self.use_rope:
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:
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
@@ -0,0 +1,193 @@
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',
]
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)
@@ -0,0 +1,135 @@
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}'
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)
@@ -0,0 +1,691 @@
"""
Sparse Tensor Implementation for TRELLIS
----------------------------------------
This file implements a unified sparse tensor interface that supports multiple backends (torchsparse and spconv).
Sparse tensors are efficient representations of tensors where most values are zero, storing only non-zero values
and their coordinates. This is particularly useful for 3D point clouds and voxel grids in computer vision and
robotics applications where data is naturally sparse.
The main components of this file are:
- SparseTensor: Core class providing a unified API over different sparse tensor backends
- Utility functions for sparse tensor operations (concatenation, unbinding, broadcasting, etc.)
- Backend-agnostic arithmetic operations for sparse tensors
The implementation abstracts away backend-specific details to allow seamless switching between
torchsparse and spconv while maintaining a consistent interface.
"""
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 to avoid circular imports and improve startup time
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
# print(SparseTensorData)
# exit(0)
# Determine initialization method based on arguments (method 0: from tensors, method 1: from existing data)
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
self.old_index = None # Placeholder for old indices, if needed
if method_id == 0:
# Initialize from feature and coordinate tensors
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])
# Create backend-specific tensor representation
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:
# Initialize from existing sparse tensor data
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])
# Store metadata
self._shape = shape
self._layout = layout
self._scale = kwargs.get('scale', (1, 1, 1))
self._spatial_cache = kwargs.get('spatial_cache', {})
# Validate tensor properties in debug mode
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):
"""
Calculate the shape of the sparse tensor from features and coordinates.
This method determines the overall shape of the sparse tensor by examining:
- The batch dimension (from max coordinate value in first column + 1)
- The feature dimensions (from the feature tensor shape)
Args:
feats (torch.Tensor): Feature tensor of shape (N, C1, C2, ...)
coords (torch.Tensor): Coordinate tensor of shape (N, D+1) where
first column contains batch indices
Returns:
torch.Size: Shape of the sparse tensor as (batch_size, C1, C2, ...)
"""
shape = []
# First dimension is the batch size (max batch index + 1)
shape.append(coords[:, 0].max().item() + 1)
# Remaining dimensions match the feature tensor's dimensions
shape.extend([*feats.shape[1:]])
return torch.Size(shape)
def __cal_layout(self, coords, batch_size):
"""
Calculate the layout of each batch in the sparse tensor.
This method computes slice objects to efficiently index into specific batches
within the coordinate and feature tensors. It assumes that coordinates are
sorted by batch index (first column).
Algorithm:
1. Count how many elements belong to each batch using bincount
2. Calculate cumulative sums to find ending offsets for each batch
3. Create slice objects representing the range of indices for each batch
Args:
coords (torch.Tensor): Coordinate tensor with first column as batch indices
batch_size (int): Number of batches in the sparse tensor
Returns:
List[slice]: List of slice objects where layout[i] indexes all elements
belonging to batch i
"""
# Count number of points in each batch
seq_len = torch.bincount(coords[:, 0], minlength=batch_size)
# Calculate ending position of each batch
offset = torch.cumsum(seq_len, dim=0)
# Create slices for each batch from (end_prev_batch, end_current_batch)
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 the shape of the sparse tensor"""
return self._shape
def dim(self) -> int:
"""Return the number of dimensions of the sparse tensor"""
return len(self.shape)
@property
def layout(self) -> List[slice]:
"""Return the layout of each batch in the sparse tensor"""
return self._layout
@property
def feats(self) -> torch.Tensor:
"""Return the features tensor with backend-specific access"""
if BACKEND == 'torchsparse':
return self.data.F
elif BACKEND == 'spconv':
return self.data.features
@feats.setter
def feats(self, value: torch.Tensor):
"""Set the features tensor with backend-specific access"""
if BACKEND == 'torchsparse':
self.data.F = value
elif BACKEND == 'spconv':
self.data.features = value
@property
def coords(self) -> torch.Tensor:
"""Return the coordinates tensor with backend-specific access"""
if BACKEND == 'torchsparse':
return self.data.C
elif BACKEND == 'spconv':
return self.data.indices
@coords.setter
def coords(self, value: torch.Tensor):
"""Set the coordinates tensor with backend-specific access"""
if BACKEND == 'torchsparse':
self.data.C = value
elif BACKEND == 'spconv':
self.data.indices = value
@property
def dtype(self):
"""Return the data type of the sparse tensor's features"""
return self.feats.dtype
@property
def device(self):
"""Return the device of the sparse tensor's features"""
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':
"""
Move the sparse tensor to the specified device and/or change its data type.
Mimics the PyTorch tensor.to() method.
"""
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']
# print(self.feats)
# print(self.coords)
# print(SparseTensorData)
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):
"""Convert the sparse tensor to the specified data type"""
new_feats = self.feats.type(dtype)
return self.replace(new_feats)
def cpu(self) -> 'SparseTensor':
"""Move the sparse tensor to CPU memory"""
new_feats = self.feats.cpu()
new_coords = self.coords.cpu()
return self.replace(new_feats, new_coords)
def cuda(self) -> 'SparseTensor':
"""Move the sparse tensor to CUDA memory"""
new_feats = self.feats.cuda()
new_coords = self.coords.cuda()
return self.replace(new_feats, new_coords)
def half(self) -> 'SparseTensor':
"""Convert the sparse tensor to half precision"""
new_feats = self.feats.half()
return self.replace(new_feats)
def float(self) -> 'SparseTensor':
"""Convert the sparse tensor to single precision"""
new_feats = self.feats.float()
return self.replace(new_feats)
def detach(self) -> 'SparseTensor':
"""Detach the sparse tensor from the computation graph"""
new_coords = self.coords.detach()
new_feats = self.feats.detach()
return self.replace(new_feats, new_coords)
def dense(self) -> torch.Tensor:
"""Convert the sparse tensor to a dense tensor representation"""
if BACKEND == 'torchsparse':
return self.data.dense()
elif BACKEND == 'spconv':
return self.data.dense()
def reshape(self, *shape) -> 'SparseTensor':
"""Reshape the feature dimensions of the sparse tensor"""
new_feats = self.feats.reshape(self.feats.shape[0], *shape)
return self.replace(new_feats)
def unbind(self, dim: int) -> List['SparseTensor']:
"""Unbind the sparse tensor along the specified dimension"""
return sparse_unbind(self, dim)
def replace(self, feats: torch.Tensor, coords: Optional[torch.Tensor] = None) -> 'SparseTensor':
"""
Create a new sparse tensor with the specified features and optionally new coordinates.
Preserves other properties like stride, spatial range, and caches.
"""
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':
"""
Create a sparse tensor with uniform values within an axis-aligned bounding box.
Args:
aabb: [x_min, y_min, z_min, x_max, y_max, z_max] defining the bounding box
dim: (batch_size, feature_dim) tuple defining tensor dimensions
value: Value to fill the tensor with
dtype: Data type for features
device: Device to create the tensor on
"""
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:
"""Merge the spatial caches of two sparse tensors"""
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':
"""Negate the sparse tensor's values"""
return self.replace(-self.feats)
def __elemwise__(self, other: Union[torch.Tensor, 'SparseTensor'], op: callable) -> 'SparseTensor':
"""
Apply an elementwise operation between this sparse tensor and another tensor.
Handles broadcasting when necessary.
"""
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':
"""Add a tensor or value to this sparse tensor"""
return self.__elemwise__(other, torch.add)
def __radd__(self, other: Union[torch.Tensor, 'SparseTensor', float]) -> 'SparseTensor':
"""Add this sparse tensor to a tensor or value (reversed)"""
return self.__elemwise__(other, torch.add)
def __sub__(self, other: Union[torch.Tensor, 'SparseTensor', float]) -> 'SparseTensor':
"""Subtract a tensor or value from this sparse tensor"""
return self.__elemwise__(other, torch.sub)
def __rsub__(self, other: Union[torch.Tensor, 'SparseTensor', float]) -> 'SparseTensor':
"""Subtract this sparse tensor from a tensor or value (reversed)"""
return self.__elemwise__(other, lambda x, y: torch.sub(y, x))
def __mul__(self, other: Union[torch.Tensor, 'SparseTensor', float]) -> 'SparseTensor':
"""Multiply this sparse tensor by a tensor or value"""
return self.__elemwise__(other, torch.mul)
def __rmul__(self, other: Union[torch.Tensor, 'SparseTensor', float]) -> 'SparseTensor':
"""Multiply a tensor or value by this sparse tensor (reversed)"""
return self.__elemwise__(other, torch.mul)
def __truediv__(self, other: Union[torch.Tensor, 'SparseTensor', float]) -> 'SparseTensor':
"""Divide this sparse tensor by a tensor or value"""
return self.__elemwise__(other, torch.div)
def __rtruediv__(self, other: Union[torch.Tensor, 'SparseTensor', float]) -> 'SparseTensor':
"""Divide a tensor or value by this sparse tensor (reversed)"""
return self.__elemwise__(other, lambda x, y: torch.div(y, x))
def __getitem__(self, idx):
"""
Extract a batch or subset of batches from the sparse tensor.
Support for integer, slice, and tensor indexing.
"""
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 = []
old_index_list = []
for new_idx, old_idx in enumerate(idx):
coords.append(self.coords[self.layout[old_idx]].clone())
# print(f"slice: old index{old_idx}, new index: {new_idx}")
old_index_list.append(old_idx)
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()
self.old_index = old_index_list
return SparseTensor(feats=feats, coords=coords)
# def get_item_preserve_batch(self, idx):
# """
# Extract a batch or subset of batches from the sparse tensor without renumbering batch indices.
# Unlike __getitem__, this method preserves the original batch IDs in the coords tensor.
# Args:
# idx: Integer, slice, torch.Tensor, or tuple specifying which batch(es) to extract.
# When a tuple is provided, it's used for direct slicing of the underlying data.
# Returns:
# SparseTensor: A new sparse tensor with the selected batches and original batch IDs
# """
# if isinstance(idx, tuple):
# # Direct slice-based indexing
# coords_slice = self.coords[idx]
# feats_slice = self.feats[idx]
# return SparseTensor(feats=feats_slice, coords=coords_slice)
# 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 old_idx in idx:
# coords.append(self.coords[self.layout[old_idx]].clone())
# # Keep original batch ID (don't modify coords[:, 0])
# feats.append(self.feats[self.layout[old_idx]])
# coords = torch.cat(coords, dim=0).contiguous()
# feats = torch.cat(feats, dim=0).contiguous()
# # Create new SparseTensor with preserved batch IDs
# 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.
If key is None, return all caches for the current scale.
Otherwise, return the cache associated with the specified key.
"""
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 tensor to a sparse tensor along the batch dimension.
Args:
input (SparseTensor): Sparse tensor to broadcast to
other (torch.Tensor): Tensor to broadcast
Returns:
torch.Tensor: Broadcasted tensor matching the sparse tensor's layout
"""
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 (SparseTensor): Sparse tensor to operate on
other (torch.Tensor): 1D tensor to broadcast
op (callable): Operation to perform after broadcasting. Defaults to torch.add.
Returns:
SparseTensor: Result of the operation
"""
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 along a specified dimension.
This function handles two types of concatenation:
1. Batch concatenation (dim=0): Combines multiple sparse tensors by stacking their batches,
adjusting batch indices to maintain proper batch ordering.
2. Feature concatenation (dim>0): Combines features while maintaining the same coordinate structure,
useful for concatenating different feature channels for the same spatial locations.
Args:
inputs (List[SparseTensor]): List of sparse tensors to concatenate. All tensors must have
compatible shapes for the requested concatenation dimension.
dim (int): Dimension along which to concatenate.
- If 0, batches are concatenated (increasing batch indices)
- If >0, features are concatenated (same coordinates, more features)
Returns:
SparseTensor: A new sparse tensor with concatenated data
"""
if dim == 0:
# Concatenate batches - requires adjusting batch indices in coordinates
start = 0
coords = []
# Process each input sparse tensor
for input in inputs:
# Create a copy of coordinates to avoid modifying the original
current_coords = input.coords.clone()
# print("current coords", current_coords[:, 0])
# Adjust batch indices (first column of coordinates) to maintain proper batch ordering
# Each tensor's batch indices are offset by the sum of previous tensors' batch sizes
current_coords[:, 0] += start
# print("current coords", current_coords[:, 0])
# Add to coordinate list and update the batch counter
coords.append(current_coords)
# print("shape of input", input.shape)
start += input.shape[0]
# print("start number", start)
# Concatenate all adjusted coordinates into a single tensor
coords = torch.cat(coords, dim=0)
# Concatenate feature values in the same order as coordinates
feats = torch.cat([input.feats for input in inputs], dim=0)
# Create a new sparse tensor with combined coordinates and features
output = SparseTensor(
coords=coords,
feats=feats,
)
else:
# Concatenate features only - coordinates remain unchanged
# This works when all input tensors share the same coordinate structure
# but have different feature dimensions to combine
# Combine features along the specified dimension
feats = torch.cat([input.feats for input in inputs], dim=dim)
# Create new sparse tensor using the first input's coordinates
# but with the concatenated features
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
Returns:
List[SparseTensor]: List of sparse tensors, each representing a slice along the dimension
"""
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]
@@ -0,0 +1,21 @@
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 *
@@ -0,0 +1,80 @@
import torch
import torch.nn as nn
from .. import SparseTensor
from .. import DEBUG
from . import SPCONV_ALGO
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)
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,
)
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
@@ -0,0 +1,38 @@
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
@@ -0,0 +1,15 @@
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))
@@ -0,0 +1,35 @@
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))
@@ -0,0 +1,58 @@
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)
@@ -0,0 +1,110 @@
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'
)
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
@@ -0,0 +1,2 @@
from .blocks import *
from .modulated import *
@@ -0,0 +1,151 @@
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.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, mod: torch.Tensor, 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)
@@ -0,0 +1,166 @@
from typing import *
import torch
import torch.nn as nn
from ..basic import SparseTensor
from ..attention import SparseMultiHeadAttention, SerializeMode
from ...norm import LayerNorm32
from .blocks import SparseFeedForwardNet
class ModulatedSparseTransformerBlock(nn.Module):
"""
Sparse Transformer block (MSA + FFN) with adaptive layer norm conditioning.
"""
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,
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=False, 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,
)
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) -> 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.attn(h)
h = h * gate_msa
x = x + h
h = x.replace(self.norm2(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) -> SparseTensor:
if self.use_checkpoint:
return torch.utils.checkpoint.checkpoint(self._forward, x, mod, use_reentrant=False)
else:
return self._forward(x, mod)
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)
@@ -0,0 +1,48 @@
import torch
def pixel_shuffle_3d(x: torch.Tensor, scale_factor: int) -> torch.Tensor:
"""
3D pixel shuffle.
"""
B, C, H, W, D = x.shape
C_ = C // scale_factor**3
x = x.reshape(B, C_, scale_factor, scale_factor, scale_factor, H, W, D)
x = x.permute(0, 1, 5, 2, 6, 3, 7, 4)
x = x.reshape(B, C_, H*scale_factor, W*scale_factor, D*scale_factor)
return x
def patchify(x: torch.Tensor, patch_size: int):
"""
Patchify a tensor.
Args:
x (torch.Tensor): (N, C, *spatial) tensor
patch_size (int): Patch size
"""
DIM = x.dim() - 2
for d in range(2, DIM + 2):
assert x.shape[d] % patch_size == 0, f"Dimension {d} of input tensor must be divisible by patch size, got {x.shape[d]} and {patch_size}"
x = x.reshape(*x.shape[:2], *sum([[x.shape[d] // patch_size, patch_size] for d in range(2, DIM + 2)], []))
x = x.permute(0, 1, *([2 * i + 3 for i in range(DIM)] + [2 * i + 2 for i in range(DIM)]))
x = x.reshape(x.shape[0], x.shape[1] * (patch_size ** DIM), *(x.shape[-DIM:]))
return x
def unpatchify(x: torch.Tensor, patch_size: int):
"""
Unpatchify a tensor.
Args:
x (torch.Tensor): (N, C, *spatial) tensor
patch_size (int): Patch size
"""
DIM = x.dim() - 2
assert x.shape[1] % (patch_size ** DIM) == 0, f"Second dimension of input tensor must be divisible by patch size to unpatchify, got {x.shape[1]} and {patch_size ** DIM}"
x = x.reshape(x.shape[0], x.shape[1] // (patch_size ** DIM), *([patch_size] * DIM), *(x.shape[-DIM:]))
x = x.permute(0, 1, *(sum([[2 + DIM + i, 2 + i] for i in range(DIM)], [])))
x = x.reshape(x.shape[0], x.shape[1], *[x.shape[2 + 2 * i] * patch_size for i in range(DIM)])
return x
@@ -0,0 +1,2 @@
from .blocks import *
from .modulated import *
@@ -0,0 +1,182 @@
from typing import *
import torch
import torch.nn as nn
from ..attention import MultiHeadAttention
from ..norm import LayerNorm32
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 FeedForwardNet(nn.Module):
def __init__(self, channels: int, mlp_ratio: float = 4.0):
super().__init__()
self.mlp = nn.Sequential(
nn.Linear(channels, int(channels * mlp_ratio)),
nn.GELU(approximate="tanh"),
nn.Linear(int(channels * mlp_ratio), channels),
)
def forward(self, x: torch.Tensor) -> torch.Tensor:
return self.mlp(x)
class TransformerBlock(nn.Module):
"""
Transformer block (MSA + FFN).
"""
def __init__(
self,
channels: int,
num_heads: int,
mlp_ratio: float = 4.0,
attn_mode: Literal["full", "windowed"] = "full",
window_size: Optional[int] = None,
shift_window: Optional[int] = 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 = MultiHeadAttention(
channels,
num_heads=num_heads,
attn_mode=attn_mode,
window_size=window_size,
shift_window=shift_window,
qkv_bias=qkv_bias,
use_rope=use_rope,
qk_rms_norm=qk_rms_norm,
)
self.mlp = FeedForwardNet(
channels,
mlp_ratio=mlp_ratio,
)
def _forward(self, x: torch.Tensor) -> torch.Tensor:
h = self.norm1(x)
h = self.attn(h)
x = x + h
h = self.norm2(x)
h = self.mlp(h)
x = x + h
return x
def forward(self, x: torch.Tensor) -> torch.Tensor:
if self.use_checkpoint:
return torch.utils.checkpoint.checkpoint(self._forward, x, use_reentrant=False)
else:
return self._forward(x)
class TransformerCrossBlock(nn.Module):
"""
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", "windowed"] = "full",
window_size: Optional[int] = None,
shift_window: Optional[Tuple[int, int, int]] = 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.self_attn = MultiHeadAttention(
channels,
num_heads=num_heads,
type="self",
attn_mode=attn_mode,
window_size=window_size,
shift_window=shift_window,
qkv_bias=qkv_bias,
use_rope=use_rope,
qk_rms_norm=qk_rms_norm,
)
self.cross_attn = MultiHeadAttention(
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 = FeedForwardNet(
channels,
mlp_ratio=mlp_ratio,
)
def _forward(self, x: torch.Tensor, context: torch.Tensor):
h = self.norm1(x)
h = self.self_attn(h)
x = x + h
h = self.norm2(x)
h = self.cross_attn(h, context)
x = x + h
h = self.norm3(x)
h = self.mlp(h)
x = x + h
return x
def forward(self, x: torch.Tensor, 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)
@@ -0,0 +1,252 @@
"""
This file implements modulated transformer blocks for conditional generation.
These blocks extend standard transformer architectures by incorporating adaptive layer normalization (adaLN),
which modulates the transformer's behavior based on conditioning information.
The modulation is applied through shift and scale parameters derived from a condition vector,
allowing the model to adapt its processing to different inputs or conditions.
The file provides two main components:
1. ModulatedTransformerBlock: A standard transformer block with self-attention and FFN, modified with adaLN
2. ModulatedTransformerCrossBlock: An extended transformer block with self-attention, cross-attention, and FFN with adaLN
"""
from typing import *
import torch
import torch.nn as nn
from ..attention import MultiHeadAttention
from ..norm import LayerNorm32
from .blocks import FeedForwardNet
class ModulatedTransformerBlock(nn.Module):
"""
Transformer block (MSA + FFN) with adaptive layer norm conditioning.
This block combines multi-head self-attention with a feed-forward network,
and uses adaptive layer normalization to condition the processing on external information.
"""
def __init__(
self,
channels: int, # Number of input/output channels
num_heads: int, # Number of attention heads
mlp_ratio: float = 4.0, # Ratio determining MLP hidden dimension size
attn_mode: Literal["full", "windowed"] = "full", # Attention computation mode
window_size: Optional[int] = None, # Size of attention window if using windowed attention
shift_window: Optional[Tuple[int, int, int]] = None, # Parameters for shifted window attention
use_checkpoint: bool = False, # Whether to use gradient checkpointing to save memory
use_rope: bool = False, # Whether to use Rotary Position Embedding
qk_rms_norm: bool = False, # Whether to use RMS normalization for query and key
qkv_bias: bool = True, # Whether to use bias in QKV projection
share_mod: bool = False, # Whether to share modulation parameters externally
):
super().__init__()
self.use_checkpoint = use_checkpoint
self.share_mod = share_mod
# Layer normalization without affine parameters (will be modulated)
self.norm1 = LayerNorm32(channels, elementwise_affine=False, eps=1e-6)
self.norm2 = LayerNorm32(channels, elementwise_affine=False, eps=1e-6)
# Multi-head self-attention layer
self.attn = MultiHeadAttention(
channels,
num_heads=num_heads,
attn_mode=attn_mode,
window_size=window_size,
shift_window=shift_window,
qkv_bias=qkv_bias,
use_rope=use_rope,
qk_rms_norm=qk_rms_norm,
)
# Feed-forward network
self.mlp = FeedForwardNet(
channels,
mlp_ratio=mlp_ratio,
)
# Modulation network to generate adaptive parameters if not shared
if not share_mod:
self.adaLN_modulation = nn.Sequential(
nn.SiLU(),
nn.Linear(channels, 6 * channels, bias=True) # 6 channels: shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp
)
def _forward(self, x: torch.Tensor, mod: torch.Tensor) -> torch.Tensor:
"""
Internal forward function for the modulated transformer block.
Args:
x: Input tensor [batch, seq_len, channels]
mod: Modulation tensor [batch, channels]
Returns:
Processed tensor with same shape as input
"""
# Split modulation vector into shift, scale, and gate parameters for MSA and FFN
if self.share_mod:
# Use externally provided modulation parameters
shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = mod.chunk(6, dim=1)
else:
# Generate modulation parameters from the conditioning vector
shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = self.adaLN_modulation(mod).chunk(6, dim=1)
# Apply modulated self-attention
h = self.norm1(x) # Normalize
h = h * (1 + scale_msa.unsqueeze(1)) + shift_msa.unsqueeze(1) # Apply modulation
h = self.attn(h) # Self-attention
h = h * gate_msa.unsqueeze(1) # Apply gate
x = x + h # Residual connection
# Apply modulated feed-forward network
h = self.norm2(x) # Normalize
h = h * (1 + scale_mlp.unsqueeze(1)) + shift_mlp.unsqueeze(1) # Apply modulation
h = self.mlp(h) # Feed-forward
h = h * gate_mlp.unsqueeze(1) # Apply gate
x = x + h # Residual connection
return x
def forward(self, x: torch.Tensor, mod: torch.Tensor) -> torch.Tensor:
"""
Forward pass with optional gradient checkpointing to save memory.
Args:
x: Input tensor [batch, seq_len, channels]
mod: Modulation tensor [batch, channels]
Returns:
Processed tensor with same shape as input
"""
if self.use_checkpoint:
return torch.utils.checkpoint.checkpoint(self._forward, x, mod, use_reentrant=False)
else:
return self._forward(x, mod)
class ModulatedTransformerCrossBlock(nn.Module):
"""
Transformer cross-attention block (MSA + MCA + FFN) with adaptive layer norm conditioning.
This block extends the standard transformer block with an additional cross-attention
layer, allowing it to attend to a separate context input.
"""
def __init__(
self,
channels: int, # Number of input/output channels
ctx_channels: int, # Number of context channels
num_heads: int, # Number of attention heads
mlp_ratio: float = 4.0, # Ratio determining MLP hidden dimension size
attn_mode: Literal["full", "windowed"] = "full", # Attention computation mode
window_size: Optional[int] = None, # Size of attention window if using windowed attention
shift_window: Optional[Tuple[int, int, int]] = None, # Parameters for shifted window attention
use_checkpoint: bool = False, # Whether to use gradient checkpointing to save memory
use_rope: bool = False, # Whether to use Rotary Position Embedding
qk_rms_norm: bool = False, # Whether to use RMS normalization for query and key in self-attention
qk_rms_norm_cross: bool = False, # Whether to use RMS normalization for query and key in cross-attention
qkv_bias: bool = True, # Whether to use bias in QKV projection
share_mod: bool = False, # Whether to share modulation parameters externally
):
super().__init__()
self.use_checkpoint = use_checkpoint
self.share_mod = share_mod
# Layer normalizations
self.norm1 = LayerNorm32(channels, elementwise_affine=False, eps=1e-6) # For self-attention, will be modulated
self.norm2 = LayerNorm32(channels, elementwise_affine=True, eps=1e-6) # For cross-attention, standard normalization
self.norm3 = LayerNorm32(channels, elementwise_affine=False, eps=1e-6) # For FFN, will be modulated
# Self-attention layer
self.self_attn = MultiHeadAttention(
channels,
num_heads=num_heads,
type="self",
attn_mode=attn_mode,
window_size=window_size,
shift_window=shift_window,
qkv_bias=qkv_bias,
use_rope=use_rope,
qk_rms_norm=qk_rms_norm,
)
# Cross-attention layer
self.cross_attn = MultiHeadAttention(
channels,
ctx_channels=ctx_channels,
num_heads=num_heads,
type="cross",
attn_mode="full", # Cross-attention always uses full attention
qkv_bias=qkv_bias,
qk_rms_norm=qk_rms_norm_cross,
)
# Feed-forward network
self.mlp = FeedForwardNet(
channels,
mlp_ratio=mlp_ratio,
)
# Modulation network to generate adaptive parameters if not shared
if not share_mod:
self.adaLN_modulation = nn.Sequential(
nn.SiLU(),
nn.Linear(channels, 6 * channels, bias=True) # 6 channels: shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp
)
def _forward(self, x: torch.Tensor, mod: torch.Tensor, context: torch.Tensor):
"""
Internal forward function for the modulated transformer cross-attention block.
Args:
x: Input tensor [batch, seq_len, channels]
mod: Modulation tensor [batch, channels]
context: Context tensor for cross-attention [batch, context_len, ctx_channels]
Returns:
Processed tensor with same shape as input
"""
# Split modulation vector into shift, scale, and gate parameters for MSA and FFN
if self.share_mod:
# Use externally provided modulation parameters
shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = mod.chunk(6, dim=1)
else:
# Generate modulation parameters from the conditioning vector
shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = self.adaLN_modulation(mod).chunk(6, dim=1)
# Apply modulated self-attention
h = self.norm1(x) # Normalize
h = h * (1 + scale_msa.unsqueeze(1)) + shift_msa.unsqueeze(1) # Apply modulation
h = self.self_attn(h) # Self-attention
h = h * gate_msa.unsqueeze(1) # Apply gate
x = x + h # Residual connection
# Apply cross-attention (not modulated)
h = self.norm2(x) # Normalize
h = self.cross_attn(h, context) # Cross-attention with context
x = x + h # Residual connection
# Apply modulated feed-forward network
h = self.norm3(x) # Normalize
h = h * (1 + scale_mlp.unsqueeze(1)) + shift_mlp.unsqueeze(1) # Apply modulation
h = self.mlp(h) # Feed-forward
h = h * gate_mlp.unsqueeze(1) # Apply gate
x = x + h # Residual connection
return x
def forward(self, x: torch.Tensor, mod: torch.Tensor, context: torch.Tensor):
"""
Forward pass with optional gradient checkpointing to save memory.
Args:
x: Input tensor [batch, seq_len, channels]
mod: Modulation tensor [batch, channels]
context: Context tensor for cross-attention [batch, context_len, ctx_channels]
Returns:
Processed tensor with same shape as input
"""
if self.use_checkpoint:
return torch.utils.checkpoint.checkpoint(self._forward, x, mod, context, use_reentrant=False)
else:
return self._forward(x, mod, context)
@@ -0,0 +1,54 @@
import torch.nn as nn
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 scale_module(module, scale):
"""
Scale the parameters of a module and return it.
"""
for p in module.parameters():
p.detach().mul_(scale)
return module
def modulate(x, shift, scale):
return x * (1 + scale.unsqueeze(1)) + shift.unsqueeze(1)
@@ -0,0 +1,24 @@
from . import samplers
from .omnipart_image_to_parts import OmniPartImageTo3DPipeline
def from_pretrained(path: str):
"""
Load a pipeline from a model folder or a Hugging Face model hub.
Args:
path: The path to the model. Can be either local path or a Hugging Face model name.
"""
import os
import json
is_local = os.path.exists(f"{path}/pipeline.json")
if is_local:
config_file = f"{path}/pipeline.json"
else:
from huggingface_hub import hf_hub_download
config_file = hf_hub_download(path, "pipeline.json")
with open(config_file, 'r') as f:
config = json.load(f)
return globals()[config['name']].from_pretrained(path)

Some files were not shown because too many files have changed in this diff Show More