diff --git a/monai/networks/nets/fullyconnectednet.py b/monai/networks/nets/fullyconnectednet.py index be179e5b59..d7354d4bd2 100644 --- a/monai/networks/nets/fullyconnectednet.py +++ b/monai/networks/nets/fullyconnectednet.py @@ -172,12 +172,25 @@ def decode_forward(self, z: torch.Tensor, use_sigmoid: bool = True) -> torch.Ten return x def reparameterize(self, mu: torch.Tensor, logvar: torch.Tensor) -> torch.Tensor: - std = torch.exp(0.5 * logvar) + """Sample a latent code using the reparameterization trick. + + At inference (eval mode) the posterior mean is returned directly. During + training, returns ``mu + eps * std`` with ``eps ~ N(0, I)``. - if self.training: # multiply random noise with std only during training - std = torch.randn_like(std).mul(std) + Args: + mu: Posterior mean, shape ``(batch, latent_size)``. + logvar: Log-variance of the posterior, same shape as ``mu``. - return std.add_(mu) + Returns: + Sampled latent code, same shape as ``mu``. + """ + if not self.training: + # At inference the latent code is the posterior mean; the random + # term is only added during training (the reparameterization trick). + return mu + std = torch.exp(0.5 * logvar) + eps = torch.randn_like(std) + return mu + eps * std def forward(self, x: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]: mu, logvar = self.encode_forward(x) diff --git a/monai/networks/nets/varautoencoder.py b/monai/networks/nets/varautoencoder.py index 0674094aa7..23c8aa37dc 100644 --- a/monai/networks/nets/varautoencoder.py +++ b/monai/networks/nets/varautoencoder.py @@ -142,12 +142,25 @@ def decode_forward(self, z: torch.Tensor, use_sigmoid: bool = True) -> torch.Ten return x def reparameterize(self, mu: torch.Tensor, logvar: torch.Tensor) -> torch.Tensor: + """Sample a latent code using the reparameterization trick. + + At inference (eval mode) the posterior mean is returned directly. During + training, returns ``mu + eps * std`` with ``eps ~ N(0, I)``. + + Args: + mu: Posterior mean, shape ``(batch, latent_size)``. + logvar: Log-variance of the posterior, same shape as ``mu``. + + Returns: + Sampled latent code, same shape as ``mu``. + """ + if not self.training: + # At inference the latent code is the posterior mean; the random + # term is only added during training (the reparameterization trick). + return mu std = torch.exp(0.5 * logvar) - - if self.training: # multiply random noise with std only during training - std = torch.randn_like(std).mul(std) - - return std.add_(mu) + eps = torch.randn_like(std) + return mu + eps * std def forward(self, x: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]: mu, logvar = self.encode_forward(x) diff --git a/tests/networks/nets/test_fullyconnectednet.py b/tests/networks/nets/test_fullyconnectednet.py index 863d1399a9..b1ee315585 100644 --- a/tests/networks/nets/test_fullyconnectednet.py +++ b/tests/networks/nets/test_fullyconnectednet.py @@ -64,6 +64,30 @@ def test_vfc_shape(self, input_param, input_shape, expected_shape): result = net.forward(torch.randn(input_shape).to(device))[0] self.assertEqual(result.shape, expected_shape) + def test_vfc_reparameterize_eval_returns_mu(self): + """A VFC latent code is deterministic at eval (equals mu) and stochastic at train. + + Regression test for the #8413 reparameterize bug, which returned ``mu + std`` + at inference instead of ``mu``. + """ + net = VarFullyConnectedNet( + in_channels=10, out_channels=10, latent_size=30, encode_channels=(15, 20, 25), decode_channels=(15, 20, 25) + ).to(device) + data = torch.randn(3, 10).to(device) + + with eval_mode(net): + _, mu1, _, z1 = net(data) + _, _, _, z2 = net(data) + self.assertTrue(torch.allclose(z1, mu1)) + self.assertTrue(torch.allclose(z1, z2)) + + net.train() + with torch.no_grad(): + _, mu_t, _, zt1 = net(data) + _, _, _, zt2 = net(data) + self.assertFalse(torch.allclose(zt1, mu_t)) + self.assertFalse(torch.allclose(zt1, zt2)) + if __name__ == "__main__": unittest.main() diff --git a/tests/networks/nets/test_varautoencoder.py b/tests/networks/nets/test_varautoencoder.py index 459c537c55..ce1e2e50e7 100644 --- a/tests/networks/nets/test_varautoencoder.py +++ b/tests/networks/nets/test_varautoencoder.py @@ -122,6 +122,29 @@ def test_script(self): test_data = torch.randn(2, 1, 32, 32) test_script_save(net, test_data, rtol=1e-3, atol=1e-3) + def test_reparameterize_eval_returns_mu(self): + """A VarAutoEncoder latent code is deterministic at eval (equals mu) and stochastic at train. + + Regression test for #8413, where eval returned ``mu + std`` instead of ``mu``. + """ + net = VarAutoEncoder( + spatial_dims=2, in_shape=(1, 32, 32), out_channels=1, latent_size=4, channels=(4, 8), strides=(2, 2) + ).to(device) + data = torch.randn(2, 1, 32, 32).to(device) + + with eval_mode(net): + _, mu1, _, z1 = net(data) + _, _, _, z2 = net(data) + self.assertTrue(torch.allclose(z1, mu1)) + self.assertTrue(torch.allclose(z1, z2)) + + net.train() + with torch.no_grad(): + _, mu_t, _, zt1 = net(data) + _, _, _, zt2 = net(data) + self.assertFalse(torch.allclose(zt1, mu_t)) + self.assertFalse(torch.allclose(zt1, zt2)) + if __name__ == "__main__": unittest.main()