You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何用Decision Transformer替换RL环境中的RNN并解决报错

用Decision Transformer替换RNN智能体的解决方案

问题背景

现有一个返回动作和隐藏状态的RNN强化学习智能体,需要用Decision Transformer(DT)替换,保持接口和输出一致,但DT无隐藏状态,尝试虚拟隐藏状态时触发错误。

报错原因分析

当前报错TypeError: 'int' object is not subscriptable,是因为原RNN的input_shape是整数(状态特征维度),但MultiAgentDecisionTransformer中错误地将其当作列表/张量取self.state_shape[0],直接访问整数的索引导致报错。此外还有其他潜在问题:

  • DT的输入适配问题:原RNN是单步输入,DT默认处理序列,需要调整适配单步推理
  • 虚拟隐藏状态的形状不匹配原RNN输出
  • 动作嵌入层的输入类型错误:原代码中传入inputs给action_embed,但Embedding层需要整数类型的动作索引,而非连续状态特征

修正后的完整代码

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

class DTagent(nn.Module):
    def __init__(self, input_shape, args):
        super(DTagent, self).__init__()
        self.args = args
        self.n_actions = args.n_actions
        self.n_agents = args.n_agents
        self.input_shape = input_shape  # input_shape是整数,对应状态特征维度
        self.output_type = "p"

        self.dt = MultiAgentDecisionTransformer(input_shape, args)

    def forward(self, inputs, hidden_state):
        # DT不需要维护隐藏状态,所以忽略传入的hidden_state
        # 适配原RNN的单步输入:添加序列维度(DT期望序列输入)
        inputs_seq = inputs.unsqueeze(1)  # 形状从[batch*agents, dim]变为[batch*agents, 1, dim]
        # 由于是单步推理,传入全0的动作序列作为占位(DT需要历史动作输入)
        dummy_actions = torch.zeros(inputs_seq.shape[0], 1, dtype=torch.long, device=inputs.device)
        
        q = self.dt(inputs_seq, dummy_actions)
        # 去掉序列维度,匹配原RNN的输出形状
        q = q.squeeze(1)

        # 生成和原RNN隐藏状态形状一致的虚拟隐藏状态
        # 原RNN的hidden_state形状是[batch*agents, rnn_hidden_dim]
        dummy_hidden = torch.zeros(
            inputs.shape[0], self.args.rnn_hidden_dim, 
            device=inputs.device, dtype=inputs.dtype
        )

        return q, dummy_hidden

    def init_hidden(self):
        # 返回和原RNN同形状的初始隐藏状态,保持接口一致
        return torch.zeros(1, self.args.rnn_hidden_dim)

class MultiAgentDecisionTransformer(nn.Module):
    def __init__(self, input_shape, args):
        super(MultiAgentDecisionTransformer, self).__init__()
        self.state_dim = input_shape  # input_shape是整数,直接作为状态维度
        self.n_actions = args.n_actions
        self.n_agents = args.n_agents
        self.embed_dim = args.rnn_hidden_dim
        self.num_heads = 5
        self.num_layers = 2

        # 状态嵌入:从状态维度映射到Transformer嵌入维度
        self.state_embed = nn.Linear(self.state_dim, self.embed_dim)
        # 动作嵌入:针对离散动作,将动作索引映射到嵌入维度
        self.action_embed = nn.Embedding(self.n_actions, self.embed_dim)

        # 使用TransformerEncoder(单步推理不需要Decoder)
        self.transformer_encoder = nn.TransformerEncoder(
            encoder_layer=nn.TransformerEncoderLayer(
                d_model=self.embed_dim,
                nhead=self.num_heads,
                batch_first=True  # 设置batch_first=True,输入形状为[batch, seq_len, dim]
            ),
            num_layers=self.num_layers
        )

        self.action_decoder = nn.Linear(self.embed_dim, self.n_actions)

    def forward(self, state_seq, action_seq):
        # state_seq形状:[batch, seq_len, state_dim]
        # action_seq形状:[batch, seq_len](整数动作索引)
        state_emb = self.state_embed(state_seq)  # [batch, seq_len, embed_dim]
        action_emb = self.action_embed(action_seq)  # [batch, seq_len, embed_dim]

        # 融合状态和动作嵌入
        combined_emb = state_emb + action_emb

        # Transformer编码器前向传播
        transformer_out = self.transformer_encoder(combined_emb)

        # 解码得到动作logits
        action_logits = self.action_decoder(transformer_out)

        return action_logits

关键修正点

  • 修复输入维度错误:将self.state_shape[0]改为self.state_dim(直接使用传入的整数input_shape)
  • 适配单步推理:在DTagent.forward中给输入添加序列维度,传入虚拟动作序列作为DT的历史动作输入
  • 匹配输出形状:去掉DT输出的序列维度,确保和原RNN的动作输出形状一致;生成和原RNN隐藏状态形状完全匹配的虚拟隐藏状态
  • 简化Transformer结构:使用TransformerEncoder替代完整Transformer,因为单步推理不需要Decoder,减少不必要的复杂度
  • 统一设备类型:虚拟张量使用和输入相同的设备,避免设备不匹配错误

内容的提问来源于stack exchange,提问作者ashah

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.07 05:33:15