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:

  1. Direct regression losses

    • Vertex L2 loss (lambda_vertice)
    • Intersection BCE loss (lambda_intersected)
  2. Subdivision mask loss — binary cross-entropy on each predicted subdivision level (lambda_subdiv)

  3. Rendering losses — decoder output rasterized with MeshRenderer and compared to ground-truth renders:

    • Mask L1 (lambda_mask)
    • Depth L1 (lambda_depth)
    • Normal L1 (lambda_normal)
    • SSIM (lambda_ssim)
    • LPIPS (lambda_lpips)
  4. 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 of Mesh objects 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=7 for albedo + normal + roughness (the base FlexiDualGridVaeEncoder hard-codes in_channels=6 for shape, so subclass similarly for texture)
  • Dataset: Return 7-channel SparseTensor using trellis2/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/Decoder architecture from trellis2/models/sc_vaes/fdg_vae.py
  • Core backbone resides in trellis2/models/sc_vaes/sparse_unet_vae.py with SparseConvNeXtBlock3d providing representation power
  • Training orchestration happens through ShapeVaeTrainer in trellis2/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:

Share the following with your agent to get started:
curl -s "https://instagit.com/install.md"

Works with
Claude Codex Cursor VS Code OpenClaw Any MCP Client

Maintain an open-source project? Get it listed too →