尧图网络 高端网站定制 · 原创设计
免费咨询热线
400-888-6620
免费获取方案
Transformers 中的 Decision Transformer:将离线强化学习重构为条件序列建模
Transformers 中的 Decision Transformer将离线强化学习重构为条件序列建模【免费下载链接】transformers Transformers: the model-definition framework for state-of-the-art machine learning models in text, vision, audio, and multimodal models, for both inference and training.项目地址: https://gitcode.com/GitHub_Trending/tra/transformersDecision Transformer 是 Transformers 仓库中一个把离线强化学习Offline RL抽象为条件序列建模问题的模型实现。它以 GPT-2 为骨干网络通过因果掩码 Transformer 直接输出最优动作而无需拟合价值函数或计算策略梯度。阅读本文后你将掌握 Decision Transformer 的核心思想、DecisionTransformerConfig全部配置参数的含义、模型前向计算流程以及如何基于 models/decision_transformer 模块加载预训练权重完成自回归动作预测。核心思想把强化学习当作序列建模问题原论文Decision Transformer: Reinforcement Learning via Sequence ModelingLili Chen、Kevin Lu、Aravind Rajeswaran 等人2021提出的框架将强化学习抽象为一个序列建模问题从而可以直接借用 Transformer 架构的简洁性与可扩展性以及 GPT-x、BERT 等语言建模领域的相关进展。与以往拟合价值函数或计算策略梯度的 RL 方法不同Decision Transformer 借助因果掩码 Transformer直接输出最优动作将自回归模型以期望回报return-to-go、过去的状态与动作为条件模型即可生成能够达成该期望回报的未来动作序列。尽管实现简单该模型在 Atari、OpenAI Gym 和 Key-to-Door 等任务上取得了与最先进的无模型离线 RL 基线相当甚至更优的表现。需要特别注意的是当前仓库中的这一版本模型只适用于状态为向量的任务如 OpenAI Gym 的连续控制环境不支持图像像素等非向量状态的场景。该模型由 edbeeching 贡献到 Hugging Face Transformers官方日语文档位于 docs/source/ja/model_doc/decision_transformer.md。整个模块由三个可独立理解的部分组成组件职责DecisionTransformerConfig模型全部超参数与默认值定义 RL 环境维度与 GPT 骨干规模DecisionTransformerGPT2Model去掉位置嵌入的 GPT-2 骨干含注意力、MLP、LayerNorm 堆叠DecisionTransformerModel完整模型模态嵌入层 GPT-2 骨干 状态/动作/回报预测头DecisionTransformerConfig完整参数详解DecisionTransformerConfig继承自PreTrainedConfig见 configuration_decision_transformer.pymodel_type为decision_transformer并声明了推理时忽略past_key_values键。它定义了两组参数一组描述 RL 环境状态、动作、episode 长度另一组描述 GPT-2 骨干网络结构。全部参数及其默认值如下RL 环境相关参数参数默认值说明state_dim17RL 环境的状态向量维度act_dim4输出动作空间的维度max_ep_len4096环境中一个 episode 的最大长度决定时间步嵌入表的大小action_tanhTrue动作预测输出后是否施加 tanh 激活用于把动作约束到合法范围GPT 骨干相关参数参数默认值说明hidden_size128隐藏层维度也是各模态嵌入的目标维度vocab_size1词表大小本模型实际不使用 token 输入仅保持与 GPT-2 兼容n_positions1024最大位置嵌入长度映射为max_position_embeddingsn_layer3Transformer 层数映射为num_hidden_layersn_head1注意力头数映射为num_attention_headsn_innerNoneMLP 中间维度None时取4 * hidden_sizeactivation_functionreluMLP 激活函数通过ACT2FN映射resid_pdrop/embd_pdrop/attn_pdrop0.1残差、嵌入、注意力三个位置的 dropout 概率layer_norm_epsilon1e-5LayerNorm 的 epsiloninitializer_range0.02权重初始化标准差scale_attn_weightsTrue是否将注意力权重除以sqrt(hidden_size)缩放scale_attn_by_inverse_layer_idxFalse是否额外按1 / (layer_idx 1)缩放注意力reorder_and_upcast_attnFalse混合精度训练时是否在计算点积前缩放 K 并将 softmax 上转为 float32add_cross_attentionFalse是否加入交叉注意力层默认关闭use_cacheTrue是否启用 KV 缓存bos_token_id/eos_token_id50256起始/结束 token idGPT-2 兼容保留字段配置中还定义了attribute_map将通用命名max_position_embeddings、num_attention_heads、num_hidden_layers映射到 GPT-2 风格命名n_positions、n_head、n_layer这样从预训练权重加载时可以正确对齐。从配置创建随机初始化模型的官方写法 from transformers import DecisionTransformerConfig, DecisionTransformerModel # 初始化一个默认配置 configuration DecisionTransformerConfig() # 由配置创建随机权重模型 model DecisionTransformerModel(configuration) # 读取模型配置 configuration model.config模型架构DecisionTransformerGPT2ModelDecisionTransformerGPT2Model源码见 modeling_decision_transformer.py是完整的 GPT-2 骨干实现包含wtetoken 嵌入层保留 GPT-2 结构但 Decision Transformer 实际不喂 token idwpe位置嵌入层hnum_hidden_layers个DecisionTransformerGPT2Block堆叠ln_f最终 LayerNorm。每个DecisionTransformerGPT2Block由ln_1 → attention → 残差 → ln_2 → MLP → 残差组成若开启add_cross_attention还会插入交叉注意力。其注意力实现DecisionTransformerGPT2Attention直接复用 GPT-2 的代码结构Conv1D投影、多头拆分、缩放点积注意力并支持三种注意力量化开关scale_attn_weights、scale_attn_by_inverse_layer_idx、reorder_and_upcast_attn。权重初始化遵循 GPT-2 论文方案——残差路径上的c_proj层权重按initializer_range / sqrt(2 * n_layer)缩放以抵消深度残差网络中的梯度累积效应。与标准 GPT-2 唯一的关键差异是DecisionTransformerGPT2Model虽然保留wpe权重但在DecisionTransformerModel中调用时传入的位置 id 恒为 0位置信息改由时间步嵌入timestep embedding提供。此外模型通过create_causal_mask构造因果注意力掩码保证每个 token 只能看到其之前的位置。DecisionTransformerModel完整模型与前向流程DecisionTransformerModel在 GPT-2 骨干之上叠加了 RL 特有的模态嵌入层与预测头self.embed_timestep nn.Embedding(config.max_ep_len, config.hidden_size) # 时间步嵌入 self.embed_return torch.nn.Linear(1, config.hidden_size) # 回报嵌入 self.embed_state torch.nn.Linear(config.state_dim, config.hidden_size) # 状态嵌入 self.embed_action torch.nn.Linear(config.act_dim, config.hidden_size) # 动作嵌入 self.embed_ln nn.LayerNorm(config.hidden_size) self.predict_state torch.nn.Linear(config.hidden_size, config.state_dim) # 状态预测头 self.predict_action nn.Sequential( # 动作预测头 *([nn.Linear(config.hidden_size, config.act_dim)] ([nn.Tanh()] if config.action_tanh else [])) ) self.predict_return torch.nn.Linear(config.hidden_size, 1) # 回报预测头输入格式forward接收五个必需张量另加可选attention_mask参数形状含义states(batch_size, episode_length, state_dim)轨迹中每一步的状态actions(batch_size, episode_length, act_dim)专家策略在当前状态采取的动作自回归预测时被掩码rewards(batch_size, episode_length, 1)每一步的奖励returns_to_go(batch_size, episode_length, 1)每一步的待实现回报期望回报减去已获奖励的累积timesteps(batch_size, episode_length)轨迹中每一步的时间步编号attention_mask(batch_size, episode_length)1 表示可被注意力关注0 表示忽略前向计算流程源码级拆解模态嵌入状态、动作、回报分别经各自的线性层投影到hidden_size时间步经embed_timestep得到位置式嵌入并加到上述三类嵌入上代码注释明确指出 time embeddings are treated similar to positional embeddings。序列堆叠将三个模态的嵌入堆叠为(R_1, s_1, a_1, R_2, s_2, a_2, ...)的顺序序列长度变为3 * episode_length随后过embed_ln做 LayerNorm。由于 GPT 骨干在自回归意义下由状态预测动作最为自然这种交错排列是模型设计的核心。掩码堆叠attention_mask同样按三个模态复制堆叠为3 * episode_length长度。骨干前向以inputs_embedsstacked_inputs而非 token id、全零position_ids送入DecisionTransformerGPT2Model得到last_hidden_state。重排与预测输出重排为(batch, episode_length, 3, hidden_size)再按模态维度 permute。其中action_preds predict_action(x[:, 1])—— 由状态 token预测下一步动作state_preds predict_state(x[:, 2])—— 由状态动作 token 预测下一状态return_preds predict_return(x[:, 2])—— 由状态动作 token 预测下一回报。输出格式return_dictFalse时返回三元组(state_preds, action_preds, return_preds)否则返回DecisionTransformerOutput数据类包含state_preds形状(batch_size, sequence_length, state_dim)、action_preds形状(batch_size, sequence_length, act_dim)、return_preds形状(batch_size, sequence_length, 1)以及last_hidden_state、hidden_states、attentions。实战加载预训练权重进行自回归评估源码 docstring 中给出了完整的评估循环示例基于 OpenAI Gym 的 Hopper-v3 环境。核心步骤是先以目标回报初始化returns_to_go每步用模型预测动作执行环境步进后更新状态与待实现回报 from transformers import DecisionTransformerModel import torch model DecisionTransformerModel.from_pretrained(edbeeching/decision-transformer-gym-hopper-medium) model model.to(device) model.eval() env gym.make(Hopper-v3) state_dim env.observation_space.shape[0] act_dim env.action_space.shape[0] state env.reset() states torch.from_numpy(state).reshape(1, 1, state_dim).to(devicedevice, dtypetorch.float32) actions torch.zeros((1, 1, act_dim), devicedevice, dtypetorch.float32) rewards torch.zeros(1, 1, devicedevice, dtypetorch.float32) target_return torch.tensor(TARGET_RETURN, dtypetorch.float32).reshape(1, 1) timesteps torch.tensor(0, devicedevice, dtypetorch.long).reshape(1, 1) attention_mask torch.zeros(1, 1, devicedevice, dtypetorch.float32) # 前向传播 with torch.no_grad(): ... state_preds, action_preds, return_preds model( ... statesstates, ... actionsactions, ... rewardsrewards, ... returns_to_gotarget_return, ... timestepstimesteps, ... attention_maskattention_mask, ... return_dictFalse, ... )测试与集成验证输出形状与自回归行为仓库的测试套件 tests/models/decision_transformer/test_modeling_decision_transformer.py 从三个层面验证了上述行为形状校验create_and_check_model断言state_preds、action_preds、return_preds分别与输入状态、动作、回报形状一致且last_hidden_state形状为(batch_size, seq_length * 3, hidden_size)——正是三个模态堆叠的结果测试注释明确写着 seq length *3 as there are 3 modalities: states, returns and actions。前向签名校验test_forward_signature通过inspect.signature断言forward的前六个位置参数依次为states、actions、rewards、returns_to_go、timesteps、attention_mask。自回归集成测试test_autoregressive_prediction加载edbeeching/decision-transformer-gym-hopper-expert预训练权重模拟两个时间步的闭环评估——每步用模型预测动作更新状态、returns_to_gopred_return returns_to_go - reward与timesteps并将预测动作与期望输出做数值比对rtol/atol 均为 1e-4。该测试完整展示了预测-步进-拼接的在线推理闭环。总结在 Transformers 仓库中Decision Transformer 的实现把条件序列建模这一思想落到了可复用的工程代码上DecisionTransformerConfig让环境维度与网络规模完全参数化DecisionTransformerGPT2Model提供经过时间步嵌入改造的 GPT-2 骨干而DecisionTransformerModel通过三模态嵌入堆叠与三个预测头实现了给定期望回报与历史轨迹自回归生成最优动作的完整能力。理解这一实现是进一步阅读其代码、改造适配新环境或复现离线 RL 实验的起点。【免费下载链接】transformers Transformers: the model-definition framework for state-of-the-art machine learning models in text, vision, audio, and multimodal models, for both inference and training.项目地址: https://gitcode.com/GitHub_Trending/tra/transformers创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
RELATED

