How to Train SC-VAE Models (ShapeVAE and TexVAE) from Scratch in TRELLIS.2
You can train the SC-VAE models in TRELLIS.2 by instantiating FlexiDualGridVaeEncoder and FlexiDualGridVaeDecoder from trellis2/models/sc_vaes/fdg_vae.py, wrapping them in ShapeVaeTrainer from trellis2/trainers/vae/shape_vae.py, and launching via python train.py --config configs/shape_vae.yaml.
TRELLIS.2 implements two Sparse-Contrastive Variational Auto-Encoders (SC-VAEs): ShapeVAE for 3D geometry and TexVAE (also called PBR-VAE) for appearance. Both models share a common Sparse-Unet VAE backbone defined in trellis2/models/sc_vaes/sparse_unet_vae.py, making it straightforward to train either variant from scratch with minimal configuration changes.
SC-VAE Architecture Overview
Model Components
The backbone in sparse_unet_vae.py consists of:
- SparseResBlock3d — residual block for 3D sparse tensors
- SparseResBlockDownsample3d / Upsample3d — down/up-sampling layers
- SparseResBlockS2C3d / C2S3d — spatial-to-channel and channel-to-spatial conversions for subdivision mask prediction
- SparseConvNeXtBlock3d — ConvNeXt-style blocks providing the bulk of representation power
| Model | Purpose | Core Classes | Primary File |
|---|---|---|---|
| ShapeVAE | Compact latent representation of voxel geometry + occupancy | FlexiDualGridVaeEncoder / FlexiDualGridVaeDecoder |
trellis2/models/sc_vaes/fdg_vae.py |
| TexVAE | Latent map of appearance (albedo, normals, shading) | Same encoder/decoder with different channel layout | Same file, with dataset from trellis2/datasets/structured_latent_pbr.py |
Encoder and Decoder Outputs
The encoder maps a flexi-dual-grid (vertices + intersected flags) to a latent vector z.
The decoder in fdg_vae.py predicts three outputs:
- Vertex positions (
pred_vertice) — sigmoid-scaled 3-channel output - Intersection logits (
pred_intersected) — 3-channel binary classification head - Subdivision mask (
subdivision) — binary mask for upsampling the sparse grid during decoding
Training Loss Composition
ShapeVaeTrainer in trellis2/trainers/vae/shape_vae.py combines four loss families:
-
Direct regression losses
- Vertex L2 loss (
lambda_vertice) - Intersection BCE loss (
lambda_intersected)
- Vertex L2 loss (
-
Subdivision mask loss — binary cross-entropy on each predicted subdivision level (
lambda_subdiv) -
Rendering losses — decoder output rasterized with
MeshRendererand compared to ground-truth renders:- Mask L1 (
lambda_mask) - Depth L1 (
lambda_depth) - Normal L1 (
lambda_normal) - SSIM (
lambda_ssim) - LPIPS (
lambda_lpips)
- Mask L1 (
-
KL regularization (
lambda_kl)
The trainer randomly perturbs the camera for each batch via _randomize_camera to make the rendering loss robust to viewpoint changes.
Step-by-Step: Training ShapeVAE from Scratch
1. Prepare Your Dataset
ShapeVAE requires a dataset yielding dictionaries with these keys:
vertices—SparseTensor(6-channel: xyz + features)intersected—SparseTensor(binary occupancy flag)mesh— list ofMeshobjects for rendering-based loss
Use the provided StructuredLatentShape class in trellis2/datasets/structured_latent_shape.py or subclass it:
from trellis2.datasets.structured_latent_shape import StructuredLatentShape
dataset = StructuredLatentShape(root_dir="/path/to/shape_dataset")
2. Configure Model Hyperparameters
from trellis2.models.sc_vaes.fdg_vae import FlexiDualGridVaeEncoder, FlexiDualGridVaeDecoder
model_channels = [64, 128, 256, 512] # feature maps per level
latent_channels = 32 # latent vector dimensionality
num_blocks = [2, 2, 2, 2] # blocks per resolution
block_type = ["SparseResBlock3d"] * 4
down_block_type = ["SparseResBlockDownsample3d"] * 3
up_block_type = ["SparseResBlockUpsample3d"] * 3
block_args = [{"use_checkpoint": False}] * 4
encoder = FlexiDualGridVaeEncoder(
model_channels=model_channels,
latent_channels=latent_channels,
num_blocks=num_blocks,
block_type=block_type,
down_block_type=down_block_type,
block_args=block_args,
use_fp16=False,
)
decoder = FlexiDualGridVaeDecoder(
resolution=64, # voxel grid resolution
model_channels=model_channels,
latent_channels=latent_channels,
num_blocks=num_blocks,
block_type=block_type,
up_block_type=up_block_type,
block_args=block_args,
voxel_margin=0.5,
use_fp16=False,
)
3. Assemble the Trainer
from trellis2.trainers.vae.shape_vae import ShapeVaeTrainer
trainer = ShapeVaeTrainer(
models={"encoder": encoder, "decoder": decoder},
dataset=dataset,
output_dir="output/shape_vae",
batch_size=8,
max_steps=250_000,
optimizer={"type": "AdamW", "lr": 2e-4, "weight_decay": 1e-5},
lr_scheduler={"type": "CosineAnnealing", "T_max": 250_000},
lambda_subdiv=0.1,
lambda_intersected=0.1,
lambda_vertice=1e-2,
lambda_kl=1e-6,
lambda_mask=1.0,
lambda_depth=10.0,
lambda_normal=1.0,
lambda_ssim=0.2,
lambda_lpips=0.2,
fp16_mode=None,
)
4. Launch Training
Use the generic train.py entry point with a YAML configuration:
python train.py \
--config configs/shape_vae.yaml \
--output_dir output/shape_vae \
--device cuda
The configs/shape_vae.yaml file should mirror the Python parameters shown above. The trainer automatically handles data loading, camera randomization, rendering, loss computation, checkpointing, and EMA updates.
Training TexVAE (Texture/PBR VAE) from Scratch
TexVAE uses the same Sparse-Unet backbone with two key differences:
- Input channels: Set
in_channels=7for albedo + normal + roughness (the baseFlexiDualGridVaeEncoderhard-codesin_channels=6for shape, so subclass similarly for texture) - Dataset: Return 7-channel
SparseTensorusingtrellis2/datasets/structured_latent_pbr.py(create based on the shape version)
Adjust loss weights to emphasize albedo reconstruction:
python train.py \
--config configs/tex_vae.yaml \
--output_dir output/tex_vae \
--device cuda
Complete Minimal Training Script
import torch
from trellis2.models.sc_vaes.fdg_vae import FlexiDualGridVaeEncoder, FlexiDualGridVaeDecoder
from trellis2.trainers.vae.shape_vae import ShapeVaeTrainer
from trellis2.datasets.structured_latent_shape import StructuredLatentShape
# Dataset
dataset = StructuredLatentShape(root_dir="/path/to/shape_dataset")
# Model configuration
model_channels = [64, 128, 256, 512]
latent_channels = 32
num_blocks = [2, 2, 2, 2]
block_type = ["SparseResBlock3d"] * 4
down_block_type = ["SparseResBlockDownsample3d"] * 3
up_block_type = ["SparseResBlockUpsample3d"] * 3
block_args = [{"use_checkpoint": False}] * 4
encoder = FlexiDualGridVaeEncoder(
model_channels=model_channels,
latent_channels=latent_channels,
num_blocks=num_blocks,
block_type=block_type,
down_block_type=down_block_type,
block_args=block_args,
use_fp16=False,
)
decoder = FlexiDualGridVaeDecoder(
resolution=64,
model_channels=model_channels,
latent_channels=latent_channels,
num_blocks=num_blocks,
block_type=block_type,
up_block_type=up_block_type,
block_args=block_args,
voxel_margin=0.5,
use_fp16=False,
)
# Trainer
trainer = ShapeVaeTrainer(
models={"encoder": encoder, "decoder": decoder},
dataset=dataset,
output_dir="output/shape_vae",
batch_size=8,
max_steps=200_000,
optimizer={"type": "AdamW", "lr": 3e-4, "weight_decay": 1e-5},
lr_scheduler={"type": "CosineAnnealing", "T_max": 200_000},
lambda_subdiv=0.1,
lambda_intersected=0.1,
lambda_vertice=1e-2,
lambda_kl=1e-6,
lambda_mask=1.0,
lambda_depth=5.0,
lambda_normal=1.0,
)
# Single training step for debugging
batch = next(iter(trainer.dataloader))
batch = {k: v.to(trainer.device) for k, v in batch.items()}
losses, _ = trainer.training_losses(
vertices=batch["vertices"],
intersected=batch["intersected"],
mesh=batch["mesh"],
)
losses.loss.backward()
trainer.optimizer.step()
trainer.optimizer.zero_grad()
print(f"Current loss: {losses.loss.item():.4f}")
Key Source Files Reference
| File | Purpose |
|---|---|
trellis2/models/sc_vaes/sparse_unet_vae.py |
Sparse-Unet VAE backbone with residual blocks and ConvNeXt-style layers |
trellis2/models/sc_vaes/fdg_vae.py |
Flexi-Dual-Grid encoder/decoder implementations |
trellis2/trainers/vae/shape_vae.py |
ShapeVaeTrainer with loss composition and rendering |
trellis2/datasets/structured_latent_shape.py |
Dataset for ShapeVAE training |
train.py |
Generic entry point for all trainer configurations |
trellis2/renderers/mesh_renderer.py |
Differentiable renderer for mask/normal/depth maps |
trellis2/utils/data_utils.py |
Helper utilities including BalancedResumableSampler |
Summary
- ShapeVAE and TexVAE share the
FlexiDualGridVaeEncoder/Decoderarchitecture fromtrellis2/models/sc_vaes/fdg_vae.py - Core backbone resides in
trellis2/models/sc_vaes/sparse_unet_vae.pywithSparseConvNeXtBlock3dproviding representation power - Training orchestration happens through
ShapeVaeTrainerintrellis2/trainers/vae/shape_vae.py, combining direct regression, subdivision, rendering, and KL losses - Command-line launch uses
python train.py --config configs/{shape,tex}_vae.yaml - TexVAE adaptation requires only changing input channels (7 vs 6) and providing a PBR-featured dataset
Frequently Asked Questions
What is the difference between ShapeVAE and TexVAE in TRELLIS.2?
ShapeVAE learns compact latent representations of 3D geometry (voxels + occupancy), while TexVAE learns appearance properties including albedo, normals, and roughness. Both use identical FlexiDualGridVaeEncoder and FlexiDualGridVaeDecoder classes but differ in input channel count: 6 channels for shape versus 7 for texture.
How does the rendering loss work in SC-VAE training?
The rendering loss in ShapeVaeTrainer._render_batch uses the built-in MeshRenderer to rasterize both ground-truth and reconstructed meshes, then computes L1, SSIM, and LPIPS losses on mask, depth, and normal maps. Random camera perturbation during training ensures viewpoint robustness.
Can I train both VAEs with the same configuration file?
No, you need separate YAML configurations because the dataset classes differ (StructuredLatentShape vs StructuredLatentPBR) and loss weights should be tuned differently—texture training typically emphasizes albedo reconstruction with higher lambda_mask values.
What resolution should I use for the decoder?
The FlexiDualGridVaeDecoder accepts a resolution parameter (default 64 in the official implementation) that determines the voxel grid resolution. Higher resolutions capture finer geometric detail but require more memory and longer training times.
Where are subdivision masks used in the decoding process?
Subdivision masks predicted by the decoder at each resolution level control adaptive upsampling of the sparse grid, as implemented in the SparseResBlockS2C3d and SparseResBlockC2S3d blocks. The binary mask loss (lambda_subdiv) trains this mechanism via binary cross-entropy.
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 →