训练DQN时GPU利用率0%CPU满负载,添加to.device后设备不匹配报错
DQN训练GPU利用率为0%及设备不匹配错误的解决方案
核心问题分析
- 初始GPU利用率为0:模型虽已移至GPU,但训练数据(状态、经验池内的张量)仍在CPU上,所有计算实际在CPU执行,导致GPU闲置。
- 添加
to.device后报错:部分张量停留在CPU,部分在GPU,运算时出现设备不匹配冲突。
具体修复步骤
1. 采样后统一张量设备
从经验池采样后,立即将所有张量转移到指定设备,避免模型与数据设备不一致:
if len(replay_buffer) > batch_size: states, actions, rewards, next_states, dones = replay_buffer.sample(batch_size) # 新增:将所有张量转移到目标设备 states = states.float().to(device) actions = actions.to(device) rewards = rewards.to(device) next_states = next_states.float().to(device) dones = dones.to(device) action_indices = torch.multinomial(actions, 1).squeeze(1).long() current_q = dqn(states).gather(1, action_indices.unsqueeze(1)).squeeze(1) next_q = target_dqn(next_states).max(1)[0].detach() expected_q = rewards + gamma * next_q * (~dones) loss = nn.functional.mse_loss(current_q, expected_q.float()) optimizer.zero_grad() loss.backward() optimizer.step()
2. 环境输出数据强制转移到GPU
确保环境生成的状态、奖励等数据均为PyTorch张量并转移到GPU:
- 修改
state_reset()函数:
def state_reset(): # 原状态生成逻辑 state = ... # 转为张量并移至设备 return torch.tensor(state, dtype=torch.float32).to(device)
environment_step()返回的next_state、get_reward()返回的reward需做同样处理,转为张量后移至GPU。
3. 经验池存储统一设备数据
修改经验池的push方法,确保存入的所有数据均在目标设备上:
def push(self, state, action, reward, next_state, done): # 非张量数据转为张量并移至设备 state = torch.tensor(state, dtype=torch.float32).to(device) if not isinstance(state, torch.Tensor) else state.to(device) action = torch.tensor(action, dtype=torch.float32).to(device) if not isinstance(action, torch.Tensor) else action.to(device) reward = torch.tensor(reward, dtype=torch.float32).to(device) if not isinstance(reward, torch.Tensor) else reward.to(device) next_state = torch.tensor(next_state, dtype=torch.float32).to(device) if not isinstance(next_state, torch.Tensor) else next_state.to(device) done = torch.tensor(done, dtype=torch.bool).to(device) if not isinstance(done, torch.Tensor) else done.to(device) self.buffer.append((state, action, reward, next_state, done))
4. 检查take_action函数输入设备
确保输入模型的状态已在GPU上:
def take_action(state, epsilon, dqn): if not isinstance(state, torch.Tensor): state = torch.tensor(state, dtype=torch.float32).to(device) # 原动作选择逻辑...
验证GPU利用率
修复后运行代码,若GPU仍闲置,需排查:
- 是否所有环境生成的状态、奖励都已转为张量并移至GPU;
- 经验池内存储的所有数据是否均在目标设备上。
内容的提问来源于stack exchange,提问作者Airen Chan
相关产品推荐
相关产品推荐