相关推荐

STM32嵌入式开发中Claude Code的工程化应用实践

STM32嵌入式开发中Claude Code的工程化应用实践

1. 项目概述:当嵌入式开发遇上AI编程助手,不是替代工程师,而是重构工作流“嵌入式软件AI编程”这个标题乍看像概念炒作,但如果你最近在Keil里调了三天串口DMA收发却始终丢帧、在CubeMX生成的HAL库初始化代码里反复注释/反注释时钟…

📅 2026/9/11 18:50:44
CMSIS-FreeRTOS源码审计:从适配层到内核调度的嵌入式RTOS实践

CMSIS-FreeRTOS源码审计:从适配层到内核调度的嵌入式RTOS实践

做嵌入式开发这些年,我越来越觉得有一类源码值得每个 MCU 工程师静下心来从头到尾读一遍:CMSIS-FreeRTOS。尤其是用 STM32CubeMX、Keil RTE 或 ARM 官方仓库生成过工程的朋友,你们每天打开的就是这个东西,但它内部到底怎么跑起来的…

📅 2026/9/11 18:45:43
运放同相与反相怎么选?从输入阻抗、偏置到噪声增益讲透

运放同相与反相怎么选?从输入阻抗、偏置到噪声增益讲透

前阵子画一块光电检测板的信号调理电路,原理图出来后,同事指着一个运放级问我:这一级为什么用反相放大?信号相位在这里需要翻转吗?我当时愣了一下,然后诚实地说:其实不是因为相位,是…

