feat(rewards): improve single-stage distributional VF training

Add component-specific learning rates, regularization, cosine decay, and updated training documentation.
This commit is contained in:
Khalil Meftah
2026-08-02 11:29:02 +02:00
parent 9947afa0db
commit af443a4071
4 changed files with 161 additions and 15 deletions
@@ -19,6 +19,40 @@ lerobot-train \
This initializes SigLIP2 and Gemma3-270M from unimodal checkpoints and creates a
fresh Gemma3 multimodal connector.
### One-stage RECAP optimization
Train the connector, value readout, and Gemma together from the start with
component-specific learning rates, while keeping SigLIP2 frozen:
```bash
lerobot-train \
--reward_model.type=distributional_value_function \
--reward_model.freeze_vision_encoder=true \
--reward_model.freeze_language_model=false \
--reward_model.target_method=hl_gauss \
--reward_model.hl_gauss_sigma_ratio=0.75 \
--reward_model.use_one_hot_terminal=false \
--reward_model.value_dropout=0.1 \
--reward_model.optimizer_language_model_lr=1e-5 \
--reward_model.optimizer_multimodal_projector_lr=5e-5 \
--reward_model.optimizer_value_query_lr=1e-4 \
--reward_model.optimizer_value_head_lr=1e-4 \
--reward_model.optimizer_weight_decay=0.01 \
--reward_model.scheduler_warmup_steps=300 \
--reward_model.scheduler_decay_steps=6000 \
--reward_model.scheduler_decay_lr=1e-6 \
--dataset.repo_id=<dataset_repo_id> \
--dataset.eval_split=0.1 \
--eval_steps=500 \
--save_freq=500 \
--steps=6000
```
Defaults for this VF are head/query `1e-4`, connector `5e-5`, Gemma `1e-5`,
weight decay `0.01`, dropout `0.1`, gradient clipping `1.0`, and cosine decay
that brings the highest learning-rate group to `1e-6` while preserving group
ratios.
## Temporal SigLIP2
```bash
@@ -85,10 +85,21 @@ class DistributionalVFConfig(RewardModelConfig):
tokenizer_max_length: int = 200
# Training controls
value_dropout: float = 0.0
value_dropout: float = 0.1
freeze_vision_encoder: bool = False
freeze_language_model: bool = False
stop_gradient_to_vlm: bool = False
optimizer_vision_lr: float = 1e-6
optimizer_language_model_lr: float = 1e-5
optimizer_multimodal_projector_lr: float = 5e-5
optimizer_value_query_lr: float = 1e-4
optimizer_value_head_lr: float = 1e-4
optimizer_weight_decay: float = 1e-2
scheduler_warmup_steps: int = 500
scheduler_decay_steps: int = 40000
scheduler_decay_lr: float = 1e-6
# Deprecated compatibility field. Component-specific learning rates above
# now control optimization directly.
vision_encoder_lr_multiplier: float = 0.5
# Normalization
@@ -98,19 +109,46 @@ class DistributionalVFConfig(RewardModelConfig):
}
)
def __post_init__(self) -> None:
super().__post_init__()
learning_rates = {
"optimizer_vision_lr": self.optimizer_vision_lr,
"optimizer_language_model_lr": self.optimizer_language_model_lr,
"optimizer_multimodal_projector_lr": self.optimizer_multimodal_projector_lr,
"optimizer_value_query_lr": self.optimizer_value_query_lr,
"optimizer_value_head_lr": self.optimizer_value_head_lr,
}
for name, learning_rate in learning_rates.items():
if learning_rate <= 0:
raise ValueError(f"{name} must be > 0, got {learning_rate}")
if self.optimizer_weight_decay < 0:
raise ValueError(f"optimizer_weight_decay must be >= 0, got {self.optimizer_weight_decay}")
if not 0 <= self.value_dropout <= 1:
raise ValueError(f"value_dropout must be in [0,1], got {self.value_dropout}")
if self.scheduler_warmup_steps < 0 or self.scheduler_decay_steps < 1:
raise ValueError("scheduler_warmup_steps must be >= 0 and scheduler_decay_steps must be >= 1")
if self.scheduler_decay_lr < 0:
raise ValueError(f"scheduler_decay_lr must be >= 0, got {self.scheduler_decay_lr}")
def get_optimizer_preset(self) -> AdamWConfig:
return AdamWConfig(
lr=5e-5,
weight_decay=1e-10,
lr=self.optimizer_value_head_lr,
weight_decay=self.optimizer_weight_decay,
grad_clip_norm=1.0,
)
def get_scheduler_preset(self) -> CosineDecayWithWarmupSchedulerConfig:
return CosineDecayWithWarmupSchedulerConfig(
num_warmup_steps=500,
num_decay_steps=40000,
peak_lr=5e-5,
decay_lr=5e-5,
num_warmup_steps=self.scheduler_warmup_steps,
num_decay_steps=self.scheduler_decay_steps,
peak_lr=max(
self.optimizer_vision_lr,
self.optimizer_language_model_lr,
self.optimizer_multimodal_projector_lr,
self.optimizer_value_query_lr,
self.optimizer_value_head_lr,
),
decay_lr=self.scheduler_decay_lr,
)
def validate_features(self) -> None:
@@ -229,21 +229,40 @@ class DistributionalVFRewardModel(PreTrainedRewardModel):
def get_optim_params(self) -> list[dict]:
"""Optimizer param groups with per-component learning rates."""
vision_params = []
other_params = []
groups: dict[str, list[nn.Parameter]] = {
"vision_encoder": [],
"language_model": [],
"multimodal_projector": [],
"value_query": [],
"value_head": [],
}
for name, param in self.named_parameters():
if not param.requires_grad:
continue
if name.startswith("vision_encoder"):
vision_params.append(param)
groups["vision_encoder"].append(param)
elif name.startswith("language_model"):
groups["language_model"].append(param)
elif name.startswith("multi_modal_projector"):
groups["multimodal_projector"].append(param)
elif name.startswith("value_query"):
groups["value_query"].append(param)
elif name.startswith("value_head"):
groups["value_head"].append(param)
else:
other_params.append(param)
raise ValueError(f"Unrecognized trainable VF parameter: {name}")
base_lr = self.config.get_optimizer_preset().lr
learning_rates = {
"vision_encoder": self.config.optimizer_vision_lr,
"language_model": self.config.optimizer_language_model_lr,
"multimodal_projector": self.config.optimizer_multimodal_projector_lr,
"value_query": self.config.optimizer_value_query_lr,
"value_head": self.config.optimizer_value_head_lr,
}
return [
{"params": other_params},
{"params": vision_params, "lr": base_lr * self.config.vision_encoder_lr_multiplier},
{"params": params, "lr": learning_rates[name], "name": name}
for name, params in groups.items()
if params
]
def embed_image(self, image: Tensor) -> Tensor:
@@ -121,6 +121,22 @@ def test_config_defaults_match_pi06_gemma3_layout():
assert config.num_image_tokens == 256
assert config.target_method == "dirac_delta"
assert config.hl_gauss_sigma_ratio == 0.75
assert config.value_dropout == 0.1
assert config.optimizer_weight_decay == 0.01
assert config.scheduler_decay_lr == 1e-6
def test_optimizer_preset_matches_regularized_vf_recipe():
config = DistributionalVFConfig(device="cpu")
optimizer = config.get_optimizer_preset()
scheduler = config.get_scheduler_preset()
assert optimizer.lr == config.optimizer_value_head_lr
assert optimizer.weight_decay == 0.01
assert optimizer.grad_clip_norm == 1.0
assert scheduler.peak_lr == 1e-4
assert scheduler.decay_lr == 1e-6
# ------------------------------------------------------------------
@@ -508,6 +524,45 @@ def test_freeze_language_model():
assert p.requires_grad
@skip_if_package_missing("transformers")
def test_frozen_towers_optimizer_groups():
from lerobot.rewards.distributional_value_function.modeling_distributional_value_function import (
DistributionalVFRewardModel,
)
model = DistributionalVFRewardModel(_make_config(freeze_vision_encoder=True, freeze_language_model=True))
groups = {group["name"]: group for group in model.get_optim_params()}
assert set(groups) == {"multimodal_projector", "value_query", "value_head"}
assert groups["multimodal_projector"]["lr"] == 5e-5
assert groups["value_query"]["lr"] == 1e-4
assert groups["value_head"]["lr"] == 1e-4
@skip_if_package_missing("transformers")
def test_one_stage_optimizer_groups():
from lerobot.rewards.distributional_value_function.modeling_distributional_value_function import (
DistributionalVFRewardModel,
)
model = DistributionalVFRewardModel(_make_config(freeze_vision_encoder=True, freeze_language_model=False))
groups = {group["name"]: group for group in model.get_optim_params()}
assert set(groups) == {
"language_model",
"multimodal_projector",
"value_query",
"value_head",
}
assert groups["language_model"]["lr"] == 1e-5
optimizer = model.config.get_optimizer_preset().build(model.get_optim_params())
optimizer_groups = {group["name"]: group for group in optimizer.param_groups}
assert optimizer_groups["language_model"]["lr"] == 1e-5
assert optimizer_groups["multimodal_projector"]["lr"] == 5e-5
assert optimizer_groups["value_head"]["weight_decay"] == 0.01
@skip_if_package_missing("transformers")
def test_stop_gradient_to_vlm_preserves_value_query_grad():
"""With stop_gradient_to_vlm, the value query still gets gradients."""