PyTorch强化学习实战——用预训练语言模型和ChatGPT玩转文本游戏

0. 前言

文本互动小说 (interactive fiction) 是强化学习研究中一个独特而富有挑战性的领域。与图形丰富的街机游戏不同,这类游戏通过纯文本描述呈现状态,要求智能体理解自然语言、进行长期规划并在复杂的语义空间中决策。我们已经以微软 TextWorld 为实验平台,学习了如何使用自然语言处理 (Natural Language Processing, NLP) 工具处理复杂的文本数据,并在交互式小说游戏环境中进行实验,在本节中,我们将借助 Hugging Face 预训练 TransformerChatGPT API 展现大语言模型在文本游戏中的强大能力。通过从手工特征到预训练模型的演进,我们将见证深度 NLP 技术如何赋能强化学习智能体。

1. Transformers

接下来我们将尝试使用预训练语言模型,这已成为现代自然语言处理领域的事实标准。得益于 Hugging Face Hub 等公共模型库,我们无需承担从零训练模型的高昂成本,只需将预训练模型接入现有架构,并对网络的一小部分进行微调以适应我们的数据集。
现有模型种类繁多——尺寸规格、预训练数据集、训练技术等各不相同。但所有模型都采用统一 API 接口,因此可以简单直接地集成到代码中。
首先需要安装相关库。针对我们的任务,需手动安装 sentence-transformers 包。安装完成后,即可使用该库计算任意字符串句子的嵌入向量:

>>> from sentence_transformers import SentenceTransformer 
>>> tr = SentenceTransformer("sentence-transformers/all-MiniLM-L6-v2") 
>>> tr.get_sentence_embedding_dimension() 
384 
>>> r = tr.encode("You’re standing in an ordinary boring room") 
>>> type(r) 
<class ’numpy.ndarray’> 
>>> r.shape 
(384,) 
>>> r2 = tr.encode(["sentence 1", "sentence 2"], convert_to_tensor=True) 
>>> type(r2) 
<class ’torch.Tensor’> 
>>> r2.shape 
torch.Size([2, 384])

在本节中,我们使用了 all-MiniLM-L6-v2 模型,它相对较小——有 2200 万个参数,训练数据为 12 亿个词元。
在本节中,我们将使用高级接口,直接输入字符串语句,由库和模型完成所有转换工作。但该方案在需要时仍能提供充分的灵活性。
preproc.TransformerPreprocessor 类实现了与原有 Preprocessor 类(使用长短期记忆 (Long Short-Term Memory, LSTM) 进行嵌入)相同的接口。
要使用 Transformers 训练智能体,需要运行 train_tr.py 模块。在训练过程中,Transformer 模型的处理速度较慢,这是因为 Transformer 模型比 LSTM 模型复杂得多,但在 20 个和 200 个游戏上的训练动态表现更优。对比 Transformer基准模型的训练奖励和回合步数,基准版本需要 1000 回合才能达到 15 步,而 Transformer 模型只需要 400 回合。但在 20 个游戏的验证中,奖励低于基准版本(最高分为 2)。
200 个游戏上的训练也呈现相同情况——智能体学习效率更高(以游戏数量衡量),但验证效果不佳。这可能是因为 Transformer 模型的容量要大得多——其生成的嵌入向量维度几乎是基线模型的 20 倍( 384 维对比 20 维),导致智能体更容易直接记忆正确的步骤序列,而非尝试寻找高层次通用观测特征到动作的映射关系。

2. ChatGPT

为了完成对 TextWorld 的讨论,我们继续尝试另一种方法——使用大语言模型 (Large Language Model, LLM)。自从 2022 年底公开发布后,ChatGPT 迅速流行起来,彻底改变了聊天机器人和文本助手领域。接下来,我们尝试将这项技术应用于解决 TextWorld 游戏问题。

2.1 设置

首先需要注册 OpenAI 账号。我们将从基于网页的交互式聊天开始实验,但后续示例将使用 ChatGPT API,这需要在 https://platform.openai.com 生成 API 密钥。创建密钥后,需将其设置到所用 shell 环境的 OPENAI_API_KEY 变量中。
同时我们将使用 langchain 库与 ChatGPT 进行通信,通过以下命令安装:

$ pip install langchain langchain-openai

2.2 交互模式

