引言

强化学习是一种机器学习分支,通过智能体与环境的交互来学习决策策略。DQN(Deep Q-Network)算法作为一种经典的强化学习算法,被成功应用于训练智能体玩各种游戏。本文将介绍如何使用DQN算法训练智能体玩CartPole游戏,展示强化学习的魅力和潜力。

在过去的几年中,强化学习在解决复杂问题方面取得了显著的进展。其中,DQN算法作为深度强化学习的代表之一,通过结合深度神经网络和经验回放技术,成功地解决了许多挑战性的任务。本文将基于Intel OneAPI围绕DQN算法在训练智能体玩CartPole游戏中的应用展开讨论。

Intel oneAPI 是一个跨行业、开放、基于标准的统一的编程模型,它为跨 CPU、GPU、FPGA、专用加速器的开发者提供统一的体验。它由一项行业计划和一款英特尔beta产品组成。oneAPI 开放规范基于行业标准和现有开发者编程模型,广泛适用于不同架构和来自不同供应商的硬件。oneAPI 行业计划鼓励生态系统内基于oneAPI规范的合作以及兼容 oneAPI的实践。通过oneAPI,我们可以最大程度地忽略硬件之间的差异,而使用统一的编程模型编写程序,有效地提高我们的开发效率

CartPole游戏简介:

首先,我们将介绍CartPole游戏的背景和规则。CartPole是一个简单的强化学习环境,在该游戏中,智能体需要通过控制小车的移动来保持杆子的平衡。

游戏中的状态由四个连续变量组成:车的水平位置、车的速度、杆子的角度和杆子的角速度。玩家可以选择向左或向右施加力来控制车的移动方向。

游戏有几个重要的限制条件:如果杆子与垂直线的夹角超过15度,或者车的位置超过屏幕边界,或者游戏运行时间(步数)超过200步,游戏就会结束。

我们充分利用oneAPI提供的统一的编程模型,基于pytorch开发的CartPole游戏可以·使用强化学习算法来训练智能体(agent)学习如何控制车来保持杆子的平衡。利用了英特尔CPU上的AVX-512矢量神经网络指令(AVX512 VNNI)和英特尔Xe高级矩阵扩展,以及英特尔GPU上的Xe矩阵扩展(XMX)AI引擎对AI训练进行加速。使得我们可以通过简洁的API接口来实现高性能的并行计算和加速器编程。充分利用多核CPU、GPU和FPGA等硬件加速器的计算能力,提高CartPole游戏的性能和效率。

DQN算法原理:

接下来,我们将介绍DQN算法的原理和核心概念。DQN(Deep Q-Network)是一种基于深度神经网络的强化学习算法,用于解决具有高维状态空间和离散动作空间的问题。DQN算法结合了Q-learning算法和深度神经网络的思想,通过学习状态-动作值函数(Q函数)来指导智能体做出最优决策。

以下是DQN算法的基本原理:

  1. 状态表示:DQN算法将环境的状态作为输入,通常使用向量、图像或其他形式的特征表示。这些特征可以包括环境的观测值、历史状态等信息。
  2. Q函数的近似:DQN算法使用一个深度神经网络来近似状态-动作值函数Q(s, a),其中s是状态,a是动作。网络的输入是状态表示,输出是每个动作的Q值估计。通过训练神经网络,我们可以学习到状态下各个动作的预期回报。
  3. 经验回放:为了解决数据相关性的问题,DQN算法采用经验回放(experience replay)技术。在每次与环境交互时,将状态、动作、奖励、下一个状态等转换存储到一个经验回放缓冲区中。然后从缓冲区中随机采样一批数据进行训练,以打破时间上的相关性。
  4. 目标Q值:DQN算法引入一个目标网络(target network)来计算目标Q值。目标网络是与主网络(即Q网络)结构相同但参数独立的网络。在每个训练步骤中,通过固定一段时间的目标网络参数,以减少Q值估计的更新目标的变化。
  5. 动作选择:为了平衡探索和利用的关系,DQN算法使用ε-贪婪策略。在训练过程中,以ε的概率随机选择一个动作,以(1-ε)的概率选择具有最高Q值的动作。随着训练的进行,ε逐渐减小,使智能体在开始时更多地进行探索,在后期更多地进行利用。
  6. Q-learning更新:DQN算法使用Q-learning的更新规则来优化网络参数。通过最小化预测Q值与目标Q值之间的误差来调整网络参数。

