How to Fine-Tune TRELLIS.2 on a Custom Dataset: A Complete Technical Guide
Fine-tuning TRELLIS.2 on a custom dataset requires implementing a subclass of StandardDatasetBase, configuring a JSON file with your dataset parameters, and launching the training script with checkpoint resumption support.
Microsoft's TRELLIS.2 framework provides a modular, configuration-driven pipeline for 3D generative modeling that simplifies fine-tuning through declarative JSON configs and standardized dataset interfaces. Unlike monolithic training scripts, TRELLIS.2 separates data loading, model instantiation, and training logic into distinct layers, allowing you to swap datasets without modifying core training code. By inheriting from the base dataset class and updating a configuration file, you can fine-tune pretrained shape VAEs or other models on your proprietary data.
Understanding the TRELLIS.2 Training Architecture
The TRELLIS.2 repository organizes its training pipeline into four distinct layers that communicate through configuration dictionaries.
Dataset Layer: StandardDatasetBase
All data providers in TRELLIS.2 inherit from StandardDatasetBase located in trellis2/datasets/components.py (lines 13-30). This base class implements the PyTorch Dataset interface and provides visualization mix-ins for debugging. To create a custom dataset, you only need to implement __len__ and __getitem__, or override mix-ins for augmented functionality.
Model Layer: Modular Torch Modules
Model definitions reside under trellis2/models/ as standard torch.nn.Module instances. Each model instantiates from a config entry (e.g., "shape_vae"), allowing you to specify architecture parameters like latent_dim and num_blocks without touching the model code.
Trainer Layer: BasicTrainer
The BasicTrainer class in trellis2/trainers/basic.py orchestrates the training loop. It receives a dictionary of models and a dataset instance, then handles forward/backward passes, logging, checkpointing, and optional distributed training across multiple GPUs.
Training Entry Point: train.py
The train.py script (lines 59-95) serves as the main entry point. It parses JSON configurations, builds the dataset and models, and delegates execution to BasicTrainer. The script supports automatic checkpoint resumption and multi-GPU launch via torch.multiprocessing.spawn.
Implementing a Custom Dataset
To fine-tune on your data, create a Python file that subclasses StandardDatasetBase. The following example demonstrates a minimal implementation expecting a directory with images and a metadata JSON file:
# my_dataset.py
import torch
from trellis2.datasets.components import StandardDatasetBase
import json
import pathlib
class MyCustomDataset(StandardDatasetBase):
"""
Example custom dataset for fine-tuning TRELLIS.2.
Expects folder layout:
data/
├─ images/
└─ metadata.json
"""
def __init__(self, data_root: str, transform=None):
super().__init__()
self.data_root = data_root
self.transform = transform
meta_path = pathlib.Path(data_root) / "metadata.json"
with open(meta_path) as f:
self.samples = json.load(f)
def __len__(self):
return len(self.samples)
def __getitem__(self, idx):
sample = self.samples[idx]
img_path = pathlib.Path(self.data_root) / "images" / sample["image"]
img = torch.load(img_path) # Replace with PIL/OpenCV loader as needed
if self.transform:
img = self.transform(img)
return {"image": img, "label": sample.get("label", -1)}
Place this file in your project directory or within the trellis2/datasets/ namespace. The __getitem__ method must return a dictionary containing the fields your model expects—typically an "image" tensor and optional labels.
Configuring the Fine-Tuning Job
TRELLIS.2 uses JSON configuration files to specify all hyperparameters, dataset selections, and model architectures. To fine-tune on your custom dataset, create a config file that references your dataset class and points to pretrained checkpoints:
// my_finetune_config.json
{
"dataset": {
"name": "MyCustomDataset",
"args": {
"data_dir": "./my_data"
}
},
"models": {
"vae": {
"name": "ShapeVAE",
"args": {
"latent_dim": 256,
"num_blocks": 12
}
}
},
"trainer": {
"name": "BasicTrainer",
"args": {
"batch_size": 8,
"num_epochs": 100,
"lr": 1e-4,
"log_interval": 50
}
},
"output_dir": "./finetune_out",
"load_dir": "./pretrained_checkpoints",
"ckpt": "latest",
"auto_retry": 3,
"num_gpus": -1
}
The load_dir and ckpt fields enable fine-tuning by loading pretrained weights before training begins. Set num_gpus to -1 to utilize all available GPUs, or specify a positive integer for limited multi-GPU training.
Launching the Fine-Tuning Process
Execute the training script with your configuration file. The CLI arguments override JSON values where specified:
python train.py \
--config my_finetune_config.json \
--output_dir ./finetune_out \
--load_dir ./pretrained_checkpoints \
--ckpt latest \
--data_dir ./my_data
Under the hood, train.py performs the following sequence:
- Parses the JSON config and CLI arguments into an
edictconfiguration object. - Dynamically instantiates
MyCustomDatasetviagetattr(datasets, cfg.dataset.name)(lines 70-71 intrain.py). - Constructs the model (e.g.,
ShapeVAE) and moves it to CUDA. - Passes both the model and dataset to
BasicTrainer, which runs the training loop and saves checkpoints tocfg.output_dir.
Key Source Files for Fine-Tuning
Understanding these core files helps debug and extend the fine-tuning pipeline:
trellis2/datasets/components.py– ContainsStandardDatasetBase(lines 13-30) and visualization mix-ins that your custom dataset inherits.trellis2/datasets/sparse_voxel_pbr.py– Reference implementation showing how concrete datasets extend the base class.trellis2/trainers/basic.py– ImplementsBasicTrainer, handling the training loop, optimizer steps, and checkpoint management.train.py– Entry point (lines 59-95) that parses configs, instantiates components, and supports distributed training viatorch.multiprocessing.spawn.configs/scvae/shape_vae_next_dc_f16c32_fp16.json– Example configuration illustrating dataset specification and model hyperparameters for shape VAE training.
Summary
Fine-tuning TRELLIS.2 on custom data follows a declarative, modular workflow:
- Inherit from
StandardDatasetBaseintrellis2/datasets/components.pyto implement your data loader with only__len__and__getitem__. - Configure via JSON by specifying your dataset class name, model architecture, and pretrained checkpoint paths without modifying training code.
- Launch with
train.pyusing the--configargument; the script handles instantiation, distributed training setup, and automatic checkpoint resumption. - Leverage
BasicTrainerfor production-ready training loops with built-in logging and multi-GPU support.
Frequently Asked Questions
Do I need to modify the original TRELLIS.2 source code to add a custom dataset?
No. The framework uses dynamic class loading via getattr(datasets, cfg.dataset.name) in train.py (lines 70-71). As long as your custom dataset inherits from StandardDatasetBase and is importable in the Python path, you only need to update the JSON configuration file to reference your class name.
What data format should the __getitem__ method return?
The method must return a dictionary compatible with your model's forward pass. For TRELLIS.2 VAEs, this typically includes an "image" key containing a torch.Tensor and optional metadata keys such as "label". Inspect the specific model implementation in trellis2/models/ to verify expected input fields.
How do I resume fine-tuning from a specific checkpoint?
Set the ckpt field in your JSON config to the checkpoint filename (e.g., "model_10000.pt") or use "latest" to automatically load the most recent checkpoint from load_dir. The BasicTrainer handles state restoration for model weights, optimizer states, and training step counters.
Can I fine-tune using multiple GPUs on a single machine?
Yes. Set "num_gpus": -1 in your config to utilize all available GPUs, or specify a positive integer to limit GPU usage. The train.py script automatically spawns processes via torch.multiprocessing.spawn and handles distributed synchronization when num_gpus is greater than 1.
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 →