跳转至

attention_utils.py — 注意力计算工具

文件路径: verl/utils/attention_utils.py

文件概述

统一封装 CUDA 和 NPU 后端的注意力相关函数(pad_input、unpad_input、index_first_axis、rearrange),使上层代码不需要关心底层硬件。

核心机制

def _get_attention_functions():
    """根据硬件动态导入注意力函数"""
    if is_torch_npu_available(check_device=False):
        from verl.utils.npu_flash_attn_utils import index_first_axis, pad_input, rearrange, unpad_input
    else:
        from flash_attn.bert_padding import index_first_axis, pad_input, rearrange, unpad_input
    return _index_first_axis, _pad_input, _rearrange, _unpad_input

四个导出函数(index_first_axis、pad_input、rearrange、unpad_input)都是延迟加载:首次调用时检测硬件并导入对应实现。

核心函数

函数 说明
unpad_input() 去除 padding,返回紧凑张量和索引
pad_input() 将紧凑张量恢复为带 padding 的批次
index_first_axis() 按索引选取第一维
rearrange() 张量重排(einops 风格)

与其他模块的关系

  • 被 torch_functional.py 中的 log-prob 计算使用
  • CUDA 后端依赖 flash_attn 库,NPU 后端依赖 npu_flash_attn_utils.py

小结

硬件无关的注意力工具层,确保相同的代码在 GPU 和 NPU 上都能运行。