# How to Fine-Tune TRELLIS.2 on a Custom Dataset: A Complete Technical Guide

> Learn how to fine-tune TRELLIS.2 on your custom dataset. This guide details subclassing DatasetBase, configuring JSON, and launching training for optimal results.

- Repository: [Microsoft/TRELLIS.2](https://github.com/microsoft/TRELLIS.2)
- Tags: how-to-guide
- Published: 2026-08-03

---

**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`](https://github.com/microsoft/TRELLIS.2/blob/main/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`](https://github.com/microsoft/TRELLIS.2/blob/main/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`](https://github.com/microsoft/TRELLIS.2/blob/main/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:

```python

# 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:

```json
// 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:

```bash
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`](https://github.com/microsoft/TRELLIS.2/blob/main/train.py) performs the following sequence:

1. Parses the JSON config and CLI arguments into an `edict` configuration object.
2. Dynamically instantiates `MyCustomDataset` via `getattr(datasets, cfg.dataset.name)` (lines 70-71 in [`train.py`](https://github.com/microsoft/TRELLIS.2/blob/main/train.py)).
3. Constructs the model (e.g., `ShapeVAE`) and moves it to CUDA.
4. Passes both the model and dataset to `BasicTrainer`, which runs the training loop and saves checkpoints to `cfg.output_dir`.

## Key Source Files for Fine-Tuning

Understanding these core files helps debug and extend the fine-tuning pipeline:

- **[`trellis2/datasets/components.py`](https://github.com/microsoft/TRELLIS.2/blob/main/trellis2/datasets/components.py)** – Contains `StandardDatasetBase` (lines 13-30) and visualization mix-ins that your custom dataset inherits.
- **[`trellis2/datasets/sparse_voxel_pbr.py`](https://github.com/microsoft/TRELLIS.2/blob/main/trellis2/datasets/sparse_voxel_pbr.py)** – Reference implementation showing how concrete datasets extend the base class.
- **[`trellis2/trainers/basic.py`](https://github.com/microsoft/TRELLIS.2/blob/main/trellis2/trainers/basic.py)** – Implements `BasicTrainer`, handling the training loop, optimizer steps, and checkpoint management.
- **[`train.py`](https://github.com/microsoft/TRELLIS.2/blob/main/train.py)** – Entry point (lines 59-95) that parses configs, instantiates components, and supports distributed training via `torch.multiprocessing.spawn`.
- **[`configs/scvae/shape_vae_next_dc_f16c32_fp16.json`](https://github.com/microsoft/TRELLIS.2/blob/main/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 `StandardDatasetBase`** in [`trellis2/datasets/components.py`](https://github.com/microsoft/TRELLIS.2/blob/main/trellis2/datasets/components.py) to 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.py`](https://github.com/microsoft/TRELLIS.2/blob/main/train.py)** using the `--config` argument; the script handles instantiation, distributed training setup, and automatic checkpoint resumption.
- **Leverage `BasicTrainer`** for 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`](https://github.com/microsoft/TRELLIS.2/blob/main/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`](https://github.com/microsoft/TRELLIS.2/blob/main/train.py) script automatically spawns processes via `torch.multiprocessing.spawn` and handles distributed synchronization when `num_gpus` is greater than 1.