Thanks for this design pattern.
Question: Suppose Network had 2 linear layers. Then, this approach (see below) fails with the error RuntimeError: Multiple sample sites named 'weight'. I think this is because now, in the original Network, two different parameters have the name ``weight’’. What would be a solution to this? Thanks a bunch.
class Network(nn.Module):
Linear = nn.Linear # this can be overridden by derived classes
def __init__(self, in_features=2, out_features=1):
super().__init__()
self.in_features = in_features
self.out_features = out_features
self.linear_1 = self.Linear(self.in_features, self.out_features)
self.linear_2 = self.Linear(self.in_features, self.out_features)
class RandomNetwork(Network, PyroModule):
Linear = PyroModule[nn.Linear]
def __init__(self, in_features, out_features):
super().__init__(in_features, out_features)
self.linear_1.weight = PyroSample(
lambda self: dist.Normal(0, 1)
.expand([self.out_features,
self.in_features])
.to_event(2))
self.linear_2.weight = PyroSample(
lambda self: dist.Normal(0, 1)
.expand([self.out_features,
self.in_features])
.to_event(2))
rand_net = RandomNetwork(2,1)