init
This commit is contained in:
@@ -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
|
||||
@@ -0,0 +1,6 @@
|
||||
__pycache__/
|
||||
output/
|
||||
ckpt/
|
||||
.DS_Store
|
||||
tmp/
|
||||
debug_images/
|
||||
@@ -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
|
||||
|
||||
@@ -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.
|
||||
@@ -0,0 +1,99 @@
|
||||
# OmniPart: Part-Aware 3D Generation with Semantic Decoupling and Structural Cohesion [SIGGRAPH Asia 2025]
|
||||
|
||||
<div align="center">
|
||||
|
||||
[](https://omnipart.github.io/)
|
||||
[](https://arxiv.org/abs/2507.06165)
|
||||
[](https://huggingface.co/omnipart)
|
||||
[](https://huggingface.co/spaces/omnipart/OmniPart)
|
||||
|
||||
</div>
|
||||
|
||||

|
||||
|
||||
## 🔥 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
@@ -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)
|
||||
@@ -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 |
@@ -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
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
@@ -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
|
||||
+12
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
Reference in New Issue
Block a user