Add files via upload

This commit is contained in:
ImpactFrames
2025-01-13 18:16:49 +00:00
committed by GitHub
parent 93a98c5778
commit 07d9cbe2bb
17 changed files with 3135 additions and 416 deletions
+4 -3
View File
@@ -193,7 +193,7 @@ class IF_TrellisImageTo3D:
video_path = os.path.join(out_dir, f"{project_name}_preview.mp4")
imageio.mimsave(video_path, video, fps=fps)
full_video_path = os.path.abspath(video_path)
video_path = get_subpath_after_dir(full_video_path, "output")
video_path = os.path.abspath(video_path)
logger.info(f"Full video path: {full_video_path}, Processed video path: {video_path}")
if save_glb:
@@ -308,7 +308,7 @@ class IF_TrellisImageTo3D:
pipeline_params = self.get_pipeline_params(
seed, ss_sampling_steps, ss_guidance_strength,
slat_sampling_steps, slat_guidance_strength
slat_sampling_steps, slat_guidance_strength,
)
# Handle single vs multi mode differently
@@ -316,7 +316,8 @@ class IF_TrellisImageTo3D:
# Take just the first image regardless of how many were input
images = images[0:1]
pil_imgs = self.torch_to_pil_batch(images, masks)
outputs = model.run(pil_imgs[0], **pipeline_params)
outputs = model.run(pil_imgs[0],
**pipeline_params)
else:
# In multi mode, treat the whole list as a batch
pil_imgs = self.torch_to_pil_batch(images, masks)
+104 -204
View File
@@ -1,261 +1,161 @@
# IF_TrellisCheckpointLoader.py
import os
import sys
import importlib
import torch
import logging
import torch
import folder_paths
from huggingface_hub import hf_hub_download, snapshot_download
from pathlib import Path
import json
from trellis_model_manager import TrellisModelManager
from trellis.pipelines.trellis_image_to_3d import TrellisImageTo3DPipeline
from trellis.modules import set_attention_backend
from trellis.backend_config import (
set_attention_backend,
set_sparse_backend,
get_available_backends,
get_available_sparse_backends
)
from typing import Literal
from trellis.modules.attention_utils import enable_sage_attention, disable_sage_attention
from torchvision import transforms
logger = logging.getLogger("IF_Trellis")
def set_backend(backend: Literal['spconv', 'torchsparse']):
# Example helper if you wish to call the underlying global set_backend from trellis.modules.sparse:
from trellis.modules.sparse import set_backend as _set_sparse_backend
# Also handle spconv algo if desired, e.g. os.environ['SPCONV_ALGO'] = ...
_set_sparse_backend(backend)
class TrellisConfig:
"""Global configuration for Trellis"""
def __init__(self):
self.logger = logger
self.attention_backend = "sage"
self.spconv_algo = "implicit_gemm"
self.smooth_k = True
self.device = "cuda"
self.use_fp16 = True
# Added new configuration dictionary
self._config = {
"dinov2_size": "large", # Default model size
"dinov2_model": "dinov2_vitg14" # Default model name
}
# Added new methods
def get(self, key, default=None):
"""Get configuration value with fallback"""
return self._config.get(key, default)
def set(self, key, value):
"""Set configuration value"""
self._config[key] = value
def setup_environment(self):
"""Set up all environment variables and backends"""
import os
from trellis.modules import set_attention_backend
from trellis.modules.sparse import set_backend
# Set attention backend
set_attention_backend(self.attention_backend)
# Set smooth k for sage attention
os.environ['SAGEATTN_SMOOTH_K'] = '1' if self.smooth_k else '0'
# Set spconv algorithm
os.environ['SPCONV_ALGO'] = self.spconv_algo
# Always use spconv as backend for now
set_backend('spconv')
logger.info(f"Environment configured - Backend: spconv, "
f"Attention: {self.attention_backend}, "
f"Smooth K: {self.smooth_k}, "
f"SpConv Algo: {self.spconv_algo}")
# Global config instance
TRELLIS_CONFIG = TrellisConfig()
class IF_TrellisCheckpointLoader:
"""
Node to manage the loading of the TRELLIS model.
Follows ComfyUI conventions for model management.
Node to manage the loading of the TRELLIS model with lazy backend selection.
"""
def __init__(self):
self.logger = logger
self.model_manager = None
# Check for available devices
self.device = self._get_device()
def _get_device(self):
"""Determine the best available device."""
if torch.cuda.is_available():
return "cuda"
elif hasattr(torch.backends, 'mps') and torch.backends.mps.is_available():
return "mps"
return "cpu"
self.device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
# We might call these to figure out what's actually installed,
# if we want to populate UI dropdowns:
self.attn_backends = get_available_backends() # e.g. { 'xformers': True, 'flash_attn': False, ... }
self.sparse_backends = get_available_sparse_backends()# e.g. { 'spconv': True, 'torchsparse': True }
@classmethod
def INPUT_TYPES(cls):
"""Define input types with device-specific options."""
device_options = []
if torch.cuda.is_available():
device_options.append("cuda")
if hasattr(torch.backends, 'mps') and torch.backends.mps.is_available():
device_options.append("mps")
device_options.append("cpu")
# Filter only available backends
attn_backends = get_available_backends()
sparse_backends = get_available_sparse_backends()
# e.g. create a list of names that are True:
available_attn = [k for k, v in attn_backends.items() if v]
if not available_attn:
available_attn = ['sdpa'] # fallback
available_sparse = [k for k, v in sparse_backends.items() if v]
if not available_sparse:
available_sparse = ['spconv'] # fallback
return {
"required": {
"model_name": (["TRELLIS-image-large"],),
"dinov2_model": (["dinov2_vitl14_reg", "dinov2_vitg14_reg"], {"default": "dinov2_vitl14_reg", "tooltip": "Select the Dinov2 model to use for the image to 3D conversion. Smaller models work but better results with larger models."}),
"dinov2_model": (["dinov2_vitl14_reg"],
{"default": "dinov2_vitl14_reg",
"tooltip": "Select which Dinov2 model to use."}),
"use_fp16": ("BOOLEAN", {"default": True}),
"attn_backend": (["sage", "xformers", "flash_attn", "sdpa", "naive"], {"default": "sage", "tooltip": "Select the attention backend to use for the image to 3D conversion. Sage is experimental but faster"}),
"smooth_k": ("BOOLEAN", {"default": True, "tooltip": "Smooth k for sage attention. This is a hyperparameter that controls the smoothness of the attention distribution. It is a boolean value that determines whether to use smooth k or not. Smooth k is a hyperparameter that controls the smoothness of the attention distribution. It is a boolean value that determines whether to use smooth k or not."}),
"spconv_algo": (["implicit_gemm", "native"], {"default": "implicit_gemm", "tooltip": "Select the spconv algorithm to use for the image to 3D conversion. Implicit gemm is the best but slower. Native is the fastest but less accurate."}),
"main_device": (device_options, {"default": device_options[0]}),
#
# The user picks from the actually installed backends
#
"attn_backend": (available_attn,
{"default": "sdpa" if "sdpa" in available_attn else available_attn[0],
"tooltip": "Select attention backend."}),
"sparse_backend": (available_sparse,
{"default": "spconv" if "spconv" in available_sparse else available_sparse[0],
"tooltip": "Select sparse backend."}),
"spconv_algo": (["implicit_gemm", "native", "auto"],
{"default": "implicit_gemm",
"tooltip": "Spconv algorithm. 'implicit_gemm' is slower but more robust."}),
"smooth_k": ("BOOLEAN",
{"default": True,
"tooltip": "Smooth-k for SageAttention. Only relevant if attn_backend=sage."}),
},
}
RETURN_TYPES = ("TRELLIS_MODEL",)
RETURN_NAMES = ("model",)
FUNCTION = "load_model"
CATEGORY = "ImpactFrames💥🎞️/Trellis"
@classmethod
def _check_backend_availability(cls, backend: str) -> bool:
"""Check if a specific attention backend is available"""
try:
if backend == 'sage':
import sageattention
elif backend == 'xformers':
import xformers.ops
elif backend == 'flash_attn':
import flash_attn
elif backend in ['sdpa', 'naive']:
# These are always available in PyTorch
pass
else:
return False
return True
except ImportError:
return False
@classmethod
def _initialize_backend(cls, requested_backend: str = None) -> str:
"""Initialize attention backend with fallback logic"""
# Priority order for backends
backend_priority = ['sage', 'flash_attn', 'xformers', 'sdpa']
# If a specific backend is requested, try it first
if requested_backend:
if cls._check_backend_availability(requested_backend):
logger.info(f"Using requested attention backend: {requested_backend}")
return requested_backend
else:
logger.warning(f"Requested backend '{requested_backend}' not available, falling back")
# Try backends in priority order
for backend in backend_priority:
if cls._check_backend_availability(backend):
logger.info(f"Using attention backend: {backend}")
return backend
# Final fallback to SDPA
logger.info("All optimized attention backends unavailable, using PyTorch SDPA")
return 'sdpa'
def _setup_environment(self):
def _setup_environment(self, attn_backend: str, sparse_backend: str, spconv_algo: str, smooth_k: bool):
"""
Set up environment variables based on the global TRELLIS_CONFIG.
Set up environment variables and backends lazily.
This is the main difference: we call our new lazy set_*_backend funcs.
"""
import os
from trellis.modules import set_attention_backend
from trellis.modules.sparse import set_backend
from trellis.modules.sparse.conv import SPCONV_ALGO
# Try attention
success = set_attention_backend(attn_backend)
if not success:
self.logger.warning(f"Failed to set {attn_backend} or not installed, fallback to sdpa.")
# Set attention backend
os.environ['ATTN_BACKEND'] = TRELLIS_CONFIG.attention_backend
set_attention_backend(TRELLIS_CONFIG.attention_backend)
# Try sparse
success2 = set_sparse_backend(sparse_backend, spconv_algo)
if not success2:
self.logger.warning(f"Failed to set {sparse_backend} or not installed, fallback to default.")
# Set smooth k for sage attention
os.environ['SAGEATTN_SMOOTH_K'] = '1' if TRELLIS_CONFIG.smooth_k else '0'
# If user wants SageAttn smooth_k, we set environment var (if they'd want that):
os.environ['SAGEATTN_SMOOTH_K'] = '1' if smooth_k else '0'
# Set spconv algorithm
os.environ['SPCONV_ALGO'] = TRELLIS_CONFIG.spconv_algo
# Always use spconv as backend for now
set_backend('spconv')
def _initialize_transforms(self):
"""Initialize image transforms if needed."""
return transforms.Compose([
transforms.Normalize(
mean=[0.485, 0.456, 0.406],
std=[0.229, 0.224, 0.225]
)
])
logger.info(f"Environment configured - Backend: spconv, "
f"Attention: {TRELLIS_CONFIG.attention_backend}, "
f"Smooth K: {TRELLIS_CONFIG.smooth_k}, "
f"SpConv Algo: {TRELLIS_CONFIG.spconv_algo}")
def optimize_pipeline(self, pipeline, use_fp16=True, attn_backend='sage'):
"""Apply optimizations to the pipeline if available"""
if self.device == "cuda":
def _optimize_pipeline(self, pipeline, use_fp16: bool = True):
"""
Apply typical optimizations, half-precision, etc.
"""
if self.device.type == "cuda":
try:
if hasattr(pipeline, 'cuda'):
pipeline.cuda()
if use_fp16:
if hasattr(pipeline, 'enable_attention_slicing'):
pipeline.enable_attention_slicing()
pipeline.enable_attention_slicing(slice_size="auto")
if hasattr(pipeline, 'half'):
pipeline.half()
# Only enable xformers if using xformers backend
if attn_backend == 'xformers' and hasattr(pipeline, 'enable_xformers_memory_efficient_attention'):
pipeline.enable_xformers_memory_efficient_attention()
except Exception as e:
logger.warning(f"Some optimizations failed: {str(e)}")
logger.warning(f"Some pipeline optimizations failed: {str(e)}")
return pipeline
def load_model(self, model_name, dinov2_model="dinov2_vitg14", attn_backend="sage", use_fp16=True,
smooth_k=True, spconv_algo="implicit_gemm", main_device="cuda"):
"""Load and configure the TRELLIS model."""
def load_model(
self,
model_name: str,
dinov2_model: str = "dinov2_vitl14_reg",
attn_backend: str = "sdpa",
sparse_backend: str = "spconv",
spconv_algo: str = "implicit_gemm",
use_fp16: bool = True,
smooth_k: bool = True,
) -> tuple:
"""
Load and configure the TRELLIS pipeline.
This is typically the main function invoked by ComfyUI at node execution time.
"""
try:
# Update global config
TRELLIS_CONFIG.attention_backend = attn_backend
TRELLIS_CONFIG.spconv_algo = spconv_algo
TRELLIS_CONFIG.smooth_k = smooth_k
TRELLIS_CONFIG.device = main_device
TRELLIS_CONFIG.use_fp16 = use_fp16
TRELLIS_CONFIG.set("dinov2_model", dinov2_model)
# 1) Setup environment + backends
self._setup_environment(attn_backend, sparse_backend, spconv_algo, smooth_k)
# Set up environment
self._setup_environment()
# Configure attention backend
set_attention_backend(attn_backend)
if attn_backend == 'sage':
enable_sage_attention()
else:
disable_sage_attention()
# Get model path
# 2) Get model path
model_path = folder_paths.get_full_path("checkpoints", model_name)
if model_path is None:
model_path = os.path.join(folder_paths.models_dir, "checkpoints", model_name)
if not os.path.exists(model_path):
raise FileNotFoundError(f"Model not found: {model_path}")
# Create pipeline with specified dinov2 model
pipeline = TrellisImageTo3DPipeline.from_pretrained(model_path, dinov2_model=dinov2_model)
# Configure pipeline after loading
pipeline._device = torch.device(main_device)
pipeline.attention_backend = attn_backend
# Store configuration in pipeline
pipeline.config = {
'device': main_device,
'use_fp16': use_fp16,
'attention_backend': attn_backend,
'dinov2_model': dinov2_model,
'spconv_algo': spconv_algo,
'smooth_k': smooth_k
}
# 3) Create pipeline with the config
pipeline = TrellisImageTo3DPipeline.from_pretrained(
model_path,
dinov2_model=dinov2_model
)
pipeline._device = self.device # ensure pipeline uses our same device
# Apply optimizations
pipeline = self.optimize_pipeline(pipeline, use_fp16, attn_backend)
# 4) Apply optimizations
pipeline = self._optimize_pipeline(pipeline, use_fp16)
return (pipeline,)
+117 -117
View File
@@ -1,117 +1,117 @@
<h1 align="center">ComfyUI-IF_Trellis</h1>
<p align="center"><a href="https://arxiv.org/abs/2412.01506"><img src='https://img.shields.io/badge/arXiv-Paper-red?logo=arxiv&logoColor=white' alt='arXiv'></a>
<a href='https://trellis3d.github.io'><img src='https://img.shields.io/badge/Project_Page-Website-green?logo=googlechrome&logoColor=white' alt='Project Page'></a>
<a href='https://huggingface.co/spaces/JeffreyXiang/TRELLIS'><img src='https://img.shields.io/badge/%F0%9F%A4%97%20Hugging%20Face-Live_Demo-blue'></a>
</p>
ComfyUI TRELLIS is a large 3D asset generation in various formats, such as Radiance Fields, 3D Gaussians, and meshes. The cornerstone of TRELLIS is a unified Structured LATent (SLAT) representation that allows decoding to different output formats and Rectified Flow Transformers tailored for SLAT as the powerful backbones.
![teaser (1)](https://github.com/user-attachments/assets/6eee56bd-0936-44a5-b843-be4e87be649f)
### Prerequisites
- **System**: The original code is currently tested only on **Linux**. For windows setup, you may refer to [#3](https://github.com/microsoft/TRELLIS/issues/3)
- (This comfyui node is following this suggeted installation steps It works but you need to follow the steps as described in part one and two of the guide).
- Windows users need to use the Win_requirements.txt.
- Linux Users (Not tested Yet) use the linux_requirements.txt but I am still testing if it works with comfy in Linux I just added the original repo requirements.
- **Hardware**: This repo improve memmory mangement with the PR made by Amorano you will need NVIDIA GPU with at least 8GB.
- **Software**:
- The [CUDA Toolkit](https://developer.nvidia.com/cuda-toolkit-archive) is needed to compile certain submodules. The code has been tested with CUDA versions 11.8 and 12.2. This repo use **CUDA 12.4**.
- [Conda](https://docs.anaconda.com/miniconda/install/#quick-command-line-install) is recommended for managing dependencies.
- Python version 3.8 or higher is required.
Give unrestricted script access to powershell so venv can work:
- Open an administrator powershell window
- Type `Set-ExecutionPolicy Unrestricted` and answer A
- Close admin powershell window
### Installation Steps
1. Clone the repo:
```
cd ComfyUI/Custom_nodes
git clone --recurse-submodules https://github.com/if-ai/ComfyUI-IF_Trellis.git
```
## MUST HAVE `--recurse-submodules`
ONLY tested on windows but it should work easier in Linux without any issues or needing such specific stuffs.
***NOT tested or compatible with PORTABLE comfy embeded python env***
watch quick overview of setting the env if needed
[![Watch the video](https://img.youtube.com/vi/-vEpuYL9I3g/hqdefault.jpg)](https://www.youtube.com/watch?v=-vEpuYL9I3g)
You need to set up the environment first
follow this guide for the first part
Set the VSCode Cpp Envirronment as in the guide
[Installing Triton and Sage Attention Flash Attention](https://ko-fi.com/post/Installing-Triton-and-Sage-Attention-Flash-Attenti-P5P8175434)
Setting up ComfyUI with the Xformers, flash attention, Sage-attention(Optional Recommended for Hunyuan and other Video models)
*** You can Also Try This Other Guide ***
[![Watch this guide if you need extra help](https://img.youtube.com/vi/FjNfDsX-jR0/hqdefault.jpg)](https://www.youtube.com/watch?v=FjNfDsX-jR0)
[How to Run Micromamba in ComfyUI with Triton, SAGE Attention, Flash Attention, and xFormers for ComfyUI 3D](https://comfyuiblog.com/how-to-run-micromamba-in-comfyui-with-triton-sage-attention-flash-attention-and-x-formers-for-comfyui-3d/)
<!-- Installation -->
## 📦 Second part of the installation
Activate youur comfy environment
```
(gen) PS D:\ComfyUI\custom_nodes\ComfyUI-IF_Trellis> micromamba activate gen
```
If you haven't set your vars or for some reason it can't compile some of this specially `nvdiffrast`
it doesn't hurt if you do it again now.
```
cmd.exe /c "C:\Program Files (x86)\Microsoft Visual Studio\2022\BuildTools\VC\Auxiliary\Build\vcvarsall.bat" x64 "&&" powershell
```
You will see some message like this:
**********************************************************************
** Visual Studio 2019 Developer Command Prompt v16.11.41
** Copyright (c) 2021 Microsoft Corporation
**********************************************************************
[vcvarsall.bat] Environment initialized for: 'x64'
Windows PowerShell
Copyright (C) Microsoft Corporation. All rights reserved.
```
pip install -r win_requirements.txt
```
```bash
pip install git+https://github.com/EasternJournalist/utils3d.git@9a4eb15e4021b67b12c460c7057d642626897ec8
New-Item -ItemType Directory -Force -Path C:\tmp\extensions
git clone --recurse-submodules https://github.com/JeffreyXiang/diffoctreerast.git C:\tmp\extensions\diffoctreerast
pip install C:\tmp\extensions\diffoctreerast
git clone https://github.com/autonomousvision/mip-splatting.git C:\tmp\extensions\mip-splatting
pip install C:\tmp\extensions\mip-splatting\submodules\diff-gaussian-rasterization\
pip install kaolin -f https://nvidia-kaolin.s3.us-east-2.amazonaws.com/torch-2.4.0_cu121.html
git clone https://github.com/NVlabs/nvdiffrast.git C:\tmp\extensions\nvdiffrast
pip install C:\tmp\extensions\nvdiffrast
```
![blender_oWFuSjA6yq](https://github.com/user-attachments/assets/7d5fb61a-f2f6-4000-ab32-8555d5c6b7da)
![thorium_jtbGswydSr](https://github.com/user-attachments/assets/15f1d538-faa1-4c79-86f9-05f9b74ae794)
## 🌟 Features
- **High Quality**: It produces diverse 3D assets at high quality with intricate shape and texture details.
- **Versatility**: It takes text or image prompts and can generate various final 3D representations including but not limited to *Radiance Fields*, *3D Gaussians*, and *meshes*, accommodating diverse downstream requirements.
- **Flexible Editing**: It allows for easy editings of generated 3D assets, such as generating variants of the same object or local editing of the 3D asset.
<!-- TODO List -->
## 🚧 TODO List
- [x] Release comfyUI-IF_Trellis
- [x] 3D Viewport
- [x] OPT mode
- [x] Multiviews
- [x] Sage attn
- [ ] Installation
- [ ] colab notebook
- [ ] MarchingCubes mode
- [ ] Exposing Texture paranmeters
<h1 align="center">ComfyUI-IF_Trellis</h1>
<p align="center"><a href="https://arxiv.org/abs/2412.01506"><img src='https://img.shields.io/badge/arXiv-Paper-red?logo=arxiv&logoColor=white' alt='arXiv'></a>
<a href='https://trellis3d.github.io'><img src='https://img.shields.io/badge/Project_Page-Website-green?logo=googlechrome&logoColor=white' alt='Project Page'></a>
<a href='https://huggingface.co/spaces/JeffreyXiang/TRELLIS'><img src='https://img.shields.io/badge/%F0%9F%A4%97%20Hugging%20Face-Live_Demo-blue'></a>
</p>
ComfyUI TRELLIS is a large 3D asset generation in various formats, such as Radiance Fields, 3D Gaussians, and meshes. The cornerstone of TRELLIS is a unified Structured LATent (SLAT) representation that allows decoding to different output formats and Rectified Flow Transformers tailored for SLAT as the powerful backbones.
![teaser (1)](https://github.com/user-attachments/assets/6eee56bd-0936-44a5-b843-be4e87be649f)
### Prerequisites
- **System**: The original code is currently tested only on **Linux**. For windows setup, you may refer to [#3](https://github.com/microsoft/TRELLIS/issues/3)
- (This comfyui node is following this suggeted installation steps It works but you need to follow the steps as described in part one and two of the guide).
- Windows users need to use the Win_requirements.txt.
- Linux Users (Not tested Yet) use the linux_requirements.txt but I am still testing if it works with comfy in Linux I just added the original repo requirements.
- **Hardware**: This repo improve memmory mangement with the PR made by Amorano you will need NVIDIA GPU with at least 8GB.
- **Software**:
- The [CUDA Toolkit](https://developer.nvidia.com/cuda-toolkit-archive) is needed to compile certain submodules. The code has been tested with CUDA versions 11.8 and 12.2. This repo use **CUDA 12.4**.
- [Conda](https://docs.anaconda.com/miniconda/install/#quick-command-line-install) is recommended for managing dependencies.
- Python version 3.8 or higher is required.
Give unrestricted script access to powershell so venv can work:
- Open an administrator powershell window
- Type `Set-ExecutionPolicy Unrestricted` and answer A
- Close admin powershell window
### Installation Steps
1. Clone the repo:
```
cd ComfyUI/Custom_nodes
git clone --recurse-submodules https://github.com/if-ai/ComfyUI-IF_Trellis.git
```
## MUST HAVE `--recurse-submodules`
ONLY tested on windows but it should work easier in Linux without any issues or needing such specific stuffs.
***NOT tested or compatible with PORTABLE comfy embeded python env***
watch quick overview of setting the env if needed
[![Watch the video](https://img.youtube.com/vi/-vEpuYL9I3g/hqdefault.jpg)](https://www.youtube.com/watch?v=-vEpuYL9I3g)
You need to set up the environment first
follow this guide for the first part
Set the VSCode Cpp Envirronment as in the guide
[Installing Triton and Sage Attention Flash Attention](https://ko-fi.com/post/Installing-Triton-and-Sage-Attention-Flash-Attenti-P5P8175434)
Setting up ComfyUI with the Xformers, flash attention, Sage-attention(Optional Recommended for Hunyuan and other Video models)
*** You can Also Try This Other Guide ***
[![Watch this guide if you need extra help](https://img.youtube.com/vi/FjNfDsX-jR0/hqdefault.jpg)](https://www.youtube.com/watch?v=FjNfDsX-jR0)
[How to Run Micromamba in ComfyUI with Triton, SAGE Attention, Flash Attention, and xFormers for ComfyUI 3D](https://comfyuiblog.com/how-to-run-micromamba-in-comfyui-with-triton-sage-attention-flash-attention-and-x-formers-for-comfyui-3d/)
<!-- Installation -->
## 📦 Second part of the installation
Activate youur comfy environment
```
(gen) PS D:\ComfyUI\custom_nodes\ComfyUI-IF_Trellis> micromamba activate gen
```
If you haven't set your vars or for some reason it can't compile some of this specially `nvdiffrast`
it doesn't hurt if you do it again now.
```
cmd.exe /c "C:\Program Files (x86)\Microsoft Visual Studio\2022\BuildTools\VC\Auxiliary\Build\vcvarsall.bat" x64 "&&" powershell
```
You will see some message like this:
**********************************************************************
** Visual Studio 2019 Developer Command Prompt v16.11.41
** Copyright (c) 2021 Microsoft Corporation
**********************************************************************
[vcvarsall.bat] Environment initialized for: 'x64'
Windows PowerShell
Copyright (C) Microsoft Corporation. All rights reserved.
```
pip install -r win_requirements.txt
```
```bash
pip install git+https://github.com/EasternJournalist/utils3d.git@9a4eb15e4021b67b12c460c7057d642626897ec8
New-Item -ItemType Directory -Force -Path C:\tmp\extensions
git clone --recurse-submodules https://github.com/JeffreyXiang/diffoctreerast.git C:\tmp\extensions\diffoctreerast
pip install C:\tmp\extensions\diffoctreerast
git clone https://github.com/autonomousvision/mip-splatting.git C:\tmp\extensions\mip-splatting
pip install C:\tmp\extensions\mip-splatting\submodules\diff-gaussian-rasterization\
pip install kaolin -f https://nvidia-kaolin.s3.us-east-2.amazonaws.com/torch-2.4.0_cu121.html
git clone https://github.com/NVlabs/nvdiffrast.git C:\tmp\extensions\nvdiffrast
pip install C:\tmp\extensions\nvdiffrast
```
![blender_oWFuSjA6yq](https://github.com/user-attachments/assets/7d5fb61a-f2f6-4000-ab32-8555d5c6b7da)
![thorium_jtbGswydSr](https://github.com/user-attachments/assets/15f1d538-faa1-4c79-86f9-05f9b74ae794)
## 🌟 Features
- **High Quality**: It produces diverse 3D assets at high quality with intricate shape and texture details.
- **Versatility**: It takes text or image prompts and can generate various final 3D representations including but not limited to *Radiance Fields*, *3D Gaussians*, and *meshes*, accommodating diverse downstream requirements.
- **Flexible Editing**: It allows for easy editings of generated 3D assets, such as generating variants of the same object or local editing of the 3D asset.
<!-- TODO List -->
## 🚧 TODO List
- [x] Release comfyUI-IF_Trellis
- [x] 3D Viewport
- [x] OPT mode
- [x] Multiviews
- [x] Sage attn
- [ ] Installation
- [ ] colab notebook
- [ ] MarchingCubes mode
- [ ] Exposing Texture paranmeters
+225
View File
@@ -0,0 +1,225 @@
# trellis/backend_config.py
from typing import *
import os
import logging
import importlib
# Global variables
BACKEND = 'spconv' # Default sparse backend
DEBUG = False # Debug mode flag
ATTN = 'sdpa' # Default attention backend
SPCONV_ALGO = 'implicit_gemm' # Default algorithm
def get_spconv_algo() -> str:
"""Get current spconv algorithm."""
global SPCONV_ALGO
return SPCONV_ALGO
def set_spconv_algo(algo: Literal['implicit_gemm', 'native', 'auto']) -> bool:
"""Set spconv algorithm with validation."""
global SPCONV_ALGO
if algo not in ['implicit_gemm', 'native', 'auto']:
logger.warning(f"Invalid spconv algorithm: {algo}. Must be 'implicit_gemm', 'native', or 'auto'")
return False
SPCONV_ALGO = algo
os.environ['SPCONV_ALGO'] = algo
logger.info(f"Set spconv algorithm to: {algo}")
return True
logger = logging.getLogger(__name__)
def _try_import_xformers() -> bool:
try:
import xformers.ops
return True
except ImportError:
return False
def _try_import_flash_attn() -> bool:
try:
import flash_attn
return True
except ImportError:
return False
def _try_import_sageattention() -> bool:
try:
import torch.nn.functional as F
from sageattention import sageattn
F.scaled_dot_product_attention = sageattn
#import sageattention
return True
except ImportError:
return False
def _try_import_spconv() -> bool:
try:
import spconv
return True
except ImportError:
return False
def _try_import_torchsparse() -> bool:
try:
import torchsparse
return True
except ImportError:
return False
def get_available_backends() -> Dict[str, bool]:
"""Return dict of available attention backends and their status"""
return {
'xformers': _try_import_xformers(),
'flash_attn': _try_import_flash_attn(),
'sage': _try_import_sageattention(),
'naive': True,
'sdpa': True # Always available with PyTorch >= 2.0
}
def get_available_sparse_backends() -> Dict[str, bool]:
"""Return dict of available sparse backends and their status"""
return {
'spconv': _try_import_spconv(),
'torchsparse': _try_import_torchsparse()
}
def get_attention_backend() -> str:
"""Get current attention backend"""
global ATTN
return ATTN
def get_sparse_backend() -> str:
"""Get current sparse backend"""
global BACKEND
return BACKEND
def get_debug_mode() -> bool:
"""Get current debug mode status"""
global DEBUG
return DEBUG
def __from_env():
"""Initialize settings from environment variables"""
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_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', 'sage', 'sdpa', 'naive']:
ATTN = env_sparse_attn
os.environ['SPARSE_ATTN_BACKEND'] = env_sparse_attn
os.environ['ATTN_BACKEND'] = env_sparse_attn
logger.info(f"[SPARSE] Backend: {BACKEND}, Attention: {ATTN}")
def set_backend(backend: Literal['spconv', 'torchsparse']) -> bool:
"""Set sparse backend with validation"""
global BACKEND
backend = backend.lower().strip()
logger.info(f"Setting sparse backend to: {backend}")
if backend == 'spconv':
try:
import spconv
BACKEND = 'spconv'
os.environ['SPARSE_BACKEND'] = 'spconv'
return True
except ImportError:
logger.warning("spconv not available")
return False
elif backend == 'torchsparse':
try:
import torchsparse
BACKEND = 'torchsparse'
os.environ['SPARSE_BACKEND'] = 'torchsparse'
return True
except ImportError:
logger.warning("torchsparse not available")
return False
return False
def set_sparse_backend(backend: Literal['spconv', 'torchsparse'], algo: str = None) -> bool:
"""Alias for set_backend for backwards compatibility
Parameters:
backend: The sparse backend to use
algo: The algorithm to use (only relevant for spconv backend)
"""
# Call set_backend first
result = set_backend(backend)
# If algorithm is provided and backend was set successfully
if algo is not None and result:
set_spconv_algo(algo)
return result
def set_debug(debug: bool):
"""Set debug mode"""
global DEBUG
DEBUG = debug
if debug:
os.environ['SPARSE_DEBUG'] = '1'
else:
os.environ['SPARSE_DEBUG'] = '0'
def set_attn(attn: Literal['xformers', 'flash_attn', 'sage', 'sdpa', 'naive']) -> bool:
"""Set attention backend with validation"""
global ATTN
attn = attn.lower().strip()
logger.info(f"Setting attention backend to: {attn}")
if attn == 'xformers' and _try_import_xformers():
ATTN = 'xformers'
os.environ['SPARSE_ATTN_BACKEND'] = 'xformers'
os.environ['ATTN_BACKEND'] = 'xformers'
return True
elif attn == 'flash_attn' and _try_import_flash_attn():
ATTN = 'flash_attn'
os.environ['SPARSE_ATTN_BACKEND'] = 'flash_attn'
os.environ['ATTN_BACKEND'] = 'flash_attn'
return True
elif attn == 'sage' and _try_import_sageattention():
ATTN = 'sage'
os.environ['SPARSE_ATTN_BACKEND'] = 'sage'
os.environ['ATTN_BACKEND'] = 'sage'
return True
elif attn == 'sdpa':
ATTN = 'sdpa'
os.environ['SPARSE_ATTN_BACKEND'] = 'sdpa'
os.environ['ATTN_BACKEND'] = 'sdpa'
return True
elif attn == 'naive':
ATTN = 'naive'
os.environ['SPARSE_ATTN_BACKEND'] = 'naive'
os.environ['ATTN_BACKEND'] = 'naive'
return True
logger.warning(f"Attention backend {attn} not available")
return False
# Add alias for backwards compatibility
def set_attention_backend(backend: Literal['xformers', 'flash_attn', 'sage', 'sdpa']) -> bool:
"""Alias for set_attn for backwards compatibility"""
return set_attn(backend)
# Initialize from environment variables on module import
__from_env()
+7 -9
View File
@@ -1,18 +1,16 @@
from .attention_utils import enable_sage_attention, disable_sage_attention
#from .attention_utils import enable_sage_attention, disable_sage_attention
from .attention import (
set_attention_backend,
get_attention_op,
ATTN_BACKEND,
scaled_dot_product_attention,
BACKEND,
DEBUG,
MultiHeadAttention,
RotaryPositionEmbedder
)
__all__ = [
'enable_sage_attention',
'disable_sage_attention',
'set_attention_backend',
'get_attention_op',
'ATTN_BACKEND',
'scaled_dot_product_attention',
'BACKEND',
'DEBUG',
'MultiHeadAttention',
'RotaryPositionEmbedder'
]
+22 -49
View File
@@ -1,65 +1,38 @@
import os
import logging
from typing import Literal
from trellis.backend_config import (
get_attention_backend,
get_debug_mode,
)
import logging
logger = logging.getLogger(__name__)
# Global settings
ATTN_BACKEND = 'flash_attn' # Default backend
DEBUG = False # Add DEBUG flag here
BACKEND = 'flash_attn' # Change default to a valid backend
#ATTN = get_attention_backend()
BACKEND = get_attention_backend()
DEBUG = get_debug_mode()
def set_attention_backend(backend: Literal['xformers', 'flash_attn', 'sdpa', 'sage', 'naive']):
"""Set the global attention backend"""
global ATTN_BACKEND, BACKEND # Also update BACKEND when setting attention backend
if backend not in ['xformers', 'flash_attn', 'sdpa', 'sage', 'naive']:
raise ValueError(f"Unsupported attention backend: {backend}")
ATTN_BACKEND = backend
BACKEND = backend # Keep BACKEND in sync with ATTN_BACKEND
os.environ['ATTN_BACKEND'] = backend
logger.info(f"[ATTENTION] Set backend to: {backend}")
def get_attention_op():
"""Get the appropriate attention implementation"""
if ATTN_BACKEND == 'xformers':
try:
import xformers.ops
return xformers.ops.memory_efficient_attention
except ImportError:
logger.warning("xformers not available, falling back to naive attention")
return None
elif ATTN_BACKEND == 'flash_attn':
try:
from flash_attn import flash_attn_func
return flash_attn_func
except ImportError:
logger.warning("flash_attn not available, falling back to naive attention")
return None
elif ATTN_BACKEND == 'sage':
try:
from sageattention import sageattn
return sageattn
except ImportError:
logger.warning("sageattention not available, falling back to naive attention")
return None
elif ATTN_BACKEND == 'sdpa':
import torch
if hasattr(torch.nn.functional, 'scaled_dot_product_attention'):
return torch.nn.functional.scaled_dot_product_attention
def __from_env():
"""Read current backend configuration"""
#global ATTN
global BACKEND
global DEBUG
# Fallback to naive implementation
logger.warning("Using naive attention implementation")
return None
# Get current settings from central config
#ATTN =
BACKEND = get_attention_backend()
DEBUG = get_debug_mode()
print(f"[ATTENTION] Using backend: {BACKEND}")
# Import attention modules after defining globals
from .modules import MultiHeadAttention, RotaryPositionEmbedder
from .full_attn import scaled_dot_product_attention
__all__ = [
'set_attention_backend',
'get_attention_op',
'ATTN_BACKEND',
'DEBUG',
'scaled_dot_product_attention',
'BACKEND',
'DEBUG',
'MultiHeadAttention',
'RotaryPositionEmbedder'
]
+26 -8
View File
@@ -1,18 +1,36 @@
from typing import *
import torch
import math
from . import DEBUG, BACKEND
from trellis.backend_config import (
get_attention_backend,
get_debug_mode,
get_available_backends
)
import logging
if BACKEND == 'xformers':
logger = logging.getLogger(__name__)
# Get configuration from central config
BACKEND = get_attention_backend()
DEBUG = get_debug_mode()
# Get available backends and import if active
available_backends = get_available_backends()
if BACKEND == "xformers" and available_backends['xformers']:
import xformers.ops as xops
elif BACKEND == 'flash_attn':
elif BACKEND == "flash_attn" and available_backends['flash_attn']:
import flash_attn
elif BACKEND == 'sdpa':
elif BACKEND == "sage" and available_backends['sage']:
import torch.nn.functional as F
from sageattention import sageattn
F.scaled_dot_product_attention = sageattn
elif BACKEND == "sdpa":
from torch.nn.functional import scaled_dot_product_attention as sdpa
elif BACKEND == 'naive':
pass
elif BACKEND == "naive":
from torch.nn.functional import scaled_dot_product_attention as naive
else:
raise ValueError(f"Unknown attention backend: {BACKEND}")
raise ValueError(f"Unknown attention module: {BACKEND}")
__all__ = [
@@ -137,4 +155,4 @@ def scaled_dot_product_attention(*args, **kwargs):
else:
raise ValueError(f"Unknown attention module: {BACKEND}")
return out
return out
File diff suppressed because it is too large Load Diff
@@ -2,3 +2,30 @@ from .full_attn import *
from .serialized_attn import *
from .windowed_attn import *
from .modules import *
import os
import logging
from typing import Literal
from trellis.backend_config import (
get_attention_backend,
get_debug_mode,
)
import logging
logger = logging.getLogger(__name__)
#ATTN = get_attention_backend()
ATTN = get_attention_backend()
DEBUG = get_debug_mode()
def __from_env():
"""Read current backend configuration"""
#global ATTN
global ATTN
global DEBUG
# Get current settings from central config
#ATTN =
ATTN = get_attention_backend()
DEBUG = get_debug_mode()
print(f"[ATTENTION] sparse backend: {ATTN}")
+30 -4
View File
@@ -1,15 +1,41 @@
#trellis\modules\sparse\attention\full_attn.py
from typing import *
import torch
import math
from .. import SparseTensor
from .. import DEBUG, ATTN
from trellis.backend_config import (
get_attention_backend,
get_debug_mode,
get_available_backends
)
import logging
if ATTN == 'xformers':
logger = logging.getLogger(__name__)
# Get configuration from central config
ATTN = get_attention_backend()
DEBUG = get_debug_mode()
# Get available backends and import if active
available_backends = get_available_backends()
if ATTN == "xformers" and available_backends['xformers']:
import xformers.ops as xops
elif ATTN == 'flash_attn':
elif ATTN == "flash_attn" and available_backends['flash_attn']:
import flash_attn
elif ATTN == "sage" and available_backends['sage']:
import torch.nn.functional as F
from sageattention import sageattn
F.scaled_dot_product_attention = sageattn
elif ATTN == "sdpa":
from torch.nn.functional import scaled_dot_product_attention as sdpa
elif ATTN == "naive":
from torch.nn.functional import scaled_dot_product_attention as naive
else:
raise ValueError(f"Unknown attention module: {ATTN}")
# Log the active backend
logger.info(f"Using attention backend: {ATTN}")
__all__ = [
'sparse_scaled_dot_product_attention',
@@ -212,4 +238,4 @@ def sparse_scaled_dot_product_attention(*args, **kwargs):
if s is not None:
return s.replace(out)
else:
return out.reshape(N, L, H, -1)
return out.reshape(N, L, H, -1)
@@ -3,15 +3,39 @@ from enum import Enum
import torch
import math
from .. import SparseTensor
from .. import DEBUG, ATTN
from trellis.backend_config import (
get_attention_backend,
get_debug_mode,
get_available_backends
)
import logging
if ATTN == 'xformers':
logger = logging.getLogger(__name__)
# Get configuration from central config
ATTN = get_attention_backend()
DEBUG = get_debug_mode()
# Get available backends and import if active
available_backends = get_available_backends()
if ATTN == "xformers" and available_backends['xformers']:
import xformers.ops as xops
elif ATTN == 'flash_attn':
elif ATTN == "flash_attn" and available_backends['flash_attn']:
import flash_attn
elif ATTN == "sage" and available_backends['sage']:
import torch.nn.functional as F
from sageattention import sageattn
F.scaled_dot_product_attention = sageattn
elif ATTN == "sdpa":
from torch.nn.functional import scaled_dot_product_attention as sdpa
elif ATTN == "naive":
from torch.nn.functional import scaled_dot_product_attention as naive
else:
raise ValueError(f"Unknown attention module: {ATTN}")
# Log the active backend
logger.info(f"Using attention backend: {ATTN}")
__all__ = [
'sparse_serialized_scaled_dot_product_self_attention',
@@ -190,4 +214,4 @@ def sparse_serialized_scaled_dot_product_self_attention(
qkv_coords = qkv_coords[bwd_indices]
assert torch.equal(qkv_coords, qkv.coords), "SparseWindowedScaledDotProductSelfAttention: coordinate mismatch"
return qkv.replace(out)
return qkv.replace(out)
@@ -1,16 +1,43 @@
from typing import *
import torch
from enum import Enum
import math
import os
import logging
import torch
from .. import SparseTensor
from .. import DEBUG, ATTN
from trellis.backend_config import (
get_attention_backend,
get_debug_mode,
get_available_backends
)
import logging
if ATTN == 'xformers':
logger = logging.getLogger(__name__)
# Get configuration from central config
ATTN = get_attention_backend()
DEBUG = get_debug_mode()
# Get available backends and import if active
available_backends = get_available_backends()
if ATTN == "xformers" and available_backends['xformers']:
import xformers.ops as xops
elif ATTN == 'flash_attn':
elif ATTN == "flash_attn" and available_backends['flash_attn']:
import flash_attn
elif ATTN == "sage" and available_backends['sage']:
import torch.nn.functional as F
from sageattention import sageattn
F.scaled_dot_product_attention = sageattn
elif ATTN == "sdpa":
from torch.nn.functional import scaled_dot_product_attention as sdpa
elif ATTN == "naive":
from torch.nn.functional import scaled_dot_product_attention as naive
else:
raise ValueError(f"Unknown attention module: {ATTN}")
# Log the active backend
logger.info(f"Using attention backend: {ATTN}")
__all__ = [
'sparse_windowed_scaled_dot_product_self_attention',
+8 -1
View File
@@ -1,10 +1,17 @@
from typing import *
import torch
import torch.nn as nn
from . import BACKEND, DEBUG
from trellis.backend_config import get_debug_mode, get_spconv_algo, get_sparse_backend
SparseTensorData = None # Lazy import
DEBUG = get_debug_mode()
SPCONV_ALGO = get_spconv_algo()
BACKEND = get_sparse_backend()
__all__ = [
'SparseTensor',
'sparse_batch_broadcast',
+10 -3
View File
@@ -1,7 +1,9 @@
from .. import BACKEND
import os
import logging
from trellis.backend_config import get_sparse_backend, get_spconv_algo
SPCONV_ALGO = 'auto' # 'auto', 'implicit_gemm', 'native'
BACKEND = get_sparse_backend()
SPCONV_ALGO = get_spconv_algo()
def __from_env():
import os
@@ -19,4 +21,9 @@ if BACKEND == 'torchsparse':
from .conv_torchsparse import *
elif BACKEND == 'spconv':
from .conv_spconv import *
__all__ = [
"SparseConv3d",
"SparseInverseConv3d",
]
+8 -3
View File
@@ -2,12 +2,17 @@ import torch
import torch.nn as nn
from .. import SparseTensor
from .. import DEBUG
from . import SPCONV_ALGO
from trellis.backend_config import get_debug_mode, get_spconv_algo, get_sparse_backend
# Get configuration from central config
DEBUG = get_debug_mode()
SPCONV_ALGO = get_spconv_algo()
BACKEND = get_sparse_backend()
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():
if BACKEND == 'spconv':
import spconv.pytorch as spconv
algo = None
if SPCONV_ALGO == 'native':
@@ -52,7 +57,7 @@ class SparseConv3d(nn.Module):
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():
if BACKEND == 'spconv':
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)
@@ -1,12 +1,17 @@
import torch
import torch.nn as nn
from .. import SparseTensor
from trellis.backend_config import get_debug_mode, get_spconv_algo, get_sparse_backend
# Get configuration from central config
DEBUG = get_debug_mode()
SPCONV_ALGO = get_spconv_algo()
BACKEND = get_sparse_backend()
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():
if BACKEND == 'torchsparse':
import torchsparse
self.conv = torchsparse.nn.Conv3d(in_channels, out_channels, kernel_size, stride, 0, dilation, bias)
@@ -22,7 +27,7 @@ class SparseConv3d(nn.Module):
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():
if BACKEND == 'torchsparse':
import torchsparse
self.conv = torchsparse.nn.Conv3d(in_channels, out_channels, kernel_size, stride, 0, dilation, bias, transposed=True)
+20 -5
View File
@@ -172,12 +172,17 @@ class TrellisModelManager:
return models
def load_dinov2(self, model_name: str):
"""Load DINOv2 model with device and precision management"""
"""Load DINOv2 model with device, precision, and attention backend management"""
try:
# Get use_fp16 from config dict or object
# Get configuration values
use_fp16 = (self.config.get('use_fp16', True)
if isinstance(self.config, dict)
else getattr(self.config, 'use_fp16', True))
# Get attention backend from config
attention_backend = (self.config.get('attention_backend', 'default')
if isinstance(self.config, dict)
else getattr(self.config, 'attention_backend', 'default'))
# Try to load from local path first
model_path = folder_paths.get_full_path("classifiers", f"{model_name}.pth")
@@ -185,8 +190,11 @@ class TrellisModelManager:
if model_path is None:
print(f"Downloading {model_name} from torch hub...")
try:
# Load model architecture
model = torch.hub.load('facebookresearch/dinov2', model_name, pretrained=True)
# Load model architecture with specified attention backend
model = torch.hub.load('facebookresearch/dinov2', model_name,
pretrained=True,
force_reload=False,
trust_repo=True)
# Save model for future use
save_dir = os.path.join(folder_paths.models_dir, "classifiers")
@@ -203,7 +211,10 @@ class TrellisModelManager:
else:
# Load from local path
print(f"Loading DINOv2 model from {model_path}")
model = torch.hub.load('facebookresearch/dinov2', model_name, pretrained=False)
model = torch.hub.load('facebookresearch/dinov2', model_name,
pretrained=False,
force_reload=False,
trust_repo=True)
model.load_state_dict(torch.load(model_path))
# Move model to specified device and apply precision settings
@@ -211,6 +222,10 @@ class TrellisModelManager:
if use_fp16:
model = model.half()
# Set attention backend if specified in config
if hasattr(model, 'set_attention_backend') and attention_backend != 'default':
model.set_attention_backend(attention_backend)
model.eval()
return model