跳转至

torch_dtypes.py — PyTorch 数据类型工具

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

文件概述

torch_dtypes.py 提供字符串到 PyTorch 数据类型的映射工具。在配置文件中,数据类型通常以字符串形式指定(如 "bf16"),需要转换为 torch.bfloat16 等实际类型。

核心函数

PrecisionType = {
    "fp16": torch.float16,
    "bf16": torch.bfloat16,
    "fp32": torch.float32,
    ...
}

提供一个字典映射,将配置中常见的精度字符串转换为对应的 torch.dtype。

与其他模块的关系

  • 被模型创建和 FSDP 配置时使用
  • fsdp_utils.py 中设置混合精度时引用

小结

简单但必要的工具文件,消除了字符串和 PyTorch dtype 之间的转换样板代码。