mirror of
https://github.com/huggingface/lerobot.git
synced 2026-07-23 17:56:07 +00:00
fix(rewards): simplify VF architecture, remove CLS token, use bidirectional mean-pool or last_token readout
This commit is contained in:
+8
-5
@@ -79,6 +79,9 @@ class DistributionalVFConfig(RewardModelConfig):
|
|||||||
stop_gradient_to_vlm: bool = False
|
stop_gradient_to_vlm: bool = False
|
||||||
vision_encoder_lr_multiplier: float = 0.5
|
vision_encoder_lr_multiplier: float = 0.5
|
||||||
|
|
||||||
|
# Readout: "mean_pool" (average all tokens) or "last_token" (causal LM last position)
|
||||||
|
readout: str = "mean_pool"
|
||||||
|
|
||||||
# Normalization
|
# Normalization
|
||||||
normalization_mapping: dict[str, NormalizationMode] = field(
|
normalization_mapping: dict[str, NormalizationMode] = field(
|
||||||
default_factory=lambda: {
|
default_factory=lambda: {
|
||||||
@@ -88,17 +91,17 @@ class DistributionalVFConfig(RewardModelConfig):
|
|||||||
|
|
||||||
def get_optimizer_preset(self) -> AdamWConfig:
|
def get_optimizer_preset(self) -> AdamWConfig:
|
||||||
return AdamWConfig(
|
return AdamWConfig(
|
||||||
lr=1e-4,
|
lr=5e-5,
|
||||||
weight_decay=1e-5,
|
weight_decay=1e-10,
|
||||||
grad_clip_norm=1.0,
|
grad_clip_norm=1.0,
|
||||||
)
|
)
|
||||||
|
|
||||||
def get_scheduler_preset(self) -> CosineDecayWithWarmupSchedulerConfig:
|
def get_scheduler_preset(self) -> CosineDecayWithWarmupSchedulerConfig:
|
||||||
return CosineDecayWithWarmupSchedulerConfig(
|
return CosineDecayWithWarmupSchedulerConfig(
|
||||||
num_warmup_steps=500,
|
num_warmup_steps=500,
|
||||||
num_decay_steps=30000,
|
num_decay_steps=40000,
|
||||||
peak_lr=1e-4,
|
peak_lr=5e-5,
|
||||||
decay_lr=1e-6,
|
decay_lr=5e-5,
|
||||||
)
|
)
|
||||||
|
|
||||||
def validate_features(self) -> None:
|
def validate_features(self) -> None:
|
||||||
|
|||||||
+34
-124
@@ -54,38 +54,6 @@ else:
|
|||||||
SiglipVisionModel = None # type: ignore[assignment]
|
SiglipVisionModel = None # type: ignore[assignment]
|
||||||
|
|
||||||
|
|
||||||
def make_att_2d_masks(pad_masks: Tensor, att_masks: Tensor) -> Tensor:
|
|
||||||
"""Copied from big_vision.
|
|
||||||
|
|
||||||
Tokens can attend to valid inputs tokens which have a cumulative mask_ar
|
|
||||||
smaller or equal to theirs. This way `mask_ar` int[B, N] can be used to
|
|
||||||
setup several types of attention, for example:
|
|
||||||
|
|
||||||
[[1 1 1 1 1 1]]: pure causal attention.
|
|
||||||
|
|
||||||
[[0 0 0 1 1 1]]: prefix-lm attention. The first 3 tokens can attend between
|
|
||||||
themselves and the last 3 tokens have a causal attention. The first
|
|
||||||
entry could also be a 1 without changing behaviour.
|
|
||||||
|
|
||||||
[[1 0 1 0 1 0 0 1 0 0]]: causal attention between 4 blocks. Tokens of a
|
|
||||||
block can attend all previous blocks and all tokens on the same block.
|
|
||||||
|
|
||||||
Args:
|
|
||||||
input_mask: bool[B, N] true if its part of the input, false if padding.
|
|
||||||
mask_ar: int32[B, N] mask that's 1 where previous tokens cannot depend on
|
|
||||||
it and 0 where it shares the same attention mask as the previous token.
|
|
||||||
"""
|
|
||||||
if att_masks.ndim != 2:
|
|
||||||
raise ValueError(att_masks.ndim)
|
|
||||||
if pad_masks.ndim != 2:
|
|
||||||
raise ValueError(pad_masks.ndim)
|
|
||||||
|
|
||||||
cumsum = torch.cumsum(att_masks, dim=1)
|
|
||||||
att_2d_masks = cumsum[:, None, :] <= cumsum[:, :, None]
|
|
||||||
pad_2d_masks = pad_masks[:, None, :] * pad_masks[:, :, None]
|
|
||||||
return att_2d_masks & pad_2d_masks
|
|
||||||
|
|
||||||
|
|
||||||
class ValueHead(nn.Module):
|
class ValueHead(nn.Module):
|
||||||
"""Categorical value projection: hidden state → bin logits.
|
"""Categorical value projection: hidden state → bin logits.
|
||||||
|
|
||||||
@@ -128,11 +96,12 @@ class DistributionalVFRewardModel(PreTrainedRewardModel):
|
|||||||
Trained with cross-entropy on HL-Gauss or Dirac delta targets centered on
|
Trained with cross-entropy on HL-Gauss or Dirac delta targets centered on
|
||||||
per-task normalized Monte Carlo returns.
|
per-task normalized Monte Carlo returns.
|
||||||
|
|
||||||
Architecture: monolithic VLM — SigLIP2-so400m + Gemma3-270M (~670M params).
|
Architecture: SigLIP2-so400m + Linear(1152→640) + Gemma3-270M.
|
||||||
Multi-camera images are encoded by SigLIP2 (256 patches each), projected to
|
Multi-camera images are encoded by SigLIP2 (256 patches each), projected to
|
||||||
Gemma3's hidden dim, concatenated with tokenized language, and processed by
|
Gemma3's hidden dim, concatenated with tokenized language, and processed by
|
||||||
all 18 Gemma3 transformer layers. A [CLS] token appended at the end provides
|
all 18 Gemma3 transformer layers.
|
||||||
the value readout via a 2-layer MLP head.
|
|
||||||
|
Mean-pooled last-layer hidden states are read out through a 2-layer MLP value head.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
name = "distributional_value_function"
|
name = "distributional_value_function"
|
||||||
@@ -154,12 +123,7 @@ class DistributionalVFRewardModel(PreTrainedRewardModel):
|
|||||||
nn.init.normal_(self.image_proj.weight, std=0.02)
|
nn.init.normal_(self.image_proj.weight, std=0.02)
|
||||||
nn.init.zeros_(self.image_proj.bias)
|
nn.init.zeros_(self.image_proj.bias)
|
||||||
|
|
||||||
# Learnable [CLS] token — appended to the sequence before Gemma3.
|
# Value head: last-token hidden state → MLP → num_bins logits
|
||||||
# nn.Embedding (not nn.Parameter) for FSDP compatibility.
|
|
||||||
self.cls_embedding = nn.Embedding(1, self.gemma3_hidden)
|
|
||||||
nn.init.normal_(self.cls_embedding.weight, std=0.02)
|
|
||||||
|
|
||||||
# Value head: MLP projection → num_bins logits
|
|
||||||
self.value_head = ValueHead(
|
self.value_head = ValueHead(
|
||||||
hidden_size=self.gemma3_hidden,
|
hidden_size=self.gemma3_hidden,
|
||||||
num_bins=config.num_value_bins,
|
num_bins=config.num_value_bins,
|
||||||
@@ -407,70 +371,63 @@ class DistributionalVFRewardModel(PreTrainedRewardModel):
|
|||||||
terminal_distribution = self.one_hot_target(target_value)
|
terminal_distribution = self.one_hot_target(target_value)
|
||||||
return torch.where(is_terminal[:, None].bool(), terminal_distribution, base_distribution)
|
return torch.where(is_terminal[:, None].bool(), terminal_distribution, base_distribution)
|
||||||
|
|
||||||
def forward(self, batch: dict[str, Tensor]) -> tuple[Tensor, dict[str, Any]]:
|
def _vlm_forward(self, batch: dict[str, Tensor]) -> tuple[Tensor, Tensor]:
|
||||||
"""Training forward pass — computes cross-entropy loss against MC return targets.
|
"""Shared VLM forward: images + text → Gemma3 → last-token hidden → logits.
|
||||||
|
|
||||||
The batch is expected to be preprocessed by the processor pipeline.
|
|
||||||
Keys expected in batch:
|
|
||||||
- observation.images.*: [B, C, H, W] preprocessed images
|
|
||||||
- observation.language_tokens: [B, seq_len] tokenized task prompt
|
|
||||||
- observation.language_attention_mask: [B, seq_len] padding mask
|
|
||||||
- mc_return: [B] normalized Monte Carlo return targets in (-1, 0)
|
|
||||||
- is_terminal: [B] boolean terminal flags
|
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
(loss, output_dict) where loss is scalar cross-entropy
|
(value_logits [B, num_bins], predicted_value [B, 1])
|
||||||
"""
|
"""
|
||||||
images, img_masks, token_ids, text_pad_mask = self._get_model_inputs(batch)
|
images, img_masks, token_ids, text_pad_mask = self._get_model_inputs(batch)
|
||||||
mc_return = batch["mc_return"]
|
|
||||||
is_terminal = batch["is_terminal"]
|
|
||||||
|
|
||||||
text_embs = self.embed_text(token_ids)
|
text_embs = self.embed_text(token_ids)
|
||||||
prefix_embs, prefix_pad_masks = self.embed_prefix(images, img_masks, text_embs, text_pad_mask)
|
prefix_embs, prefix_pad_masks = self.embed_prefix(images, img_masks, text_embs, text_pad_mask)
|
||||||
|
|
||||||
# VLM forward: prefix + [CLS] through Gemma3, then value head
|
|
||||||
batch_size = prefix_embs.shape[0]
|
|
||||||
device = prefix_embs.device
|
|
||||||
|
|
||||||
if self.config.stop_gradient_to_vlm:
|
if self.config.stop_gradient_to_vlm:
|
||||||
prefix_embs = prefix_embs.detach()
|
prefix_embs = prefix_embs.detach()
|
||||||
|
|
||||||
cls_ids = torch.zeros(batch_size, 1, dtype=torch.long, device=device)
|
device = prefix_embs.device
|
||||||
cls_emb = self.cls_embedding(cls_ids)
|
|
||||||
hidden_states = torch.cat([prefix_embs, cls_emb], dim=1)
|
|
||||||
|
|
||||||
cls_pad = torch.ones(batch_size, 1, dtype=torch.bool, device=device)
|
|
||||||
pad_masks = torch.cat([prefix_pad_masks, cls_pad], dim=1)
|
|
||||||
|
|
||||||
prefix_att = torch.zeros(batch_size, prefix_embs.shape[1], dtype=torch.long, device=device)
|
|
||||||
cls_att = torch.ones(batch_size, 1, dtype=torch.long, device=device)
|
|
||||||
att_masks = torch.cat([prefix_att, cls_att], dim=1)
|
|
||||||
|
|
||||||
att_2d = make_att_2d_masks(pad_masks, att_masks)
|
|
||||||
model_dtype = next(self.gemma3.parameters()).dtype
|
model_dtype = next(self.gemma3.parameters()).dtype
|
||||||
|
|
||||||
|
# Bidirectional attention: every valid token attends to every valid token
|
||||||
|
att_2d = prefix_pad_masks[:, None, :] * prefix_pad_masks[:, :, None]
|
||||||
att_4d = torch.where(
|
att_4d = torch.where(
|
||||||
att_2d[:, None, :, :],
|
att_2d[:, None, :, :],
|
||||||
torch.tensor(0.0, dtype=model_dtype, device=device),
|
torch.tensor(0.0, dtype=model_dtype, device=device),
|
||||||
torch.tensor(_ATTENTION_MASK_VALUE, dtype=model_dtype, device=device),
|
torch.tensor(_ATTENTION_MASK_VALUE, dtype=model_dtype, device=device),
|
||||||
)
|
)
|
||||||
|
|
||||||
position_ids = torch.cumsum(pad_masks.long(), dim=1) - 1
|
position_ids = torch.cumsum(prefix_pad_masks.long(), dim=1) - 1
|
||||||
|
|
||||||
if hidden_states.dtype != model_dtype:
|
if prefix_embs.dtype != model_dtype:
|
||||||
hidden_states = hidden_states.to(model_dtype)
|
prefix_embs = prefix_embs.to(model_dtype)
|
||||||
|
|
||||||
outputs = self.gemma3.model(
|
outputs = self.gemma3.model(
|
||||||
inputs_embeds=hidden_states,
|
inputs_embeds=prefix_embs,
|
||||||
attention_mask=att_4d,
|
attention_mask=att_4d,
|
||||||
position_ids=position_ids,
|
position_ids=position_ids,
|
||||||
)
|
)
|
||||||
cls_hidden_state = outputs.last_hidden_state[:, -1, :]
|
|
||||||
|
|
||||||
value_logits = self.value_head(cls_hidden_state)
|
# Readout from last hidden layer
|
||||||
|
if self.config.readout == "mean_pool":
|
||||||
|
hidden = outputs.last_hidden_state
|
||||||
|
mask = prefix_pad_masks.unsqueeze(-1).to(dtype=hidden.dtype)
|
||||||
|
readout = (hidden * mask).sum(dim=1) / mask.sum(dim=1).clamp(min=1)
|
||||||
|
else:
|
||||||
|
readout = outputs.last_hidden_state[:, -1, :]
|
||||||
|
|
||||||
|
value_logits = self.value_head(readout)
|
||||||
value_probs = F.softmax(value_logits, dim=-1)
|
value_probs = F.softmax(value_logits, dim=-1)
|
||||||
predicted_value = (value_probs * self.value_head.bin_centers.to(dtype=value_probs.dtype)).sum(
|
predicted_value = (value_probs * self.value_head.bin_centers.to(dtype=value_probs.dtype)).sum(
|
||||||
dim=-1, keepdim=True
|
dim=-1, keepdim=True
|
||||||
)
|
)
|
||||||
|
return value_logits, predicted_value
|
||||||
|
|
||||||
|
def forward(self, batch: dict[str, Tensor]) -> tuple[Tensor, dict[str, Any]]:
|
||||||
|
"""Training forward pass — cross-entropy loss on MC return targets."""
|
||||||
|
mc_return = batch["mc_return"]
|
||||||
|
is_terminal = batch["is_terminal"]
|
||||||
|
|
||||||
|
value_logits, predicted_value = self._vlm_forward(batch)
|
||||||
|
|
||||||
# Compute target distribution from MC returns
|
# Compute target distribution from MC returns
|
||||||
target_dist = self.compute_target_distribution(
|
target_dist = self.compute_target_distribution(
|
||||||
@@ -539,52 +496,5 @@ class DistributionalVFRewardModel(PreTrainedRewardModel):
|
|||||||
Returns:
|
Returns:
|
||||||
[batch_size] tensor of predicted values V(s).
|
[batch_size] tensor of predicted values V(s).
|
||||||
"""
|
"""
|
||||||
images, img_masks, token_ids, text_pad_mask = self._get_model_inputs(batch)
|
_, predicted_value = self._vlm_forward(batch)
|
||||||
text_embs = self.embed_text(token_ids)
|
|
||||||
prefix_embs, prefix_pad_masks = self.embed_prefix(images, img_masks, text_embs, text_pad_mask)
|
|
||||||
|
|
||||||
# VLM forward: prefix + [CLS] through Gemma3, then value head
|
|
||||||
batch_size = prefix_embs.shape[0]
|
|
||||||
device = prefix_embs.device
|
|
||||||
|
|
||||||
if self.config.stop_gradient_to_vlm:
|
|
||||||
prefix_embs = prefix_embs.detach()
|
|
||||||
|
|
||||||
cls_ids = torch.zeros(batch_size, 1, dtype=torch.long, device=device)
|
|
||||||
cls_emb = self.cls_embedding(cls_ids)
|
|
||||||
hidden_states = torch.cat([prefix_embs, cls_emb], dim=1)
|
|
||||||
|
|
||||||
cls_pad = torch.ones(batch_size, 1, dtype=torch.bool, device=device)
|
|
||||||
pad_masks = torch.cat([prefix_pad_masks, cls_pad], dim=1)
|
|
||||||
|
|
||||||
prefix_att = torch.zeros(batch_size, prefix_embs.shape[1], dtype=torch.long, device=device)
|
|
||||||
cls_att = torch.ones(batch_size, 1, dtype=torch.long, device=device)
|
|
||||||
att_masks = torch.cat([prefix_att, cls_att], dim=1)
|
|
||||||
|
|
||||||
att_2d = make_att_2d_masks(pad_masks, att_masks)
|
|
||||||
model_dtype = next(self.gemma3.parameters()).dtype
|
|
||||||
att_4d = torch.where(
|
|
||||||
att_2d[:, None, :, :],
|
|
||||||
torch.tensor(0.0, dtype=model_dtype, device=device),
|
|
||||||
torch.tensor(_ATTENTION_MASK_VALUE, dtype=model_dtype, device=device),
|
|
||||||
)
|
|
||||||
|
|
||||||
position_ids = torch.cumsum(pad_masks.long(), dim=1) - 1
|
|
||||||
|
|
||||||
if hidden_states.dtype != model_dtype:
|
|
||||||
hidden_states = hidden_states.to(model_dtype)
|
|
||||||
|
|
||||||
outputs = self.gemma3.model(
|
|
||||||
inputs_embeds=hidden_states,
|
|
||||||
attention_mask=att_4d,
|
|
||||||
position_ids=position_ids,
|
|
||||||
)
|
|
||||||
cls_hidden_state = outputs.last_hidden_state[:, -1, :]
|
|
||||||
|
|
||||||
value_logits = self.value_head(cls_hidden_state)
|
|
||||||
value_probs = F.softmax(value_logits, dim=-1)
|
|
||||||
predicted_value = (value_probs * self.value_head.bin_centers.to(dtype=value_probs.dtype)).sum(
|
|
||||||
dim=-1, keepdim=True
|
|
||||||
)
|
|
||||||
|
|
||||||
return predicted_value.squeeze(-1)
|
return predicted_value.squeeze(-1)
|
||||||
|
|||||||
Reference in New Issue
Block a user