📅 2026/9/11 18:45:43
MORE NEWS

更多资讯

📰

Actual Budget CLI 完全指南:用 @actual-app/cli 在终端管理与查询预算数据

Actual Budget CLI 完全指南:用 actual-app/cli 在终端管理与查询预算数据 【免费下载链接】actual A local-first personal finance app 项目地址: https://gitcode.com/GitHub_Trending/ac/actual 本篇技术指南系统讲解 Actual Budget 官方命令行工具 actu…

📰

Android Soong 构建规则中的 Rust:从 rust_binary 到 rust_library 的完整指南(comprehensive-rust 实战)

Android Soong 构建规则中的 Rust:从 rust_binary 到 rust_library 的完整指南(comprehensive-rust 实战) 【免费下载链接】comprehensive-rust This is the Rust course used by the Android team at Google. It provides you the material …

📰

考研调剂公告怎么读?从长春工业大学样本看调剂全流程与避坑指南

每年二三月,考研初试成绩公布后的这段时间,很多学生的手机都会一刻不停刷着各类调剂信息。如果你恰好看到了“长春工业大学2025年硕士研究生招生调剂公告(一)”这样一条标题,第一反应可能是赶紧点进去找有没有自己的专…

📰

Java保险业务管理系统毕业设计:从技术选型到部署答辩全攻略

