# How Sol-RL Achieves 4.64× Faster Convergence with NVFP4 Rollout and BF16 Training

> Discover how Sol-RL achieves 4.64x faster convergence. Learn about NVFP4 rollout and BF16 training, which eliminate precision-stability trade-offs for efficient diffusion RL.

- Repository: [NVIDIA Research Projects/Sana](https://github.com/NVlabs/Sana)
- Tags: deep-dive
- Published: 2026-05-19

---

**Sol-RL decouples exploration from exploitation by using NVFP4 (4-bit) precision for cheap rollout generation and BF16 for stable gradient updates, eliminating the precision-stability trade-off that slows traditional diffusion RL loops.**

Sol-RL is the post-training reinforcement learning pipeline for the **NVlabs/Sana** diffusion models (including SANA, FLUX.1, and SD-3.5-L). By strategically separating the **exploration** (rollout) and **exploitation** (training) phases into different numerical precisions, Sol-RL achieves a reported **4.64× faster convergence** compared to conventional single-precision RL pipelines.

## The Two-Stage Precision Decoupling Strategy

### Stage 1: NVFP4 Rollout for Cheap Exploration

During the rollout phase, Sol-RL employs **NVFP4** (a 4-bit floating-point format with E2M1 mantissa) through NVIDIA's Transformer Engine (TE). In [`train_scripts/sol_rl/train_utils.py`](https://github.com/NVlabs/Sana/blob/main/train_scripts/sol_rl/train_utils.py), the `NVFP4_RECIPE` (lines 27-32) configures the FP8 emulation layer to use 4-bit arithmetic for forward passes.

The `replace_linear_with_te` function (lines 10-52) substitutes every `nn.Linear` layer with `te.Linear` equivalents, while `wrap_forward_with_fp8` (lines 70-78) wraps these layers inside `te.fp8_autocast` context managers. This compiled NVFP4 model generates candidate images with approximately **4× less memory bandwidth and compute** than FP16/BF16 equivalents, enabling massive batch rollouts without proportional GPU cost increases.

### Stage 2: BF16 Training for Stable Exploitation

For the policy gradient update, Sol-RL switches to **BF16** (bfloat16) precision to maintain numerical stability. The `BF16TELinear` class in [`train_scripts/sol_rl/train_utils.py`](https://github.com/NVlabs/Sana/blob/main/train_scripts/sol_rl/train_utils.py) (lines 83-98) implements a thin wrapper that casts inputs to BF16 before executing TE linear operations.

This separation prevents the gradient instability that would occur if 4-bit weights were used for backpropagation. The configuration in [`configs/sol_rl/sana.py`](https://github.com/NVlabs/Sana/blob/main/configs/sol_rl/sana.py) explicitly sets `preview_model="compile_nvfp4"` for rollouts and `fullrollout_model="compile"` for the BF16 training phase, enforcing this architectural boundary.

## Implementation: Key Code Components

Convert your diffusion UNet to NVFP4-compatible layers using the TE replacement utility:

```python
from train_scripts.sol_rl.train_utils import replace_linear_with_te

model = ...  # Your diffusion UNet or transformer

replaced, skipped, *_ = replace_linear_with_te(
    model,
    skip_modules=["some_module_to_keep_fp32"],
    min_dim=0,
)
print(f"NVFP4 ready: {replaced} layers replaced, {skipped} skipped")

```

Wrap specific modules for FP4 autocast during rollout generation:

```python
from train_scripts.sol_rl.train_utils import wrap_forward_with_fp8
import transformer_engine.pytorch as te

module = torch.nn.Linear(1024, 1024)
wrap_forward_with_fp8(module)  # Module now runs in NVFP4 precision

with te.fp8_autocast(enabled=True, fp8_recipe=NVFP4_RECIPE):
    output = module(torch.randn(1, 1024))

```

For the training phase, instantiate the BF16-compatible layer:

```python
from train_scripts.sol_rl.train_utils import BF16TELinear

te_linear = te.Linear(1024, 1024, bias=True)
bf16_layer = BF16TELinear(te_linear)  # Casts inputs to bfloat16 before TE ops

output = bf16_layer(torch.randn(1, 1024))

```

The complete RL pipeline orchestrates these stages in a dual-model setup:

```python

# Conceptual Sol-RL training loop

for step in range(num_steps):
    # Stage 1: Fast NVFP4 rollout (exploration)

    with te.fp8_autocast(enabled=True, fp8_recipe=NVFP4_RECIPE):
        images = nvfp4_model.generate(prompts, batch_size=64)  # 4x larger batches possible

    
    # Reward computation via diffusion/post_training/rewards.py

    rewards = compute_rewards(images, prompts)  # PickScore, CLIPScore, HPSv2, ImageReward

    
    # Stage 2: BF16 policy update (exploitation) with Advantage Weighted Matching

    loss = awm_loss(bf16_model, images, rewards)
    loss.backward()
    optimizer.step()

```

## Why NVFP4 Plus BF16 Accelerates Convergence

**Cheap exploration:** NVFP4 reduces forward-pass memory and compute costs by approximately 4×, allowing significantly more rollouts per training step without increasing wall-clock time. This higher sample density improves policy coverage and exploration efficiency.

**Accurate exploitation:** BF16 retains FP32's dynamic range while halving memory usage compared to FP32. Training the policy network in BF16 yields stable gradients that would be impossible with 4-bit backpropagation, ensuring the optimizer receives high-fidelity update signals.

**Decoupled precision architecture:** By isolating generation noise (NVFP4) from optimization stability (BF16), Sol-RL avoids the "precision drag" that forces traditional pipelines to compromise between rollout throughput and update quality. The reward models in [`diffusion/post_training/rewards.py`](https://github.com/NVlabs/Sana/blob/main/diffusion/post_training/rewards.py) compute multi-dimensional scores (PickScore, CLIPScore, HPSv2, ImageReward), while **Advantage Weighted Matching (AWM)**—referenced in [`docs/sol_rl.md`](https://github.com/NVlabs/Sana/blob/main/docs/sol_rl.md) (line 118)—converts these into sample-efficient policy gradients.

## Summary

- Sol-RL achieves **4.64× faster convergence** by decoupling NVFP4 rollout generation from BF16 gradient updates in the NVlabs/Sana codebase.
- **NVFP4** (implemented via `replace_linear_with_te` and `wrap_forward_with_fp8` in [`train_scripts/sol_rl/train_utils.py`](https://github.com/NVlabs/Sana/blob/main/train_scripts/sol_rl/train_utils.py)) enables cheap exploration with 4-bit precision.
- **BF16** (via `BF16TELinear` in lines 83-98) ensures stable policy updates without sacrificing training speed or numerical stability.
- Configuration files in `configs/sol_rl/` enforce this separation via `preview_model="compile_nvfp4"` and `fullrollout_model="compile"` settings.
- Rewards are computed via [`diffusion/post_training/rewards.py`](https://github.com/NVlabs/Sana/blob/main/diffusion/post_training/rewards.py) and optimized using **Advantage Weighted Matching (AWM)** for sample-efficient learning.

## Frequently Asked Questions

### What is NVFP4 and how does it differ from FP8?

NVFP4 is a 4-bit floating-point format (E2M1) implemented through NVIDIA Transformer Engine's FP8 emulation layer. While FP8 uses 8 bits, NVFP4 halves the precision further to 4 bits, reducing memory bandwidth by approximately 4× compared to FP16/BF16. In Sol-RL, it is configured via the `NVFP4_RECIPE` in [`train_scripts/sol_rl/train_utils.py`](https://github.com/NVlabs/Sana/blob/main/train_scripts/sol_rl/train_utils.py) (lines 27-32) and executed within `te.fp8_autocast` contexts during the rollout phase.

### Why not use NVFP4 for both rollout and training?

Using 4-bit precision for backpropagation would introduce catastrophic numerical instability and gradient compression artifacts. Sol-RL uses BF16 for the training phase (via `BF16TELinear`) because it maintains the dynamic range necessary for stable optimization while still benefiting from reduced memory usage compared to FP32. This decoupling prevents the precision-stability trade-off that would otherwise limit convergence speed.

### Where is the 4.64× speedup number documented?

The convergence acceleration figure is documented in the Sol-RL README and [`docs/sol_rl.md`](https://github.com/NVlabs/Sana/blob/main/docs/sol_rl.md) within the NVlabs/Sana repository. The speedup derives from generating 4× more rollouts per unit time with NVFP4 while maintaining efficient BF16 updates, effectively increasing sample efficiency without proportional compute cost increases. The actual training script [`train_scripts/sol_rl/train_sana.py`](https://github.com/NVlabs/Sana/blob/main/train_scripts/sol_rl/train_sana.py) implements this orchestration logic.

### How does Advantage Weighted Matching (AWM) integrate with the precision switching?

AWM converts multi-reward scores (from PickScore, CLIPScore, etc.) into policy gradients. It operates on the BF16 model's outputs during the exploitation phase, using the high-volume rollouts generated by the NVFP4 preview model. This decoupling allows AWM to benefit from abundant NVFP4 samples while computing stable updates in BF16 precision, as the reward computation in [`diffusion/post_training/rewards.py`](https://github.com/NVlabs/Sana/blob/main/diffusion/post_training/rewards.py) is independent of the model's inference precision.