How Sol-RL Achieves 4.64× Faster Convergence with NVFP4 Rollout and BF16 Training
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, 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 (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 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:
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:
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:
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:
# 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 compute multi-dimensional scores (PickScore, CLIPScore, HPSv2, ImageReward), while Advantage Weighted Matching (AWM)—referenced in 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_teandwrap_forward_with_fp8intrain_scripts/sol_rl/train_utils.py) enables cheap exploration with 4-bit precision. - BF16 (via
BF16TELinearin lines 83-98) ensures stable policy updates without sacrificing training speed or numerical stability. - Configuration files in
configs/sol_rl/enforce this separation viapreview_model="compile_nvfp4"andfullrollout_model="compile"settings. - Rewards are computed via
diffusion/post_training/rewards.pyand 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 (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 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 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 is independent of the model's inference precision.
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 →