简介:基于Java的保险业务管理系统毕业设计完整资料包,面向高校计算机或软件工程专业学生,可用于毕业设计参考、课程项目实训或Java Web开发实践。压缩包约66.12MB,内容涵盖项目报告、答辩PPT、源代码、数据库脚本、界面截图及部署…

📰

双关节机械臂自适应模糊反演控制:建模、设计与MATLAB仿真

简介:双关节机械臂的自适应模糊反演控制Matlab仿真包,面向机器人控制与智能算法方向的本科、硕士阶段教研学习。资源围绕双关节机械臂的轨迹跟踪问题,完整实现了基于模糊系统的自适应反演控制方案,通过模糊逻辑逼近系统未知非线性…

📰

从零实现 C++ AI 大模型接入 SDK(一):项目介绍与整体架构

目录 前言 一、我们要做的不是“大模型”,而是“大模型应用” 二、项目最终能做到什么 三、项目整体架构 3.1 Web UI:用户真正操作的界面 3.2 ChatServer:把 SDK 变成 HTTP 服务 四、ChatSDK:整个项目真正的核心 五、SDK …

TODAY

今日更新

THIS WEEK

本周精选

THIS MONTH

本月热门

读完文章,想聊聊您的网站?

告诉我们您的行业与需求,资深顾问一对一梳理方案与报价,全程免费。

📞 💬