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 预训练 Transformer 和 ChatGPT 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,主要实现以下功能:
- 启动命令行指定游戏
ID的TextWorld环境 - 为
ChatGPT创建包含操作说明、游戏目标和房间描述的提示词 - 将提示词输出至控制台
- 从控制台读取待执行的指令
- 在环境中执行该指令
- 重复步骤
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
在 20 个 TextWorld 游戏(种子 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强化学习实战(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强化学习实战(20)------优势演员-评论家(Advantage Actor-Critic, A2C)
PyTorch强化学习实战(21)------异步优势演员-评论家(Asynchronous Advantage Actor-Critic, A3C)