通过迭代以上步骤,DQN算法能够逐渐学习到最优的Q函数估计,实现智能体在复杂环境中做出最佳决策的能力。DQN算法的创新之处在于使用深度神经网络对高维状态空间进行建模,从而解决了传统Q-learning在高维问题上的限制。这使得DQN算法在许多领域都取得了显著的成功,并成为强化学习的重要里程碑之一。

训练步骤:

在本部分,我们将使用Intel oneAPI的pytorch extension编写并训练一个 Deep Q-Learning 网络。进行DQN算法的简单实现,构建深度神经网络模型、初始化经验回放缓冲区、定义损失函数和优化器等。

依赖

本项目需要intel pytorch extension、pytorch、matplotlib、gymnasium等库才能运行。

安装

pip install matplotlib gymnasium torch intel_extension_for_pytorch

使用XPU加速运算

# 引入oneAPI中的Intel Extension for Pytorch

import intel_extension_for_pytorch as ipex

# 选择xpu作为torch运算硬件

device = torch.device("xpu")

运行

本项目不需要额外配置,直接运行python文件即可。

python dqn.py

运行时,会显示matplot的图表,显示在一个episode中坚持了多少个step控制台中会显示当前episode的结果。

代码

import gymnasium as gym
import random
import matplotlib
import matplotlib.pyplot as plt
from collections import deque, namedtuple

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

# 引入oneAPI中的Intel Extension for Pytorch
import intel_extension_for_pytorch as ipex

batchSize = 128
minEpsilon = 0.05
dEpsilon = 1e-3
learningRate = 1e-4
gamma = 0.99

env = gym.make("CartPole-v1", render_mode="rgb_array")

# 选择xpu作为torch运算硬件
device = torch.device("xpu")

Transition = namedtuple('Transition', ('state', 'action', 'next_state', 'reward'))

class DQN(nn.Module):
    def __init__(self, stateSize, actionSize):
        super(DQN, self).__init__()
        self.fc1 = nn.Linear(stateSize, 128)
        self.fc2 = nn.Linear(128, 128)
        self.fc3 = nn.Linear(128, actionSize)

    def forward(self, x):
        # output.shape: batch_size*n_actions, state_action_value
        return self.fc3(F.relu(self.fc2(F.relu(self.fc1(x)))))


class ReplayMemory(object):

    def __init__(self, capacity):
        self.memory = deque([], maxlen=capacity)

    def push(self, *args):
        self.memory.append(Transition(*args))

    def sample(self, batch_size):
        return random.sample(self.memory, batch_size)

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


