How to Implement Custom Image Conditioning in TRELLIS.2 Using the image_feature_extractor Module
You can implement custom image conditioning in TRELLIS.2 by creating a class that inherits from ImageConditioningBase in trellis2/modules/image_feature_extractor.py, implementing the extract_features() method, and specifying your class name in the pipeline configuration's image_cond_model dictionary.
The Microsoft TRELLIS.2 framework provides a flexible architecture for 3D generation that allows developers to inject custom visual guidance through its modular image conditioning system. By extending the image_feature_extractor module, you can integrate proprietary feature extraction networks or pretrained backbones to control how image inputs influence the texturing and geometry synthesis pipelines. This guide covers the exact implementation steps, configuration formats, and integration points required to add your own conditioning models.
Understanding the Image Conditioning Architecture
The ImageConditioningBase Interface
The foundation of the conditioning system resides in trellis2/modules/image_feature_extractor.py, which defines the ImageConditioningBase abstract class. This interface requires all custom implementations to define the extract_features(image) method, which accepts an input image tensor and returns a feature representation ready for fusion with the latent code.
Concrete implementations such as CLIPImageConditioning and VGGImageConditioning demonstrate the expected pattern: they wrap a pretrained backbone with an optional projection layer and return feature tensors of shape [B, C] for global conditioning or [B, C, H, W] for spatial conditioning. The base class ensures that the pipelines can interact with any conditioning model through a uniform API without knowing the specific implementation details.
Pipeline Integration
The pipelines dynamically instantiate conditioning models based on runtime configuration. In trellis2/pipelines/trellis2_texturing.py and trellis2/pipelines/trellis2_image_to_3d.py, the pipeline stores the instantiated model as pipeline.image_cond_model. During the forward pass, the pipeline calls pipeline.image_cond_model.extract_features(img) to obtain conditioning features that are subsequently fused with the latent representation before reaching the decoder.
The instantiation logic uses Python's getattr to resolve class names dynamically:
# From trellis2/pipelines/trellis2_texturing.py
pipeline.image_cond_model = getattr(
image_feature_extractor,
args['image_cond_model']['name']
)(**args['image_cond_model']['args'])
This design allows you to add new conditioning architectures without modifying the pipeline source code.
Step-by-Step Implementation Guide
1. Create a Custom Conditioning Class
Create a new class that inherits from ImageConditioningBase. Implement __init__() to initialize your backbone network and any learnable projection layers. Implement extract_features(self, image) to process the input and return the feature tensor.
The method signature must accept image: torch.Tensor with shape [B, 3, H, W] containing values typically in the [0, 1] range, and return a tensor compatible with the pipeline's feature fusion mechanism.
2. Register the Model
Ensure your class is accessible as an attribute of the image_feature_extractor module. You can add your class directly to trellis2/modules/image_feature_extractor.py or import it in the module's __init__.py. The class name must match exactly the string you will specify in the configuration file.
3. Configure the Pipeline
Update your YAML or JSON configuration to reference your custom class. The configuration requires a dictionary with name (the class name) and args (keyword arguments passed to __init__):
{
"image_cond_model": {
"name": "MyCustomImageConditioning",
"args": {
"pretrained_path": "path/to/weights.pt",
"output_dim": 256
}
}
}
When you run app_texturing.py or app.py with this configuration, the pipeline automatically instantiates your class and integrates it into the generation workflow.
Complete Code Examples
Minimal Custom Feature Extractor
The following implementation demonstrates a simple CNN-based conditioner that maps RGB images to 256-dimensional feature vectors:
# In trellis2/modules/image_feature_extractor.py
import torch
import torch.nn as nn
from .image_feature_extractor import ImageConditioningBase
class MyCustomImageConditioning(ImageConditioningBase):
"""Simple CNN that maps an RGB image to a 256-dim feature vector."""
def __init__(self, output_dim: int = 256):
super().__init__()
self.backbone = nn.Sequential(
nn.Conv2d(3, 32, 3, stride=2, padding=1),
nn.ReLU(),
nn.Conv2d(32, 64, 3, stride=2, padding=1),
nn.ReLU(),
nn.AdaptiveAvgPool2d(1)
)
self.proj = nn.Linear(64, output_dim)
def extract_features(self, image: torch.Tensor) -> torch.Tensor:
# image: [B, 3, H, W] (values in [0, 1])
x = self.backbone(image) # -> [B, 64, 1, 1]
x = x.view(x.size(0), -1) # -> [B, 64]
return self.proj(x) # -> [B, output_dim]
Running with Custom Configuration
Save your configuration as custom_cond_config.json:
{
"image_cond_model": {
"name": "MyCustomImageConditioning",
"args": {
"output_dim": 256
}
}
}
Execute the texturing pipeline with your custom conditioner:
python app_texturing.py --config custom_cond_config.json
The pipeline loads MyCustomImageConditioning, passing output_dim=256 to the constructor, and uses the resulting features to condition the 3D synthesis process.
Fine-Tuning the Conditioning Model
If your custom model includes learnable parameters that you want to optimize during training, include them in the optimizer alongside the generator parameters:
# In your training script after pipeline initialization
optimizer = torch.optim.Adam(
list(pipeline.generator.parameters()) +
list(pipeline.image_cond_model.parameters()),
lr=1e-4
)
This enables gradient flow through the image conditioning network during the GAN training loop, allowing the conditioner to adapt to the specific domain of your training data.
Summary
- Inherit from
ImageConditioningBasedefined intrellis2/modules/image_feature_extractor.pyto ensure API compatibility with the TRELLIS.2 pipelines. - Implement
extract_features()to return feature tensors of shape[B, C]or[B, C, H, W]that the pipeline can fuse with latent representations. - Configure via
image_cond_modelin your JSON or YAML config, specifying the exact class name and constructor arguments. - Access the model at runtime through
pipeline.image_cond_model, which the texturing and image-to-3D pipelines automatically populate. - Enable training by adding the conditioner's parameters to your optimizer if you need to fine-tune the feature extraction layers.
Frequently Asked Questions
What is the expected output shape for the extract_features method?
The method should return either a global feature vector of shape [B, C] where B is batch size and C is the feature dimension, or a spatial feature map of shape [B, C, H, W] if your conditioning requires spatial alignment. The pipeline in trellis2/pipelines/trellis2_texturing.py handles both formats, applying appropriate fusion operations before the decoder stage.
Can I load pretrained weights in my custom image conditioner?
Yes. Load pretrained weights inside your __init__ method using standard PyTorch loading mechanisms such as torch.load() or load_state_dict(). The args dictionary in your configuration can pass paths to checkpoint files, which your class receives as constructor arguments. This allows you to initialize with ImageNet weights or domain-specific pretrained features before freezing or fine-tuning.
How do I enable gradient updates for the conditioning model during training?
Include pipeline.image_cond_model.parameters() in your optimizer's parameter list alongside the generator parameters. According to the training scripts in the repository, the pipeline exposes the conditioner as a standard nn.Module, so you can toggle gradients using standard PyTorch patterns. Set the model to training mode with pipeline.image_cond_model.train() to ensure dropout and batch normalization layers behave correctly during optimization.
Which pipelines support custom image conditioning?
Both trellis2/pipelines/trellis2_texturing.py and trellis2/pipelines/trellis2_image_to_3d.py implement the dynamic loading mechanism via getattr(image_feature_extractor, config['image_cond_model']['name']). You can verify support by checking for the image_cond_model attribute in the pipeline initialization code. Entry points such as app_texturing.py and app.py parse the configuration and pass it to these pipelines, making them compatible with any custom conditioner following the ImageConditioningBase interface.
Have a question about this repo?
These articles cover the highlights, but your codebase questions are specific. Give your agent direct access to the source. Share this with your agent to get started:
curl -s "https://instagit.com/install.md" Maintain an open-source project? Get it listed too →