You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.25 15:52:36