uint8类型张量能否用于需调用.backward()的损失函数?
uint8张量在PyTorch损失函数与反向传播中的问题解答
问题重现
用户代码如下:
import torch import torch.nn as nn a = torch.randn(3, 3, dtype=torch.float32, requires_grad=True) b = torch.randint(0, 256, (3, 3), dtype=torch.uint8) loss = nn.MSELoss()(a, b) print(loss.dtype) # 输出torch.float32 loss.backward() # 报错RuntimeError: Found dtype Byte but expected Float
同时发现uint8张量无法设置requires_grad=True,报错:RuntimeError: Only Tensors of floating point and complex dtype can require gradients
核心问题解答
1. uint8张量能否用于需调用.backward()的损失函数?
不能直接使用,必须先显式转换为浮点型(如float32/float64)后,才能用于需要反向传播的场景。
2. 为什么loss是float32,但backward()会报Byte类型错误?
MSELoss前向计算时确实会自动对输入做类型提升,将uint8的b转换为float32来计算损失,所以最终loss的类型是float32。但问题出在计算图的构建上:
- 原始张量
b的uint8类型信息会被保留在计算图中,当调用backward()时,PyTorch的自动求导系统会检查所有参与运算的张量类型。 - uint8属于整数类型,完全不支持梯度相关的操作(即使
b没有开启requires_grad=True),求导系统在处理对应节点时,会因为类型不兼容触发错误。
3. 为什么uint8张量无法设置requires_grad=True?
PyTorch的自动求导系统仅支持浮点型和复数型张量存储梯度——因为梯度本质是浮点数,整数类型无法表示小数形式的梯度值,这是框架的底层设计限制,所有整数类型(包括uint8、int32等)都不允许开启requires_grad=True。
解决方案
将uint8张量显式转换为浮点型后再传入损失函数:
import torch import torch.nn as nn a = torch.randn(3, 3, dtype=torch.float32, requires_grad=True) b = torch.randint(0, 256, (3, 3), dtype=torch.uint8).float() # 显式转换为float32 loss = nn.MSELoss()(a, b) print(loss.dtype) # 输出torch.float32 loss.backward() # 正常运行
内容的提问来源于stack exchange,提问作者Aria Lovelace
相关产品推荐
相关产品推荐