在第一个示例中,我们将使用基于网页的 ChatGPT 界面,要求其根据房间描述和游戏目标生成游戏指令。代码位于 chatgpt_interactive.py,主要实现以下功能:

  1. 启动命令行指定游戏 IDTextWorld 环境
  2. ChatGPT 创建包含操作说明、游戏目标和房间描述的提示词
  3. 将提示词输出至控制台
  4. 从控制台读取待执行的指令
  5. 在环境中执行该指令
  6. 重复步骤 2-5 直至达到步数限制或游戏通关。

所以,我们的任务是将生成的提示词复制并粘贴到 https://chat.openai.com 网页界面中,ChatGPT 将生成需要输入控制台的指令。

(1) 完整代码非常简洁,仅包含一个执行游戏循环的 play_game 函数:

        env_id = register_game(
            gamefile=f"games/{args.game}{index}.ulx",
            request_infos=EnvInfos(description=True, objective=True),
        )
        env = gym.make(env_id)

在创建环境时,我们仅要求获取两个额外信息:房间描述和游戏目标。原则上这些信息都包含在自由文本观察值中,因此可通过解析文本获取。但为方便起见,我们直接要求 TextWorld 显式提供这些信息。

(2)play_game 函数的开始部分,我们重置环境并生成初始提示词:

def play_game(env, max_steps: int = 20) -> bool:
    commands = []

    obs, info = env.reset()

    print(textwrap.dedent("""\
    You're playing the interactive fiction game.
    Here is the game objective: %s
    
    Here is the room description: %s
    
    What command do you want to execute next? Reply with 
    just a command in lowercase and nothing else. 
    """)  % (info['objective'], info['description']))

    print("=== Send this to chat.openai.com and type the reply...")

为了避免 ChatGPT 输出冗长的内容,我们可以要求其仅回复可输入游戏的指令。

(2) 随后我们执行循环直至游戏通关或达到步数限制:

    while len(commands) < max_steps:
        cmd = input(">>> ")
        commands.append(cmd)
        obs, r, is_done, info = env.step(cmd)
        if is_done:
            print(f"You won in {len(commands)} steps! "
                  f"Don't forget to congratulate ChatGPT!")
            return True

        print(textwrap.dedent("""\
        Last command result: %s
        Room description: %s
        
        What's the next command?
        """) % (obs, info['description']))
        print("=== Send this to chat.openai.com and type the reply...")

    print(f"Wasn't able to solve after {max_steps} steps, commands: {commands}")
    return False

后续提示词更为简洁——我们只需提供获得的观察结果(即指令执行结果)和新的房间描述。由于网页界面会保持对话上下文,无需重复传递游戏目标,聊天机器人能记住之前的指令。

(3) 查看一个游戏测试(使用种子 1):

$ python3 chatgpt_interactive.py 1

运行过程

可以看到,大语言模型能够完美解决这个任务。同时,整体任务难度实际上更高——我们要求其生成指令,而非像本节前文那样从"可用指令"列表中做出选择。

2.3 ChatGPT API

由于复制粘贴操作繁琐乏味,接下来,我们使用 ChatGPT API 实现智能体自动化。我们将采用 langchain 库,该库提供了足够的灵活性和控制力来发挥大语言模型的功能。

(1) 完整的代码位于文件 chatgpt_auto.py 中。接下来,我们介绍核心函数 play_game()

from langchain_openai import ChatOpenAI
from langchain_core.output_parsers import StrOutputParser
from langchain_core.prompts import ChatPromptTemplate, MessagesPlaceholder

def play_game(env, max_steps: int = 20) -> bool:
    prompt_init = ChatPromptTemplate.from_messages([
        ("system", "You're playing the interactive fiction game. "
                   "Reply with just a command in lowercase and nothing else"),
        ("system", "Game objective: {objective}"),
        ("user", "Room description: {description}"),
        ("user", "What command you want to execute next?"),
    ])
    llm = ChatOpenAI()
    output_parser = StrOutputParser()

初始提示词与之前相同——向聊天机器人说明游戏类型,并要求其仅回复可输入游戏的指令。

(2) 接着重置环境并生成第一条消息,传递来自 TextWorld 的信息:

    commands = []

    obs, info = env.reset()
    init_msg = prompt_init.invoke({
        "objective": info['objective'],
        "description": info['description'],
    })

    context = init_msg.to_messages()
    ai_msg = llm.invoke(init_msg)
    context.append(ai_msg)
    cmd = output_parser.invoke(ai_msg)

