# How to Train SC-VAE Models (ShapeVAE and TexVAE) from Scratch in TRELLIS.2

> Learn to train SC-VAE models from scratch in TRELLIS.2. Instantiate FlexiDualGridVaeEncoder and FlexiDualGridVaeDecoder, wrap them in ShapeVaeTrainer, and launch training with a simple command.

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

---

**You can train the SC-VAE models in TRELLIS.2 by instantiating `FlexiDualGridVaeEncoder` and `FlexiDualGridVaeDecoder` from [`trellis2/models/sc_vaes/fdg_vae.py`](https://github.com/microsoft/TRELLIS.2/blob/main/trellis2/models/sc_vaes/fdg_vae.py), wrapping them in `ShapeVaeTrainer` from [`trellis2/trainers/vae/shape_vae.py`](https://github.com/microsoft/TRELLIS.2/blob/main/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`](https://github.com/microsoft/TRELLIS.2/blob/main/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`](https://github.com/microsoft/TRELLIS.2/blob/main/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`](https://github.com/microsoft/TRELLIS.2/blob/main/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`](https://github.com/microsoft/TRELLIS.2/blob/main/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`](https://github.com/microsoft/TRELLIS.2/blob/main/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`](https://github.com/microsoft/TRELLIS.2/blob/main/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`](https://github.com/microsoft/TRELLIS.2/blob/main/trellis2/datasets/structured_latent_shape.py) or subclass it:

```python
from trellis2.datasets.structured_latent_shape import StructuredLatentShape

dataset = StructuredLatentShape(root_dir="/path/to/shape_dataset")

```

### 2. Configure Model Hyperparameters

```python
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

```python
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`](https://github.com/microsoft/TRELLIS.2/blob/main/train.py) entry point with a YAML configuration:

```bash
python train.py \
    --config configs/shape_vae.yaml \
    --output_dir output/shape_vae \
    --device cuda

```

The [`configs/shape_vae.yaml`](https://github.com/microsoft/TRELLIS.2/blob/main/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`](https://github.com/microsoft/TRELLIS.2/blob/main/trellis2/datasets/structured_latent_pbr.py) (create based on the shape version)

Adjust loss weights to emphasize albedo reconstruction:

```bash
python train.py \
    --config configs/tex_vae.yaml \
    --output_dir output/tex_vae \
    --device cuda

```

## Complete Minimal Training Script

```python
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`](https://github.com/microsoft/TRELLIS.2/blob/main/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`](https://github.com/microsoft/TRELLIS.2/blob/main/trellis2/models/sc_vaes/fdg_vae.py) | Flexi-Dual-Grid encoder/decoder implementations |
| [`trellis2/trainers/vae/shape_vae.py`](https://github.com/microsoft/TRELLIS.2/blob/main/trellis2/trainers/vae/shape_vae.py) | `ShapeVaeTrainer` with loss composition and rendering |
| [`trellis2/datasets/structured_latent_shape.py`](https://github.com/microsoft/TRELLIS.2/blob/main/trellis2/datasets/structured_latent_shape.py) | Dataset for ShapeVAE training |
| [`train.py`](https://github.com/microsoft/TRELLIS.2/blob/main/train.py) | Generic entry point for all trainer configurations |
| [`trellis2/renderers/mesh_renderer.py`](https://github.com/microsoft/TRELLIS.2/blob/main/trellis2/renderers/mesh_renderer.py) | Differentiable renderer for mask/normal/depth maps |
| [`trellis2/utils/data_utils.py`](https://github.com/microsoft/TRELLIS.2/blob/main/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`](https://github.com/microsoft/TRELLIS.2/blob/main/trellis2/models/sc_vaes/fdg_vae.py)
- **Core backbone** resides in [`trellis2/models/sc_vaes/sparse_unet_vae.py`](https://github.com/microsoft/TRELLIS.2/blob/main/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`](https://github.com/microsoft/TRELLIS.2/blob/main/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.