train_utils.py — 训练相关的工具函数¶
文件路径:
verl/experimental/vla/models/openvla_oft/train_utils.py模块路径:verl.experimental.vla.models.openvla_oft.train_utils
文件概述¶
训练相关的工具函数,主要用于计算动作 token 的 mask、token 准确率和 L1 损失。
核心函数¶
get_current_action_mask - 当前动作 mask¶
def get_current_action_mask(labels):
"""找出 labels 中当前动作 chunk 的位置
通过检测非 IGNORE_INDEX 的位置,标记出第一个动作 chunk
(前 ACTION_DIM 个动作 token)。
"""
newline_positions = labels != IGNORE_INDEX
cumsum = torch.cumsum(newline_positions, dim=1)
mask = (1 <= cumsum) & (cumsum <= ACTION_DIM)
action_tokens_only = labels > ACTION_TOKEN_BEGIN_IDX
return action_tokens_only * mask
get_next_actions_mask - 后续动作 mask¶
def get_next_actions_mask(labels):
"""找出 labels 中后续动作 chunk 的位置
标记出第一个动作 chunk 之后的所有动作 token。
"""
newline_positions = labels != IGNORE_INDEX
cumsum = torch.cumsum(newline_positions, dim=1)
mask = cumsum > ACTION_DIM
action_tokens_only = labels > ACTION_TOKEN_BEGIN_IDX
return action_tokens_only * mask
训练指标¶
def compute_token_accuracy(logits, labels, mask):
"""计算动作 token 的预测准确率"""
predicted = logits.argmax(dim=-1)
correct = (predicted == labels) & mask
return correct.sum() / mask.sum()
def compute_actions_l1_loss(predicted_actions, target_actions):
"""计算连续动作的 L1 损失(用于回归模式)"""
return F.l1_loss(predicted_actions, target_actions)
核心函数列表¶
| 名称 | 说明 |
|---|---|
get_current_action_mask |
当前 action chunk 的 mask |
get_next_actions_mask |
后续 action chunks 的 mask |
compute_token_accuracy |
动作 token 准确率 |
compute_actions_l1_loss |
L1 损失(回归模式) |
与其他模块的关系¶
- 被
modeling_prismatic.py的_process_action_masks调用 - 使用
constants.py中的ACTION_DIM、ACTION_TOKEN_BEGIN_IDX等常量
小结¶
这些工具函数解决了一个核心问题:在混合了文本和动作的序列中,如何精确定位动作 token 的位置,以便只对这些位置计算损失。