PyTorch自定义DQN报错“the derivative for '' is not implemented”求助
解决DQN中loss.backward()触发“the derivative for '' is not implemented”错误
你遇到的这个错误本质是PyTorch反向传播时,计算图里存在无法求导的操作、张量类型不匹配,或者梯度传递逻辑有问题。结合你的代码,我帮你梳理问题点并给出修复方案:
核心问题分析
1. DQN网络输出维度与动作空间不匹配
你的动作空间是4种(action = rd.randint(1, 4)),但DQN最后一层self.head = nn.Linear(20, 1)只输出1个值——常规DQN应该输出每个动作对应的Q值,也就是输出维度要和动作数量一致,这样才能关联动作与对应Q值计算loss,维度不匹配会直接导致反向传播异常。
2. State-Action Q值计算逻辑错误
在optimize_model里,你直接用policy_net(state_batch)得到所有状态的Q值,但没有根据action_batch筛选出对应动作的Q值。正确做法是根据动作索引取出对应动作的Q值,否则loss计算的是全输出与target的差异,逻辑错误且触发求导问题。
3. 张量类型与梯度设置问题
- 你用了旧版PyTorch的
Variable,现在PyTorch 0.4+已经把Variable整合进torch.Tensor,直接用张量加requires_grad参数即可。 action_batch是浮点类型的动作值(比如1.0、2.0),但DQN中动作需要是整数索引才能正确索引Q值数组,浮点类型无法完成索引,也会导致求导失败。
4. 其他潜在问题
done变量未定义,代码中if done:会直接报错;AutoDrive的_select_action返回值逻辑不明确,np.argmax(drive._select_action(0.5, 0.5))可能无法返回有效动作索引。
修复后的完整代码
import sys, math import random as rd import numpy as np import matplotlib import matplotlib.pyplot as plt from collections import namedtuple from itertools import count from PIL import Image import torch import torch.nn as nn import torch.optim as optim import torch.nn.functional as F import torchvision.transforms as T # 定义经验回放的过渡数据结构 Transition = namedtuple('Transition', ('state', 'action', 'next_state', 'reward')) class ReplayMemory(object): def __init__(self, capacity): self.capacity = capacity self.memory = [] self.position = 0 def push(self, *args): """保存一条过渡数据""" if len(self.memory) < self.capacity: self.memory.append(None) self.memory[self.position] = Transition(*args) self.position = (self.position + 1) % self.capacity def sample(self, batch_size): return rd.sample(self.memory, batch_size) def __len__(self): return len(self.memory) # 修改DQN网络输出维度,匹配4个动作的Q值输出 class DQN(nn.Module): def __init__(self): super(DQN, self).__init__() self.l1 = nn.Linear(5, 16) self.l2 = nn.Linear(16, 12) self.l3 = nn.Linear(12, 20) self.head = nn.Linear(20, 4) # 动作空间是4,输出4个动作对应的Q值 def forward(self, x): x = F.relu(self.l1(x)) x = F.relu(self.l2(x)) x = F.relu(self.l3(x)) return self.head(x.view(x.size(0), -1)) # 超参数设置 BATCH_SIZE = 5 GAMMA = 0.999 EPS_START = 0.9 EPS_END = 0.05 EPS_DECAY = 200 TARGET_UPDATE = 5 # 初始化网络与优化器 policy_net = DQN() target_net = DQN() target_net.load_state_dict(policy_net.state_dict()) target_net.eval() optimizer = optim.RMSprop(policy_net.parameters()) memory = ReplayMemory(10000) # 修正优化逻辑 def optimize_model(): if len(memory) < BATCH_SIZE: return transitions = memory.sample(BATCH_SIZE) batch = Transition(*zip(*transitions)) # 转换为张量,无需使用Variable state_batch = torch.cat(batch.state) action_batch = torch.cat(batch.action).long() # 动作转为整数索引 reward_batch = torch.cat(batch.reward) next_state_batch = torch.cat(batch.next_state) # 计算当前状态下对应动作的Q值 state_action_values = policy_net(state_batch).gather(1, action_batch) # 计算下一个状态的最大Q值(目标网络,断开梯度) next_state_values = target_net(next_state_batch).max(1)[0].detach() expected_state_action_values = (next_state_values * GAMMA) + reward_batch # 计算loss,确保维度匹配 loss = F.smooth_l1_loss(state_action_values, expected_state_action_values.unsqueeze(1)) optimizer.zero_grad() loss.backward() # 梯度裁剪防止爆炸 for param in policy_net.parameters(): param.grad.data.clamp_(-1, 1) optimizer.step() # 初始化episode记录(需要你自行实现plot_durations) episode_durations = [] num_episodes = 5 for i_episode in range(num_episodes): # 初始化你的环境 drive = AutoDrive(20, 20, 0, 16, 0) drive._make_observation(0, -1, -1, -1, -1, -1) stand = 3 # epsilon衰减 e = 1. / ((i_episode // 100) + 1) for t in range(stand): # 动作选择:epsilon-greedy if np.random.rand(1) > e: action = rd.randint(1, 4) else: # 这里需要确保_select_action返回的是4个动作的Q值 q_values = drive._select_action(0.5, 0.5) action = np.argmax(q_values) + 1 # 转为1-4的动作 state = drive.state drive._step(action) drive._calc_reward(0.5, 0.5) done = (drive.reward == -10) # 定义结束标志 # 转换张量类型,动作转为0-3的整数索引 state1 = torch.FloatTensor(state).view(1, 5) state2 = torch.FloatTensor(drive.state).view(1, 5) action_tensor = torch.LongTensor([action - 1]).view(1, 1) reward_tensor = torch.FloatTensor([drive.reward]).view(1, 1) memory.push(state1, action_tensor, state2, reward_tensor) optimize_model() if done: episode_durations.append(t + 1) # plot_durations() # 若未实现可先注释 break # 更新目标网络 if i_episode % TARGET_UPDATE == 0: target_net.load_state_dict(policy_net.state_dict())
关键修复点说明
- DQN输出维度修正:最后一层改为输出4个值,对应4种动作的Q值;
- 动作张量类型修正:动作存储为
LongTensor(整数索引),才能用gather取出对应动作的Q值; - loss计算逻辑修正:用
gather筛选对应动作的Q值,目标网络输出用detach()断开梯度,避免更新目标网络参数; - 移除过时的Variable:改用PyTorch原生张量,简化代码并避免版本兼容问题;
- 补充done变量定义:明确游戏结束的判断条件,避免代码报错。
内容的提问来源于stack exchange,提问作者박규진학부생
相关产品推荐
相关产品推荐

