【问题标题】:forward() not overridden in implementation of nn.Module in an example示例中的 nn.Module 实现中未覆盖 forward()
【发布时间】:2021-10-12 20:16:13
【问题描述】:

this 示例中,我们看到nn.Module 的以下实现:

class Net(torch.nn.Module):
    def __init__(self, in_channels, hidden_channels, out_channels):
        super().__init__()
        self.conv1 = GCNConv(in_channels, hidden_channels)
        self.conv2 = GCNConv(hidden_channels, out_channels)

    def encode(self, x, edge_index):
        x = self.conv1(x, edge_index).relu()
        return self.conv2(x, edge_index)

    def decode(self, z, edge_label_index):
        return (z[edge_label_index[0]] * z[edge_label_index[1]]).sum(dim=-1)

    def decode_all(self, z):
        prob_adj = z @ z.t()
        return (prob_adj > 0).nonzero(as_tuple=False).t()

但是,在docs 中,我们有'forward(*input)'“应该被所有子类覆盖。”

为什么在这个例子中不是这样?

【问题讨论】:

    标签: module pytorch forward


    【解决方案1】:

    这个Net 模块旨在通过两个独立的接口encoderdecode 使用,至少看起来是这样......因为它没有forward 实现,那么是的,它不正确继承自nn.Module。但是,代码仍然“有效”,并且可以正常运行,但如果您使用前向挂钩,可能会产生一些副作用。

    nn.Module 执行推理的标准方法是调用对象,调用__call__ 函数。这个__call__函数是由父类nn.Module实现的,它会依次做两件事:

    • 在推理调用之前或之后处理前向挂钩
    • 调用类的forward函数。

    __call__ 函数充当forward 的包装器。 因此,出于这个原因,forward 函数预计将被用户定义的nn.Module 覆盖。违反此设计模式的唯一警告是,它将有效地忽略应用于nn.Module 的任何钩子。

    【讨论】:

      猜你喜欢
      • 1970-01-01
      • 1970-01-01
      • 1970-01-01
      • 1970-01-01
      • 2013-09-04
      • 2016-02-11
      • 1970-01-01
      • 2017-09-27
      • 1970-01-01
      相关资源
      最近更新 更多