PyTorch强化学习代码反向传播出现Double/Float类型报错如何解决
问题根因
你看到的loss本身是float32属于正向传播的最终输出类型,报错是发生在反向传播的梯度计算阶段:计算图中存在某个中间张量、模型参数或者输入张量是torch.float64(Double类型),和其他torch.float32(Float类型)的张量运算时触发了类型不匹配,和loss最终的输出类型没有必然关联。
排查与解决步骤
- 第一步:排查输入张量的类型
打印从make_batch返回的所有输入的dtype:s、a、r、s_prime、done_mask、prob_a,重点关注prob_a、r、done_mask三个变量,只要其中一个是Double类型,运算过程中就会触发隐式类型提升,导致计算图中存在Double张量。 - 第二步:修正优势函数计算的类型隐患
你代码中初始化advantage = 0.0时用的是Python原生浮点数,默认是64位精度,和delta中的32位浮点数运算后结果会自动转成64位,哪怕后续转张量时指定了dtype=torch.float,也可能存在隐式转换问题。可以把初始化逻辑改成:advantage = torch.tensor(0.0, dtype=torch.float32) - 第三步:排查模型参数的类型
打印策略网络和价值网络的参数类型,确认没有被全局转成Double:
如果输出是print(next(self.pi_ap.parameters()).dtype) print(next(self.v_ap.parameters()).dtype)torch.float64,要么是你之前调用过.double()方法,要么是全局配置了torch.set_default_dtype(torch.float64),改回默认的float32即可。 - 第四步:临时快速兼容方案
如果不想逐个排查根因,可以在计算loss前统一把所有输入转成float32:
也可以在反向传播时强制转换loss的类型:s = s.float() r = r.float() s_prime = s_prime.float() done_mask = done_mask.float() prob_a = prob_a.float()loss.mean().float().backward()
内容的提问来源于stack exchange,提问作者Leee
相关产品推荐
相关产品推荐

