mirror of
https://github.com/huggingface/lerobot.git
synced 2026-08-08 17:39:44 +00:00
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:
@@ -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
|
||||
|
||||
+45
-7
@@ -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:
|
||||
|
||||
+27
-8
@@ -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."""
|
||||
|
||||
Reference in New Issue
Block a user