How to Train TRELLIS.2 from Scratch Using train.py and Custom JSON Configurations
You can train TRELLIS.2 from scratch by passing a custom JSON configuration to python -m train, which instantiates your dataset, builds sparse voxel models, and launches a distributed training loop with automatic mixed precision and EMA.
TRELLIS.2 is Microsoft’s open-source framework for generating 3D neural fields including voxels, meshes, and PBR textures. To train TRELLIS.2 from scratch, you utilize the train.py entry point to parse a JSON configuration file that declaratively defines your model architecture, data pipeline, and optimization hyperparameters. This guide walks through the exact source code paths, configuration schemas, and command-line workflows required to execute single-GPU or multi-node training runs.
Understanding the Training Pipeline Architecture
The training workflow follows a strict instantiation pipeline defined in train.py and trellis2/trainers/basic.py:
- CLI Argument Parsing –
train.pyprocesses flags for--config,--output_dir,--data_dir, and distributed settings (lines 97‑116). - JSON Configuration Loading – The config file specifies
models,dataset, andtrainerdictionaries that map to concrete Python classes. - Dynamic Object Construction – Using
getattrontrellis2.datasetsandtrellis2.models, the script instantiates your dataset (e.g.,SparseVoxelPbrDataset) and model architectures (e.g.,SparseUnetVaeEncoder), moving tensors to CUDA (lines 70‑78). - Trainer Initialization – The trainer class (e.g.,
PbrVaeTrainer) is instantiated with optimizers, gradient clipping, mixed-precision settings, and EMA configuration (lines 86‑94). - Distributed Execution –
trainer.run()executes the loop intrellis2/trainers/basic.py(lines 15‑99), handling gradient accumulation,DistributedDataParallelwrapping, and periodic checkpointing.
Step-by-Step Training Guide
1. Create the JSON Configuration File
Every training run requires a JSON file with three top-level keys: models, dataset, and trainer. Refer to the example configuration at configs/scvae/tex_vae_next_dc_f16c32_fp16.json as a template.
- Models: Define encoder and decoder architectures with arguments like
model_channels,latent_channels, anduse_fp16. - Dataset: Specify the dataset class name (e.g.,
SparseVoxelPbrDataset) and resolution filters. - Trainer: Set the trainer class (e.g.,
PbrVaeTrainer), batch size, optimizer, and logging intervals (i_log,i_sample,i_save).
2. Prepare Your Dataset
Ensure your data root matches the structure expected by your chosen dataset class. For SparseVoxelPbrDataset, this includes a metadata CSV and associated .pickle or .png files. The dataset loader filters instances according to aesthetic scores and voxel limits (see trellis2/datasets/sparse_voxel_pbr.py lines 94‑112).
3. Launch Training with train.py
Execute the training script from the repository root:
python -m train \
--config path/to/config.json \
--output_dir ./runs/experiment_01 \
--data_dir ./data \
--num_gpus 4
The script automatically detects available GPUs (torch.cuda.device_count()) unless overridden by --num_gpus. It initializes distributed training via setup_dist (lines 60‑65) and seeds RNGs per rank (lines 35‑40).
4. Resume and Fine-Tune from Checkpoints
Resume training using the automatic checkpoint finder:
python -m train \
--config path/to/config.json \
--output_dir ./runs/experiment_01 \
--load_dir ./runs/experiment_01 \
--ckpt latest
The find_ckpt helper (lines 17‑32 in train.py) scans ckpts/misc_*.pt to locate the highest step count.
Fine-tune specific components by adding a finetune_ckpt map inside your JSON config:
"finetune_ckpt": {
"encoder": "/path/to/encoder_step0123456.pt",
"decoder": "/path/to/decoder_step0123456.pt"
}
The trainer calls BasicTrainer.finetune_from (lines 23‑31 in trellis2/trainers/basic.py) to load matching weights while initializing new parameters randomly.
Complete Configuration Examples
Below is a fully functional configuration for training a VAE on PBR textures:
{
"models": {
"encoder": {
"name": "SparseUnetVaeEncoder",
"args": {
"in_channels": 6,
"model_channels": [64, 128, 256, 512, 1024],
"latent_channels": 32,
"num_blocks": [0, 4, 8, 16, 4],
"block_type": ["SparseConvNeXtBlock3d", "SparseConvNeXtBlock3d", "SparseConvNeXtBlock3d", "SparseConvNeXtBlock3d", "SparseConvNeXtBlock3d"],
"use_fp16": true
}
},
"decoder": {
"name": "SparseUnetVaeDecoder",
"args": {
"out_channels": 6,
"model_channels": [1024, 512, 256, 128, 64],
"latent_channels": 32,
"num_blocks": [4, 16, 8, 4, 0],
"use_fp16": true
}
}
},
"dataset": {
"name": "SparseVoxelPbrDataset",
"args": {
"resolution": 256,
"min_aesthetic_score": 4.5,
"max_active_voxels": 1000000,
"with_mesh": false,
"attrs": ["base_color", "metallic", "roughness", "alpha"]
}
},
"trainer": {
"name": "PbrVaeTrainer",
"args": {
"max_steps": 500000,
"batch_size_per_gpu": 8,
"optimizer": {
"name": "AdamW",
"args": {
"lr": 1e-4
}
},
"mix_precision_mode": "inflat_all",
"fp16_mode": "inflat_all",
"grad_clip": {
"name": "AdaptiveGradClipper",
"args": {
"max_norm": 1.0,
"clip_percentile": 95
}
},
"i_log": 500,
"i_sample": 5000,
"i_save": 5000
}
}
}
Key Training Parameters and Features
When you train TRELLIS.2 from scratch, several advanced features are configurable via the JSON trainer section:
- Mixed Precision Training: Set
fp16_modeandmix_precision_modeto"inflat_all"for automatic loss scaling and memory-efficient sparse convolution operations. - Adaptive Gradient Clipping: Use
AdaptiveGradClipper(as shown above) to stabilize training on high-resolution voxels by clipping gradients based on percentile statistics. - EMA (Exponential Moving Average): Enable EMA tracking in the trainer arguments to maintain shadow weights for improved sampling quality.
- Elastic Memory Management: The
BasicTrainersupports elastic memory allocation for large sparse voxel grids, automatically handling variable-size tensors across batches.
Summary
- JSON-driven configuration: TRELLIS.2 uses declarative JSON files to specify
models,dataset, andtrainercomponents, eliminating hard-coded training scripts. - Entry point:
train.pyhandles CLI parsing, distributed setup, and dynamic instantiation of dataset and model classes viagetattr. - Distributed training: The framework auto-detects GPU count, initializes
DistributedDataParallel, and manages RNG seeding per rank. - Checkpointing: Resume training with
--load_dirand--ckpt latest, or fine-tune specific modules using thefinetune_ckptJSON map andBasicTrainer.finetune_from. - Monitoring: TensorBoard logs write to
<output_dir>/tb_logs/, sample images save to<output_dir>/samples/, and checkpoints store in<output_dir>/ckpts/.
Frequently Asked Questions
What file format does TRELLIS.2 use for training configurations?
TRELLIS.2 uses standard JSON files containing three required top-level keys: models, dataset, and trainer. Each key maps to a dictionary specifying the class name (name) and initialization arguments (args), allowing you to compose complex pipelines without modifying Python code.
How do I resume training from a checkpoint in TRELLIS.2?
Pass --load_dir pointing to your previous output directory and --ckpt latest to the train.py script. The find_ckpt utility (lines 17‑32 in train.py) automatically locates the most recent misc_*.pt file in the ckpts/ subdirectory and restores model weights, optimizer states, and training step counters.
Can I fine-tune specific model components without retraining everything?
Yes. Add a finetune_ckpt dictionary to your JSON configuration, mapping component names (like "encoder" or "decoder") to specific .pt file paths. During initialization, BasicTrainer.finetune_from (lines 23‑31 in trellis2/trainers/basic.py) loads only the matching weight keys, leaving other parameters randomly initialized for fine-tuning.
Where are training logs, samples, and checkpoints stored?
By default, all outputs are written relative to your --output_dir argument. TensorBoard logs appear in tb_logs/, generated sample images (controlled by i_sample intervals) save to samples/, and model checkpoints (controlled by i_save) are stored in ckpts/. This structure allows easy resumption and monitoring across distributed nodes.
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 →