变量 context 至关重要,它包含当前对话中的所有消息记录(包括用户和聊天机器人的消息)。我们将这些消息传递给聊天机器人以保持游戏进程的连续性。这是必要的,因为游戏目标仅显示一次且不会重复。若没有历史记录,智能体将缺乏足够信息来执行所需操作序列。另一方面,传递大量文本可能导致成本上升( ChatGPT API 按处理 token 数量计费)。我们的游戏流程较短( 5-7 步即可完成任务),因此不是主要问题,但对于更复杂的游戏,可能需要优化历史记录。

(3) 随后进入游戏循环,其逻辑与交互版本非常相似,只是无需控制台交互:

    prompt_next = ChatPromptTemplate.from_messages([
        MessagesPlaceholder(variable_name="chat_history"),
        ("user", "Last command result: {result}"),
        ("user", "Room description: {description}"),
        ("user", "What command you want to execute next?"),
    ])

    for _ in range(max_steps):
        commands.append(cmd)
        print(">>>", cmd)
        obs, r, is_done, info = env.step(cmd)
        if is_done:
            print(f"I won in {len(commands)} steps!")
            return True

        user_msgs = prompt_next.invoke({
            "chat_history": context,
            "result": obs.strip(),
            "description": info['description'],
        })
        context = user_msgs.to_messages()
        ai_msg = llm.invoke(user_msgs)
        context.append(ai_msg)
        cmd = output_parser.invoke(ai_msg)

在后续提示中,我们传递对话历史、上条指令的执行结果、当前房间描述,并请求下一条指令。

(4) 同时我们设置了步数限制以防止智能体陷入循环(这种情况时有发生)。若游戏在 20 步内未能解决,则退出循环:

    print(f"Wasn't able to solve after {max_steps} steps, commands: {commands}")
    return False

20TextWorld 游戏(种子 1-20)上对上述代码进行了测试,成功解决了其中 9 个游戏。多数失败情况是由于智能体陷入循环——生成未被 TextWorld 正确解析的错误指令(例如使用 “take the key” 而非 “take the key from the box”),或在导航过程中卡住。
有两个游戏中,ChatGPT 因生成 “exit” 指令而失败,该指令会立即终止 TextWorld 进程。若能检测该指令或在提示中禁止其生成,很可能提高通关率。但即便如此,智能体未经任何预先训练就能解决 9 个游戏已是相当优异的结果。

相关链接

PyTorch强化学习实战(1)——强化学习(Reinforcement Learning,RL)详解
PyTorch强化学习实战(2)——强化学习环境库Gymnasium
PyTorch强化学习实战(3)——Gymnasium API扩展功能
PyTorch强化学习实战(4)——PyTorch基础
PyTorch强化学习实战(5)——PyTorch Ignite 事件驱动机制与实践
PyTorch强化学习实战(6)——交叉熵方法详解与实现
PyTorch强化学习实战(7)——表格学习与贝尔曼方程
PyTorch强化学习实战(8)——Q学习详解与实现
PyTorch强化学习实战(9)——深度Q学习
PyTorch强化学习实战(10)——强化学习高级组件
PyTorch强化学习实战(11)——N步DQN(N-step DQN)
PyTorch强化学习实战(12)——Double DQN(DDQN)
PyTorch强化学习实战(13)——噪声网络(NoisyNet-DQN)
PyTorch强化学习实战(14)——优先经验回放机制
PyTorch强化学习实战(15)——Dueling DQN
PyTorch强化学习实战(16)——Categorical DQN
PyTorch强化学习实战(17)——强化学习训练加速
PyTorch强化学习实战(18)——基于DQN处理股票交易问题
PyTorch强化学习实战(19)——策略梯度法
PyTorch强化学习实战(20)——优势演员-评论家(Advantage Actor-Critic, A2C)
PyTorch强化学习实战(21)——异步优势演员-评论家(Asynchronous Advantage Actor-Critic, A3C)
PyTorch强化学习实战(22)——将强化学习应用于TextWorld互动小说游戏

Logo

AtomGit AI 社区提供模型库、数据集、Agent、Token等资源

更多推荐