Add files via upload
This commit is contained in:
+4
-3
@@ -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
@@ -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,)
|
||||
|
||||
|
||||
@@ -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.
|
||||
|
||||

|
||||
|
||||
### 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
|
||||
[](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 ***
|
||||
|
||||
[](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
|
||||
```
|
||||
|
||||
|
||||

|
||||
|
||||

|
||||
|
||||
## 🌟 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.
|
||||
|
||||

|
||||
|
||||
### 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
|
||||
[](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 ***
|
||||
|
||||
[](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
|
||||
```
|
||||
|
||||
|
||||

|
||||
|
||||

|
||||
|
||||
## 🌟 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
|
||||
|
||||
@@ -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()
|
||||
@@ -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'
|
||||
]
|
||||
@@ -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'
|
||||
]
|
||||
|
||||
@@ -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}")
|
||||
@@ -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',
|
||||
|
||||
@@ -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',
|
||||
|
||||
@@ -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",
|
||||
]
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
Reference in New Issue
Block a user