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 上都能运行。