如何用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
相关产品推荐
相关产品推荐

