保存带 FP8 元数据的 checkpoint 时报了 CUDA 相关的错,翻到 module/base.py 发现 get_extra_state / set_extra_state 这条链路上还留着三处 CUDA 专用调用没改。
module/base.py
get_extra_state
set_extra_state
transformer_engine/pytorch/module/base.py
793 行,get_extra_state() 里:
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() 的旧格式兼容分支:
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()
map_location="cuda"
copy_tensor
dst.copy_(src, non_blocking=True)
第二种情况比第一种更麻烦,因为它不报错,症状是 resume 之后 FP8 的 scaling 从一个不对的状态接着跑,loss 会有一个说不清来源的跳变。
get_extra_state / set_extra_state 里按 fp8_meta 张量实际所在的设备取同步函数,别写死:
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(),或者干脆按当前模块的设备来。
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 不可用。
CudaRNGStatesTracker.add_()
torch.cuda.manual_seed(seed)
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 的性能。
device_type="cuda"
ops/basic/rmsnorm.py:167
device_type="npu"
torch.is_autocast_enabled("cuda")
weight.dtype
要不要把 torch.cuda.* 在整个 transformer_engine/pytorch/ 下扫一遍统一处理掉?如果方向没问题我可以整理一份完整清单。
torch.cuda.*
transformer_engine/pytorch/
感谢您的提出,这个由torch_npu的transfer_to_npu保证。
保存带 FP8 元数据的 checkpoint 时报了 CUDA 相关的错,翻到
module/base.py发现get_extra_state/set_extra_state这条链路上还留着三处 CUDA 专用调用没改。位置
transformer_engine/pytorch/module/base.py793 行,
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_serialized814 行,
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-65def 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-51def 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/下扫一遍统一处理掉?如果方向没问题我可以整理一份完整清单。