跳转至

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 的位置,以便只对这些位置计算损失。