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

PyTorch实现DQN遇矩阵形状不匹配RuntimeError的解决问询

DQN中PyTorch张量维度不匹配问题的解决与优化

问题重现

用PyTorch实现DQN算法时,环境观测经预处理后得到shape为torch.Size([1, 2, 9, 7])的张量,调用网络的act函数时触发错误:

RuntimeError: mat1 and mat2 shapes cannot be multiplied (18x7 and 126x64)

原因剖析

错误本质是线性层输入维度不匹配:

  • 你定义的第一层线性层输入维度是126(对应2*9*7,即通道数×高度×宽度的总特征数),但实际传入线性层的张量形状是(18,7)——这是因为没有正确保留batch维度,错误地将[1,2,9,7]的前三维合并成了18,导致最后一维还剩7,和线性层的126维输入要求冲突。

现有解决方案解析

你用的两行代码刚好命中了问题核心:

  • obs = obs.unsqueeze(0):针对单条观测(无batch维度,shape为[2,9,7])补充batch维度,统一成[1,2,9,7]的批量格式,确保后续处理逻辑兼容单条/批量输入。
  • obs = obs.view(obs.shape[0], -1):固定batch维度(第一维),将后面的所有维度(2,9,7)展平为一维,得到[1, 126]的张量,完美匹配线性层的输入维度要求。

优化方案

1. 模型内置展平逻辑(推荐)

把展平操作整合到神经网络的前向传播中,避免外部手动处理,用nn.Flatten()层自动处理batch维度:

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

class DQN(nn.Module):
    def __init__(self, input_shape, num_actions):
        super().__init__()
        # 自动展平:保留batch维度,展平后续所有维度
        self.flatten = nn.Flatten(start_dim=1)
        # 计算输入总特征数:input_shape是(2,9,7)
        input_features = torch.prod(torch.tensor(input_shape))
        self.fc1 = nn.Linear(input_features, 64)
        self.fc2 = nn.Linear(64, num_actions)
    
    def forward(self, x):
        x = self.flatten(x)
        x = F.relu(self.fc1(x))
        return self.fc2(x)

这样不管输入是单条观测([2,9,7])还是批量观测([N,2,9,7]),模型都能自动处理维度,无需外部调用view或unsqueeze。

2. 用flatten替代view增强可读性

如果不想修改模型,用torch.flatten(start_dim=1)替代view,语义更清晰:

# 替代obs.view(obs.shape[0], -1)
obs = obs.flatten(start_dim=1)

start_dim=1明确表示从第1维开始展平(跳过batch维度),效果和view一致,但代码意图更直观。

3. 预处理阶段统一输出格式

在观测预处理函数中,直接确保输出带batch维度:

def preprocess_observation(obs):
    # 假设原obs是numpy数组或无batch维度的张量
    obs_tensor = torch.tensor(obs, dtype=torch.float32)
    # 确保有batch维度
    if len(obs_tensor.shape) == 3:
        obs_tensor = obs_tensor.unsqueeze(0)
    return obs_tensor

这样后续调用模型时无需再手动补充batch维度,减少重复代码。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.25 08:23:11