class Agent:
    def __init__(self, stateSize, actionSize):
        self.actionSize = actionSize
        self.stateSize = stateSize
        # 一开始policyNN和targetNN参数相同
        self.policyNN = DQN(stateSize, actionSize).to(device)
        self.targetNN = DQN(stateSize, actionSize).to(device)
        self.targetNN.load_state_dict(self.policyNN.state_dict())
        # AdamW优化器
        self.optimiser = optim.AdamW(self.policyNN.parameters(), lr=learningRate, amsgrad=True)

        self.nowEpsilon = 1
        self.memory = ReplayMemory(10000)

    # 用epsilon-greedy选择explore还是exploit,state要是tensor
    def selectAction(self, nowState):
        a = random.random()
        self.nowEpsilon = self.nowEpsilon - dEpsilon if self.nowEpsilon - dEpsilon > minEpsilon else minEpsilon
        if a < self.nowEpsilon:
            # explore
            return random.randint(0, self.actionSize - 1)
        else:
            # exploit
            with torch.no_grad():
                return self.policyNN(nowState).max(1)[1].item()

    def recordTransition(self, state, action, next_state, reward):
        self.memory.push(state, action, next_state, reward)

    def updateQ(self):
        if len(self.memory) < batchSize:
            return
        # 对policyNN(Q)进行更新,隔c次再将targetNN与policyNN同步
        transitions = self.memory.sample(batchSize)
        batch = Transition(*zip(*transitions))

        # 计算目标tensor
        notNoneNextStateMask = [False] * batchSize
        for i, t in enumerate(transitions):  # 筛选出St+1不为None的state的mask
            if t.next_state is not None:
                notNoneNextStateMask[i] = True

        # 所有不为空的nextState组成的列向量
        notNoneNextStates = torch.cat([s for s in batch.next_state if s is not None])
        nextStateReward = torch.zeros(batchSize, device=device)  # fi[i]
        with torch.no_grad():
            nextStateReward[notNoneNextStateMask] = self.targetNN(notNoneNextStates).max(1)[0]

        # 计算reward列向量
        batchReward = torch.cat(batch.reward)
        batchState = torch.cat(batch.state)
        batchAction = torch.tensor(batch.action, device=device).unsqueeze(-1)

        currentQ = self.policyNN(batchState).gather(1, batchAction)
        y = batchReward + (gamma * nextStateReward)

        # 计算loss
        loss = F.huber_loss(currentQ, y.unsqueeze(1))

        # 梯度下降
        self.optimiser.zero_grad()
        loss.backward()
        self.optimiser.step()

        # soft update
        targetArgs = self.targetNN.state_dict()
        policyArgs = self.policyNN.state_dict()
        for key in policyArgs:
            targetArgs[key] = policyArgs[key] * 0.005 + targetArgs[key] * (1 - 0.005)
        self.targetNN.load_state_dict(targetArgs)


episode_durations = []

def plot_durations(show_result=False):
    plt.figure(1)
    durations_t = torch.tensor(episode_durations, dtype=torch.float)
    if show_result:
        plt.title('Result')
    else:
        plt.clf()
        plt.title('Training...')
    plt.xlabel('Episode')
    plt.ylabel('Duration')
    plt.plot(durations_t.numpy())
    # Take 100 episode averages and plot them too
    if len(durations_t) >= 100:
        means = durations_t.unfold(0, 100, 1).mean(1).view(-1)
        means = torch.cat((torch.zeros(99), means))
        plt.plot(means.numpy())

    plt.pause(0.01)  # pause a bit so that plots are updated


if __name__ == "__main__":
    # 获取action与state数
    actionCount = env.action_space.n
    state, info = env.reset()
    stateCount = len(state)

    max_episodes = 600
    complete_episodes = 0
    finished_flag = False
    agent = Agent(stateCount, actionCount)
    for nowEpisode in range(max_episodes):
        # 训练一轮初始化一次gym
        state, info = env.reset()
        state = torch.tensor(state, dtype=torch.float32, device=device).unsqueeze(0)
        nowStep = 0
        while True:
            # 选择一个操作,然后记录4个参数到memory里面,并更新网络参数
            action = agent.selectAction(state)
            observation, reward, terminated, truncated, _ = env.step(action)
            reward = torch.tensor([reward], device=device)
            done = terminated or truncated

            if terminated:
                nextState = None
            else:
                nextState = torch.tensor(observation, dtype=torch.float32, device=device).unsqueeze(0)

            agent.recordTransition(state, action, nextState, reward)

            state = nextState
            agent.updateQ()

            if done:
                episode_durations.append(nowStep + 1)
                print(str(nowEpisode) + "  " + str(nowStep))
                plot_durations()
                break

            nowStep += 1


plot_durations(show_result=True)
plt.ioff()
plt.show()

实验结果:

 可以发现,AI很好地掌握了杆平衡的方法,通过Intel oneAPI训练所需的时间也较纯CPU训练短。

结论:

DQN算法的成功应用于训练智能体玩CartPole游戏展示了强化学习在解决复杂问题方面的潜力。通过深度神经网络的近似和经验回放技术的利用,DQN算法能够在训练过程中提升智能体的表现,并为其他更复杂的任务奠定基础。同时通过oneAPI,我们可以最大程度地忽略硬件之间的差异,而使用统一的编程模型编写程序,有效地提高我们的开发效率未来,随着强化学习算法的不断发展,我们可以期待更多有趣和挑战性的任务被智能体征服。

Logo

北京人形旗下天工造物具身智能开源社区,聚焦具身天工与慧思开物两大平台

更多推荐