From 648fa28b588811d4db2c32b1cd9a00594606264a Mon Sep 17 00:00:00 2001 From: OlaRonning Date: Tue, 2 Aug 2022 16:34:26 +0200 Subject: [PATCH 1/7] added local latent model test case. --- tests/infer/test_inference.py | 71 +++++++++++++++++++++++++++++++++++ 1 file changed, 71 insertions(+) diff --git a/tests/infer/test_inference.py b/tests/infer/test_inference.py index a598ea9c54..a6d0096b54 100644 --- a/tests/infer/test_inference.py +++ b/tests/infer/test_inference.py @@ -1005,3 +1005,74 @@ def guide(data, weights): loss = svi.step(data, weights) if step % 20 == 0: logger.info("step {} loss = {:0.4g}".format(step, loss)) + + +@pytest.mark.stage("integration", "integration_batch_2") +class OneWayNormalRandomEffects(TestCase): + def setUp(self) -> None: + self.n_groups = 3 + self.n_experiments = 5 + self.data = torch.tensor( + [ + [4.1, 3.5, 0.2, -3.3, 3.3], + [2.4, -6.5, -0.7, 4.4, -4.8], + [1.1, -0.6, 1.3, -1.3, -1.1], + ] + ) + self.group_locs = torch.tensor([[3.0], [-2.0], [0.0]]) + self.group_prec = torch.tensor([[0.2], [0.1], [0.3]]) + self.obs_prec = torch.tensor(6.0) + obs_prec = self.obs_prec * self.n_experiments + self.post_group_locs = ( + self.data.mean(1, keepdim=True) * obs_prec + + self.group_locs * self.group_prec + ) / (obs_prec + self.group_prec) + + def test_renyi_reparameterized(self): + self.do_elbo_test(True, 5000, RenyiELBO(num_particles=2)) + + def test_renyi_nonreparameterized(self): + self.do_elbo_test(False, 15000, RenyiELBO(alpha=0.2, num_particles=2)) + + def test_elbo_reparameterized(self): + self.do_elbo_test(True, 5000, Trace_ELBO()) + + def test_elbo_nonreparameterized(self): + self.do_elbo_test(False, 35_000, Trace_ELBO()) + + def do_elbo_test(self, reparameterized, n_steps, loss, debug=False): + def model(): + with pyro.plate("groups", self.n_groups, dim=-2): + group_loc = pyro.sample( + "group_loc", + dist.Normal(self.group_locs, torch.pow(self.group_prec, -0.5)), + ) + with pyro.plate("data", self.n_experiments, dim=-1): + pyro.sample( + "y", + dist.Normal(group_loc, torch.pow(self.obs_prec, -0.5)), + obs=self.data, + ) + + def guide(): + gloc = pyro.param( + "group_loc_param", + self.post_group_locs + torch.tensor([[0.05], [-0.08], [0.14]]), + ) + with pyro.plate("groups", self.n_groups, dim=-2): + Normal = ( + dist.Normal if reparameterized else fakes.NonreparameterizedNormal + ) + pyro.sample("group_loc", Normal(gloc, torch.pow(self.group_prec, -0.5))) + + adam = optim.Adam({"lr": 0.0005, "betas": (0.97, 0.999)}) + svi = SVI(model, guide, adam, loss=loss) + + for k in range(n_steps): + svi.step() + + group_loc_error = param_abs_error("group_loc_param", self.post_group_locs) + assert_equal(0.0, group_loc_error, prec=0.08) + + def do_fit_prior_test(self, reparameterized, n_steps, loss, debug=False): + pass From 58d8018b89df757bbead8e59bf76867a02df880f Mon Sep 17 00:00:00 2001 From: OlaRonning Date: Mon, 8 Aug 2022 15:49:56 +0200 Subject: [PATCH 2/7] changed `get_dependent_plate_dims` to `get_nonparticle_plate_dims`. --- pyro/infer/renyi_elbo.py | 6 +++--- pyro/infer/util.py | 14 +++++++++----- tests/infer/test_inference.py | 15 ++++++++------- 3 files changed, 20 insertions(+), 15 deletions(-) diff --git a/pyro/infer/renyi_elbo.py b/pyro/infer/renyi_elbo.py index 349f7c43d4..0c44bae929 100644 --- a/pyro/infer/renyi_elbo.py +++ b/pyro/infer/renyi_elbo.py @@ -9,7 +9,7 @@ from pyro.distributions.util import is_identically_zero from pyro.infer.elbo import ELBO from pyro.infer.enum import get_importance_trace -from pyro.infer.util import get_dependent_plate_dims, is_validation_enabled, torch_sum +from pyro.infer.util import get_nonparticle_plate_dims, is_validation_enabled, torch_sum from pyro.util import check_if_enumerated, warn_if_nan @@ -104,7 +104,7 @@ def loss(self, model, guide, *args, **kwargs): # grab a vectorized trace from the generator for model_trace, guide_trace in self._get_traces(model, guide, args, kwargs): elbo_particle = 0.0 - sum_dims = get_dependent_plate_dims(model_trace.nodes.values()) + sum_dims = get_nonparticle_plate_dims(model_trace.nodes.values()) # compute elbo for name, site in model_trace.nodes.items(): @@ -152,7 +152,7 @@ def loss_and_grads(self, model, guide, *args, **kwargs): for model_trace, guide_trace in self._get_traces(model, guide, args, kwargs): elbo_particle = 0 surrogate_elbo_particle = 0 - sum_dims = get_dependent_plate_dims(model_trace.nodes.values()) + sum_dims = get_nonparticle_plate_dims(model_trace.nodes.values()) # compute elbo and surrogate elbo for name, site in model_trace.nodes.items(): diff --git a/pyro/infer/util.py b/pyro/infer/util.py index 3ee94b884d..3c8be1e0a1 100644 --- a/pyro/infer/util.py +++ b/pyro/infer/util.py @@ -99,17 +99,21 @@ def get_plate_stacks(trace): } -def get_dependent_plate_dims(sites): +def get_nonparticle_plate_dims(sites): """ - Return a list of unique dims for plates that are not common to all sites. + Return a list of unique dims of all plates except the par """ plate_sets = [ site["cond_indep_stack"] for site in sites if site["type"] == "sample" ] all_plates = set().union(*plate_sets) - common_plates = all_plates.intersection(*plate_sets) - sum_plates = all_plates - common_plates - sum_dims = sorted({f.dim for f in sum_plates if f.dim is not None}) + sum_dims = sorted( + { + f.dim + for f in all_plates + if f.dim is not None and f.name != "num_particles_vectorized" + } + ) return sum_dims diff --git a/tests/infer/test_inference.py b/tests/infer/test_inference.py index a6d0096b54..eda1d31724 100644 --- a/tests/infer/test_inference.py +++ b/tests/infer/test_inference.py @@ -1029,17 +1029,18 @@ def setUp(self) -> None: ) / (obs_prec + self.group_prec) def test_renyi_reparameterized(self): - self.do_elbo_test(True, 5000, RenyiELBO(num_particles=2)) + self.do_elbo_test(True, 10_000, RenyiELBO(num_particles=2)) + + def test_renyi_vectorized(self): + self.do_elbo_test( + True, + 15_000, + RenyiELBO(num_particles=2, vectorize_particles=True, max_plate_nesting=3), + ) def test_renyi_nonreparameterized(self): self.do_elbo_test(False, 15000, RenyiELBO(alpha=0.2, num_particles=2)) - def test_elbo_reparameterized(self): - self.do_elbo_test(True, 5000, Trace_ELBO()) - - def test_elbo_nonreparameterized(self): - self.do_elbo_test(False, 35_000, Trace_ELBO()) - def do_elbo_test(self, reparameterized, n_steps, loss, debug=False): def model(): with pyro.plate("groups", self.n_groups, dim=-2): From 2039bc9702b09c0ce8d133dcbc3805d5572f3f67 Mon Sep 17 00:00:00 2001 From: OlaRonning Date: Mon, 8 Aug 2022 15:52:12 +0200 Subject: [PATCH 3/7] fixed docstring for `get_nonparticle_plate_dims` --- pyro/infer/util.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/pyro/infer/util.py b/pyro/infer/util.py index 3c8be1e0a1..670d31495b 100644 --- a/pyro/infer/util.py +++ b/pyro/infer/util.py @@ -101,7 +101,7 @@ def get_plate_stacks(trace): def get_nonparticle_plate_dims(sites): """ - Return a list of unique dims of all plates except the par + Return a list of unique dims of all plates except vectorized particles """ plate_sets = [ site["cond_indep_stack"] for site in sites if site["type"] == "sample" From df9b73e500ee643bb3a0b2feb1c054816e65398a Mon Sep 17 00:00:00 2001 From: OlaRonning Date: Mon, 8 Aug 2022 15:58:09 +0200 Subject: [PATCH 4/7] removed `do_fit_prior_test` --- tests/infer/test_inference.py | 3 --- 1 file changed, 3 deletions(-) diff --git a/tests/infer/test_inference.py b/tests/infer/test_inference.py index eda1d31724..7ce2cd72a9 100644 --- a/tests/infer/test_inference.py +++ b/tests/infer/test_inference.py @@ -1074,6 +1074,3 @@ def guide(): group_loc_error = param_abs_error("group_loc_param", self.post_group_locs) assert_equal(0.0, group_loc_error, prec=0.08) - - def do_fit_prior_test(self, reparameterized, n_steps, loss, debug=False): - pass From e922d8ac35681db2614fabd3eab77c9cbb4bb5bd Mon Sep 17 00:00:00 2001 From: OlaRonning Date: Tue, 9 Aug 2022 22:20:47 +0200 Subject: [PATCH 5/7] added black to pyproject.toml --- pyproject.toml | 4 ++++ 1 file changed, 4 insertions(+) create mode 100644 pyproject.toml diff --git a/pyproject.toml b/pyproject.toml new file mode 100644 index 0000000000..1268a8fd4a --- /dev/null +++ b/pyproject.toml @@ -0,0 +1,4 @@ +[tool.black] +line-length = 120 +target-version = ['py37'] +include = '*.p pyro examples tests scripts profiler' \ No newline at end of file From 2b544bda901de9355676fc66d749cbe3c8307406 Mon Sep 17 00:00:00 2001 From: OlaRonning Date: Tue, 9 Aug 2022 23:01:56 +0200 Subject: [PATCH 6/7] fixed --include regex and changed to black defaults for formatting --- Makefile | 4 ++-- pyproject.toml | 12 +++++++++--- 2 files changed, 11 insertions(+), 5 deletions(-) diff --git a/Makefile b/Makefile index f98e2cae72..713070cb29 100644 --- a/Makefile +++ b/Makefile @@ -19,7 +19,7 @@ tutorial: FORCE lint: FORCE flake8 - black --check *.py pyro examples tests scripts profiler + black --check . isort --check . python scripts/update_headers.py --check mypy --install-types --non-interactive pyro scripts @@ -28,7 +28,7 @@ license: FORCE python scripts/update_headers.py format: license FORCE - black *.py pyro examples tests scripts profiler + black . isort . version: FORCE diff --git a/pyproject.toml b/pyproject.toml index 1268a8fd4a..95ba5986f9 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -1,4 +1,10 @@ [tool.black] -line-length = 120 -target-version = ['py37'] -include = '*.p pyro examples tests scripts profiler' \ No newline at end of file +include = ''' +( + pyro/.*\.py + | examples/.*\.py + | tests/.*\.py + | scripts/.*\.py + | profiler/.*\.py +) +''' \ No newline at end of file From f5492721703d23af8dcc3e64b7888c55735c7fd7 Mon Sep 17 00:00:00 2001 From: OlaRonning Date: Tue, 9 Aug 2022 23:05:28 +0200 Subject: [PATCH 7/7] added newline end to pyproject.toml --- pyproject.toml | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/pyproject.toml b/pyproject.toml index 95ba5986f9..60398beec7 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -7,4 +7,4 @@ include = ''' | scripts/.*\.py | profiler/.*\.py ) -''' \ No newline at end of file +'''