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

PyTorch实现DQN训练CartPole-v1失效问题求助

DQN训练CartPole-v1失效问题排查

问题描述

我之前训练DQN能在约65000次迭代后解决CartPole-v1环境,但现在完全失效,达不到原有效果。按过往经验调优超参数后还是没用,训练时损失会爆炸,延长训练时长、调整硬更新tau也解决不了问题。相关代码如下:

完整代码

主训练脚本

import gym
import numpy as np

import torch
from torch import nn
from torch.nn import functional as F
from torch import optim

from models import DQN
from memory import Memory
from utils import wrap_input, epsilon_greedy

def main() -> int:
    env = gym.make("CartPole-v1")
    device = torch.device("cuda" if torch.cuda.is_available() else "cpu")

    # Online and offline model for learning
    model = DQN(env.observation_space, env.action_space, 24).to(device)
    target = DQN(env.observation_space, env.action_space, 24).to(device)
    target.eval()

    # Optimizer and loss function
    optimizer = optim.Adam(model.parameters(), lr=.001)
    loss_fn = F.smooth_l1_loss


    memory = Memory(10_000)
    obs, info = env.reset()

    for it in range(65_000):
        # Do this for the batch norm
        model.eval()

        # Maybe explore
        if np.random.random() <= epsilon_greedy(1.0, .01, 15_000, it):
            state = wrap_input(obs, device).unsqueeze(0)
            action  = model(state).argmax().item()
        else:
            action = env.action_space.sample()

        # Act in environment and store the memory
        next_state, reward, done, truncated, info = env.step(action)
        if truncated or done:
            next_state = np.zeros(env.observation_space.shape)
        memory.store([obs, action, reward, int(done), next_state])
        done = done or truncated

        if done:
            obs, info = env.reset()

        # Train
        if len(memory) > 32:
            model.train()
            states, actions, rewards, dones, next_states = memory.sample(32)

            # Wrap and move all values to the cpu
            states = wrap_input(states, device)
            actions = wrap_input(actions, device, torch.int64, reshape=True)
            next_states = wrap_input(next_states, device)
            rewards = wrap_input(rewards, device, reshape=True)
            dones = wrap_input(dones, device, reshape=True)

            # Get current q-values
            qs = model(states)
            qs = torch.gather(qs, dim=1, index=actions)

            # Compute target q-values
            with torch.no_grad():
                next_qs, _ = target(next_states).max(dim=1)
                next_qs = next_qs.reshape(-1, 1)

            target_qs = rewards + .9 * (1 - dones) * next_qs.reshape(-1, 1)

            # Compute loss
            loss = loss_fn(qs, target_qs)
            optimizer.zero_grad()
            loss.backward()
            
            # Clip gradients
            nn.utils.clip_grad_norm_(model.parameters(), 1)

            # Backprop
            optimizer.step()

            # soft update
            with torch.no_grad():
                for target_param, local_param in zip(target.parameters(), model.parameters()):
                    target_param.data.copy_(1e-2 * local_param.data + (1 - 1e-2) * target_param.data)


        if it % 200 == 0:
            target.load_state_dict(model.state_dict())

models.py

class FlatExtractor(nn.Module):
    '''Does nothing but pass the input on'''
    def __init__(self, obs_space):
        super(FlatExtractor, self).__init__()

        self.n_flatten = obs_space.shape[0]

    def forward(self, obs):
        return obs


class DQN(nn.Module):
    def __init__(self, obs_space, act_space, layer_size):
        super(DQN, self).__init__()

        # Feature extractor
        if len(obs_space.shape) == 1:
            self.feature_extractor = FlatExtractor(obs_space)
        elif len(obs_space.shape) == 3:
            self.feature_extractor = NatureCnn(obs_space)
        else:
            raise NotImplementedErorr("This type of environment is not supported")

        # Neural network
        self.net = nn.Sequential(
            nn.Linear(self.feature_extractor.n_flatten, layer_size),
            nn.BatchNorm1d(layer_size),
            nn.ReLU(),
            nn.Linear(layer_size, layer_size),
            nn.BatchNorm1d(layer_size),
            nn.ReLU(),
            nn.Linear(layer_size, act_space.n),
        )

    def forward(self, obs):
        return self.net(self.feature_extractor(obs))

memory.py

import random
from collections import deque

class Memory(object):
    def __init__(self, maxlen):
        self.memory = deque(maxlen=maxlen)

    def store(self, experience):
        self.memory.append(experience)

    def sample(self, n_samples):
        return zip(*random.sample(self.memory, n_samples))

    def __len__(self):
        return len(self.memory)

utils.py

def wrap_input(arr, device, dtype=torch.float, reshape=False):
    output = torch.from_numpy(np.array(arr)).type(dtype).to(device)
    if reshape:
        output = output.reshape(-1, 1)

    return output

def epsilon_greedy(start, end, n_steps, it):
    return max(start - (start - end) * (it / n_steps), end)

关键问题排查

  • ε-贪婪策略逻辑完全颠倒:当前代码中,当随机数小于等于ε时使用模型选动作,否则随机探索。正确逻辑应该是:随机数≤ε时随机探索,否则用模型选最优动作。这会导致前期几乎不探索(ε初始为1.0,此时全用模型选动作,但模型完全没训练,输出随机),后期ε降到0.01时反而大量探索,彻底破坏探索-利用平衡,是训练失效的核心原因。
  • 同时混用软更新与硬更新:代码中既做了每步的软更新(tau=0.01),又每隔200次迭代强制硬更新target网络。两种更新方式只需二选一:要么固定频率硬更新,要么每步软更新。同时用会导致target网络更新混乱,破坏DQN目标值的稳定性,引发损失爆炸。
  • BatchNorm1d不适合当前场景:CartPole推理时是单样本输入,BatchNorm需要批量数据才能计算有效的均值和方差。训练时的批量统计量在前期样本不足时误差极大,推理时model.eval()固定的统计量会进一步导致动作选择异常,污染经验池。对于这种简单环境,建议直接去掉BatchNorm,或改用LayerNorm。
  • 终止状态next_state处理可优化:将终止状态的next_state设为全0的做法,结合BatchNorm可能加剧数值不稳定。更稳妥的方式是保留原始next_state,通过(1 - dones)掩码让终止状态的target_qs直接等于rewards,符合DQN目标公式的原始定义。
  • 经验采样维度需确认:memory.sample()返回的是zip后的元组,wrap_input中np.array(arr)是否能正确将多个obs转换为(32,4)的张量?虽然这不是核心问题,但如果维度错误也会导致训练失败,可添加打印语句验证。

内容的提问来源于stack exchange,提问作者Squeemos

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.03 20:16:08