【问题标题】:PyTorch model input shapePyTorch 模型输入形状
【发布时间】:2021-06-03 21:55:24
【问题描述】:

我加载了一个自定义 PyTorch 模型,我想找出它的输入形状。像这样的:

model.input_shape

是否有可能获得这些信息?


更新:print()summary() 不显示此模型的输入形状,因此它们不是我要查找的。​​p>

【问题讨论】:

  • 嗨,兄弟,pytorch 模型的输入形状是灵活的。唯一重要的是它的深度、RGB 或灰度。
  • 如果是卷积神经网络模型..
  • @yakhyo,所以输入可以是任何形状?
  • 是的,它可以是除深度以外的任何形状

标签: python deep-learning pytorch conv-neural-network


【解决方案1】:

PyTorch 灵活性

PyTorch 模型是非常灵活的对象,以至于它们不强制或通常不期望数据具有固定的输入形状。

如果您有某些层,可能会有限制,例如:

  • 一个展平后跟一个宽度为 N 的全连接层将强制您的原始输入 (M1 x M2 x ... Mn) 的尺寸具有等于 N 的乘积
  • N 个输入通道的 2d 卷积将强制数据为 3 维,第一个维度的大小为 N

但正如您所见,这些都不会强制数据的total 形状。

我们现在可能没有意识到这一点,但在更复杂的模型中,正确设置第一个线性层的大小有时会令人沮丧。我们听说过著名的实践者输入任意数字,然后依靠 PyTorch 的错误消息来回溯线性层的正确大小。跛脚,嗯?不,这都是合法的!

  • 使用 PyTorch 进行深度学习

调查

简单案例:第一层全连接

如果您的模型的第一层是全连接层,那么print(model) 中的第一层将详细说明单个样本的预期维度。

模棱两可的情况:CNN

但是,如果它是卷积层,由于这些是动态的,并且会在输入允许的范围内尽可能长/宽,因此没有简单的方法可以从模型本身中检索此信息。1 这灵活性意味着对于许多架构多种兼容的输入尺寸2都将被网络接受。

这是 PyTorch 的 Dynamic computational graph 的一个功能。

人工检查

你需要做的是调查网络架构,一旦你找到一个可解释的层(如果存在,例如完全连接),它的维度“向后工作”,确定之前的层(例如池和卷积)已对其进行压缩/修改。

示例

例如在 Deep Learning with PyTorch (8.5.1) 的以下模型中:

class NetWidth(nn.Module):
    def __init__(self):
        super().__init__()
        self.conv1 = nn.Conv2d(3, 32, kernel_size=3, padding=1)
        self.conv2 = nn.Conv2d(32, 16, kernel_size=3, padding=1)
        self.fc1 = nn.Linear(16 * 8 * 8, 32)
        self.fc2 = nn.Linear(32, 2)
    
    def forward(self, x):
        out = F.max_pool2d(torch.tanh(self.conv1(x)), 2)
        out = F.max_pool2d(torch.tanh(self.conv2(out)), 2)
        out = out.view(-1, 16 * 8 * 8)
        out = torch.tanh(self.fc1(out))
        out = self.fc2(out)
        return out

我们看到模型接受输入 2.d。带有3 频道的图片和:

  • Conv2d -> 将其发送到相同大小的 32 通道图像
  • max_pool2d(,2) -> 将每个维度的图像大小减半
  • Conv2d -> 发送到相同大小的 16 通道图像
  • max_pool2d(,2) -> 将每个维度的图像大小减半
  • view -> 重塑图像
  • Linear -> 接受大小为 16 * 8 * 8 的张量并发送到大小为 32
  • ...

所以向后工作,我们有:

  • 形状张量16 * 8 * 8
  • 未重塑形状(通道 x 高度 x 宽度)
  • un-max_pooled in 2d with factor 2,因此高度和宽度未减半
  • un-convolved from 16 channels to 32
    假设:产品中很可能是16,因此指的是通道数,view看到的图像是形状(频道,8,8),目前是(频道,16,16)2
  • un-max_pooled in 2d with factor 2,因此高度和宽度再次减半(通道,32,32)
  • 从 32 个通道未卷积到 3 个

因此假设 kernel_size 和 padding 足以使卷积本身保持图像尺寸,输入图像的形状可能为 (3,32,32),即 RGB 32x32 像素方形图像。


注意事项:

  1. 即使是外部包pytorch-summary 也要求您提供输入形状,以便显示每一层的输出形状。

  2. 然而,它可以是任何 2 个产生等于 8*8 的数字,例如(64,1), (32,2), (16,4) 等,但是由于代码写为 8*8,因此作者很可能使用了实际尺寸。

【讨论】:

    【解决方案2】:
    print(model)
    

    会给你一个模型的概要,在这里你可以看到每一层的形状。

    您也可以使用pytorch-summary 包。

    如果您的网络将 FC 作为第一层,您可以轻松计算其输入形状。你提到你在前面有一个卷积层。也存在全连接层,网络将只为一种特定的输入大小产生输出。我建议通过使用各种形状来解决这个问题,即喂一个具有某种形状的玩具批次,然后在 FC 层之前检查 Conv 层的输出。

    由于这取决于第一个 FC 层之前的网络架构(conv 层数、内核等),因此我无法为您提供正确输入的准确公式。如前所述,您必须通过尝试各种输入形状以及在第一个 FC 之前得到的网络输出来解决这个问题。 (几乎)总有办法用代码解决问题,但我现在想不出别的办法。

    【讨论】:

    • 但这与input_shape无关
    • 根据pytorch的doc,只有in_channelsout_channels需要指定。
    • 是的,这就是定义一个卷积网络。那么?
    • @AlexMetsai,是的,我的模型是 CNN。那么,我应该如何找出输入形状呢? “评估输入是否会'崩溃'代码”是什么意思?我应该继续猜???
    • 这个答案离题了,print()pytorch-summery 不显示输入形状。它们显示每一层的输出形状。
    【解决方案3】:

    您可以从模型参数中的第一个张量获取输入形状。

    例如创建一些模型:

    class CustomNet(nn.Module):
        def __init__(self):
            super().__init__()
            self.fc1 = nn.Linear(1568, 256)
            self.fc2 = nn.Linear(256, 256)
            self.fc3 = nn.Linear(256, 20)
    
        def forward(self, x):
            out = self.fc1(x)
            out = F.relu(out)
            out = self.fc2(out)
            out = F.relu(out)
            out = self.fc3(out)
            return out
    
    model = CustomNet()
    

    所以model.parameters() 方法返回一个迭代器,该迭代器覆盖了torch.Tensor 类的模块参数。查看文档https://pytorch.org/docs/stable/generated/torch.nn.Module.html#torch.nn.Module.parameters

    第一个参数是输入张量。

    first_parameter = next(model.parameters())
    input_shape = first_parameter.size()
    

    【讨论】:

    • 嗨亚历山大! 1. 您的代码在first_parameter = next(module.parameters()) 中有错字。它应该是 --> model.parameters()。 2.如果网络从FC层开始,似乎就可以了。
    猜你喜欢
    • 2020-09-09
    • 2020-08-21
    • 2020-10-03
    • 2021-01-10
    • 2019-01-26
    • 2019-12-12
    • 2020-09-26
    • 1970-01-01
    相关资源
    最近更新 更多