已关闭
FP8 extra_state 的存取路径仍在调用 torch.cuda.synchronize 与 map_location="cuda" #28
崇理战队创建于  8月13日关闭于  17 天前
崇理战队
8月13日 创建

保存带 FP8 元数据的 checkpoint 时报了 CUDA 相关的错,翻到 module/base.py 发现 get_extra_state / set_extra_state 这条链路上还留着三处 CUDA 专用调用没改。

位置

transformer_engine/pytorch/module/base.py

793 行,get_extra_state() 里:

        # Serialize state into byte tensor
        torch.cuda.synchronize()
        state_serialized = bytearray(pickle.dumps(state))
        state_serialized = torch.frombuffer(state_serialized, dtype=torch.uint8)
        return state_serialized

814 行,set_extra_state() 的旧格式兼容分支:

        elif isinstance(state, io.BytesIO):
            # Deprecated format with io.BytesIO
            state.seek(0)
            state = torch.load(state, map_location="cuda")

871 行,set_extra_state() 末尾:

            copy_tensor(state["amax_history_bwd"], self.fp8_meta["scaling_bwd"].amax_history)
        torch.cuda.synchronize()

分析

这三处的意图都能看出来:pickle 之前要等 device 上的 amax / scale 算完,copy 之后要等拷贝落地。问题是这个仓库的目标设备是 NPU:

  • torch.cuda.synchronize() 同步的是 CUDA 流,对 NPU 的计算流没有任何屏障作用。纯 NPU 环境(没装 CUDA / 没有 N 卡)下这句会直接抛异常,checkpoint 保存不了;就算机器上恰好有 CUDA 环境不报错,屏障也是空的——pickle 的时候 NPU 侧的 amax_history 可能还没算完,写出去的是撕裂状态。
  • map_location="cuda" 会把字节流里的张量搬到 CUDA 上。纯 NPU 环境直接报错;不报错的话恢复出来的 scale / amax 落在 CUDA,而模块权重在 NPU,后面 copy_tensor 做的是 dst.copy_(src, non_blocking=True) 跨设备异步拷贝,再配上 871 行那个同步错设备的屏障,等于完全没保证。

第二种情况比第一种更麻烦,因为它不报错,症状是 resume 之后 FP8 的 scaling 从一个不对的状态接着跑,loss 会有一个说不清来源的跳变。

建议

get_extra_state / set_extra_state 里按 fp8_meta 张量实际所在的设备取同步函数,别写死:

device = self.fp8_meta["scaling_fwd"].scale.device
if device.type == "npu":
    torch.npu.synchronize()
elif device.type == "cuda":
    torch.cuda.synchronize()

map_location="cuda" 建议改成 map_location=torch.get_default_device(),或者干脆按当前模块的设备来。

另外两处同类残留

transformer_engine/pytorch/distributed.py:58-65

def graph_safe_rng_available() -> bool:
    """Returns whether cuda graph safe RNG state manipulation is supported."""
    return (
        hasattr(torch.cuda.CUDAGraph, "register_generator_state")
        and hasattr(torch.Generator, "graphsafe_set_state")
        and hasattr(torch.Generator, "graphsafe_get_state")
        and hasattr(torch.Generator, "clone_state")
    )

检查的全是 CUDA 专属 API,NPU 上恒为 False,于是 CudaRNGStatesTracker.add_() 必走 else 分支,1346 行调 torch.cuda.manual_seed(seed)——设的是 CUDA 默认生成器的种子。TP 场景下 dropout / recompute 的随机源就没被隔离,纯 NPU 上则是直接报 CUDA 不可用。

transformer_engine/pytorch/ops/_common.py:38-51

def maybe_autocast_dtype(
    *,
    device_type: str = "cuda",
    default_dtype: Optional[torch.dtype] = None,
) -> torch.dtype:

默认 device_type="cuda",而 ops/basic/rmsnorm.py:167 调用时没传这个参数。NPU 上 autocast 是以 device_type="npu" 开的,torch.is_autocast_enabled("cuda") 恒为 False,函数直接退回 weight.dtype。结果是开了 bf16 autocast,RMSNorm 这一段还在 fp32 上算,跟同一层其它算子的精度对不齐,也白丢了 bf16 的性能。

要不要把 torch.cuda.* 在整个 transformer_engine/pytorch/ 下扫一遍统一处理掉?如果方向没问题我可以整理一份完整清单。

likedislike
clc2025成员
17 天前 评论:

感谢您的提出,这个由torch_npu的transfer_to_npu保证。

likedislike
Cclc2025成员
17 天前 issue状态由 TODO 改变为 DONE
Cclc2025成员
17 天前 关闭了 issue
ascend-robotascend-robot成员
17 天前 添加了label:resolved