如何以原始二进制格式保存bfloat16类型的torch.tensor?
保存bfloat16类型torch.Tensor为原始二进制文件的方法
因为numpy不支持bfloat16类型,所以不能通过转numpy数组的方式保存。直接使用PyTorch自身的API即可完成原始二进制文件的读写,以下是两种惯用方法:
方法一:直接写入原始字节流
通过torch.Tensor.tobytes()获取tensor的原始字节数据,直接写入文件,这是最直接的方式:
import torch # 创建bfloat16类型的tensor tensor = torch.tensor([1, 2, 3, 4, 5]).bfloat16() # 写入原始二进制文件 with open("output.bin", "wb") as f: f.write(tensor.tobytes())
方法二:读取恢复成bfloat16 tensor的配套写法
如果需要从保存的二进制文件恢复tensor,需要明确原tensor的形状和数据类型,示例代码如下:
import torch # 读取二进制文件内容 with open("output.bin", "rb") as f: raw_bytes = f.read() # 从字节数据恢复为bfloat16 tensor,并reshape到原形状 restored_tensor = torch.frombuffer(raw_bytes, dtype=torch.bfloat16).reshape(5) print(restored_tensor) # 输出:tensor([1., 2., 3., 4., 5.], dtype=torch.bfloat16)
原代码报错原因
你的代码中tensor.numpy()抛出错误,是因为numpy没有原生支持bfloat16数据类型,PyTorch无法将bfloat16类型的tensor转换为numpy数组,因此需要绕开numpy,直接使用PyTorch的文件操作API。
内容的提问来源于stack exchange,提问作者flexwang
相关产品推荐
相关产品推荐

