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

DQN智能体优先级经验回放缓冲区张量维度不匹配报错求助

优先级经验回放缓冲区适配多维状态的问题解决

问题描述

环境状态维度为(40, 40, 1),向优先级经验回放缓冲区添加transition时触发错误:

RuntimeError: expand(torch.DoubleTensor{[40, 40, 1]}, size=[3]): the number of sizes provided (1) must be greater or equal to the number of dimensions in the tensor (3)

问题根源

原缓冲区代码的__init__函数默认state_size=3,并创建了self.state = torch.empty(buffer_size, state_size, dtype=torch.float)——这是为一维状态(如3维向量)设计的结构。但你的状态是3维张量(40,40,1),直接赋值时维度不匹配,触发了expand错误。

修复方案

修改后的缓冲区代码

class PrioritizedReplayBuffer:
    def __init__(self, state_shape=(40,40,1), action_size=1, buffer_size=10000, eps=1e-2, alpha=0.1, beta=0.1):
        self.tree = SumTree(size=buffer_size)

        # PER参数
        self.eps = eps 
        self.alpha = alpha
        self.beta = beta
        self.max_priority = eps

        # 存储transition的张量,适配多维状态
        self.state = torch.empty((buffer_size, *state_shape), dtype=torch.float)
        self.action = torch.empty(buffer_size, action_size, dtype=torch.float)
        self.reward = torch.empty(buffer_size, dtype=torch.float)
        self.next_state = torch.empty((buffer_size, *state_shape), dtype=torch.float)
        self.done = torch.empty(buffer_size, dtype=torch.int)

        self.count = 0
        self.real_size = 0
        self.size = buffer_size

    def add(self, transition):
        state, action, reward, next_state, done = transition

        self.tree.add(self.max_priority, self.count)

        # 统一张量类型,避免隐式转换错误
        self.state[self.count] = torch.as_tensor(state, dtype=torch.float)
        self.action[self.count] = torch.as_tensor(action, dtype=torch.float)
        self.reward[self.count] = torch.as_tensor(reward, dtype=torch.float)
        self.next_state[self.count] = torch.as_tensor(next_state, dtype=torch.float)
        self.done[self.count] = torch.as_tensor(done, dtype=torch.int)

        self.count = (self.count + 1) % self.size
        self.real_size = min(self.size, self.real_size + 1)

关键修改点

  • 将state_size参数替换为state_shape,接收元组形式的多维状态维度(如(40,40,1))
  • 创建状态张量时,用(buffer_size, *state_shape)展开维度,确保缓冲区每个位置能容纳3维状态张量
  • 在add方法中显式指定dtype=torch.float,避免输入张量类型(如DoubleTensor)不匹配导致的错误

内容的提问来源于stack exchange,提问作者E.T.Tuna

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.26 23:48:22