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

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

    • [0. 前言](#0. 前言)
    • [1. Transformers](#1. Transformers)
    • [2. ChatGPT](#2. ChatGPT)
      • [2.1 设置](#2.1 设置)
      • [2.2 交互模式](#2.2 交互模式)
      • [2.3 ChatGPT API](#2.3 ChatGPT API)
    • 相关链接

0. 前言

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

1. Transformers

接下来我们将尝试使用预训练语言模型,这已成为现代自然语言处理领域的事实标准。得益于Hugging Face Hub等公共模型库,我们无需承担从零训练模型的高昂成本,只需将预训练模型接入现有架构,并对网络的一小部分进行微调以适应我们的数据集。

现有模型种类繁多------尺寸规格、预训练数据集、训练技术等各不相同。但所有模型都采用统一 API 接口,因此可以简单直接地集成到代码中。

首先需要安装相关库。针对我们的任务,需手动安装 sentence-transformers 包。安装完成后,即可使用该库计算任意字符串句子的嵌入向量:

shell 复制代码
>>> 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 进行通信,通过以下命令安装:

shell 复制代码
$ 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 函数:

python 复制代码
        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 函数的开始部分,我们重置环境并生成初始提示词:

python 复制代码
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) 随后我们执行循环直至游戏通关或达到步数限制:

python 复制代码
    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):

shell 复制代码
$ python3 chatgpt_interactive.py 1

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

2.3 ChatGPT API

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

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

python 复制代码
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 的信息:

python 复制代码
    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) 随后进入游戏循环,其逻辑与交互版本非常相似,只是无需控制台交互:

python 复制代码
    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 步内未能解决,则退出循环:

python 复制代码
    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互动小说游戏

相关推荐
JavaPub-rodert1 小时前
Codex 从 0 开始:安装 ChatGPT,并接入 DeepSeek
人工智能·gpt·chatgpt
Lee_jerome10 小时前
从 PyTorch 权重到 RK3588 板端推理:ResNet18 二分类模型完整部署教程
pytorch·边缘计算·rk3588·模型部署·onnx·resnet18·int8量化
菜冻鱼15 小时前
Python-pytorch-高级技巧
开发语言·人工智能·pytorch·python·深度学习·神经网络·聚类
菜冻鱼15 小时前
Python-pytorch-模型保存与加载
开发语言·人工智能·pytorch·python·深度学习·机器学习
LaughingZhu16 小时前
Product Hunt 每日热榜 | 2026-08-19
经验分享·深度学习·神经网络·百度·产品运营
小白学大数据17 小时前
CuPy vs Numba vs PyTorch:GPU 加速方案怎么选
人工智能·pytorch·爬虫·python
webor200617 小时前
<七>从3秒记忆到过目不忘——语言模型的三代进化
人工智能·语言模型·自然语言处理
聪明蛋子哟18 小时前
RL训练Agentic模型实现并行多轮上下文检索:匹配前沿模型且快10倍
人工智能·深度学习·机器学习
bulingg18 小时前
bert输入长度有限,如何处理超长文本?
人工智能·深度学习·bert