Skip to content

Commit 6182801

Browse files
authored
Approximation's sample method uses model contexts
1 parent 4e28043 commit 6182801

5 files changed

Lines changed: 75 additions & 31 deletions

File tree

pymc/sampling/mcmc.py

Lines changed: 12 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -1731,9 +1731,10 @@ def model_logp_fn(ip: PointType) -> np.ndarray:
17311731
obj_optimizer=pm.adagrad_window,
17321732
compile_kwargs=compile_kwargs,
17331733
)
1734-
approx_sample = approx.sample(
1735-
draws=chains, random_seed=random_seed_list[0], return_inferencedata=False
1736-
)
1734+
with model:
1735+
approx_sample = approx.sample(
1736+
draws=chains, random_seed=random_seed_list[0], return_inferencedata=False
1737+
)
17371738
initial_points = [
17381739
{k: np.asarray(v) for k, v in approx_sample[i].items()} for i in range(chains)
17391740
]
@@ -1757,9 +1758,10 @@ def model_logp_fn(ip: PointType) -> np.ndarray:
17571758
obj_optimizer=pm.adagrad_window,
17581759
compile_kwargs=compile_kwargs,
17591760
)
1760-
approx_sample = approx.sample(
1761-
draws=chains, random_seed=random_seed_list[0], return_inferencedata=False
1762-
)
1761+
with model:
1762+
approx_sample = approx.sample(
1763+
draws=chains, random_seed=random_seed_list[0], return_inferencedata=False
1764+
)
17631765
initial_points = [
17641766
{k: np.asarray(v) for k, v in approx_sample[i].items()} for i in range(chains)
17651767
]
@@ -1777,9 +1779,10 @@ def model_logp_fn(ip: PointType) -> np.ndarray:
17771779
obj_optimizer=pm.adagrad_window,
17781780
compile_kwargs=compile_kwargs,
17791781
)
1780-
approx_sample = approx.sample(
1781-
draws=chains, random_seed=random_seed_list[0], return_inferencedata=False
1782-
)
1782+
with model:
1783+
approx_sample = approx.sample(
1784+
draws=chains, random_seed=random_seed_list[0], return_inferencedata=False
1785+
)
17831786
initial_points = [
17841787
{k: np.asarray(v) for k, v in approx_sample[i].items()} for i in range(chains)
17851788
]

pymc/variational/approximations.py

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -329,7 +329,8 @@ def sample_approx(approx, draws=100, include_transformed=True):
329329
trace: class:`pymc.backends.base.MultiTrace`
330330
Samples drawn from variational posterior.
331331
"""
332-
return approx.sample(draws=draws, include_transformed=include_transformed)
332+
with approx.model:
333+
return approx.sample(draws=draws, include_transformed=include_transformed)
333334

334335

335336
# single group shortcuts exported to user

pymc/variational/opvi.py

Lines changed: 11 additions & 12 deletions
Original file line numberDiff line numberDiff line change
@@ -1229,7 +1229,10 @@ def __init__(self, groups, model=None):
12291229
else:
12301230
rest.__init_group__(unseen_free_RVs)
12311231
self.groups.append(rest)
1232-
self.model = model
1232+
1233+
@property
1234+
def model(self):
1235+
return modelcontext(self.groups[0].model if self.groups else None)
12331236

12341237
@property
12351238
def has_logq(self):
@@ -1503,14 +1506,8 @@ def rslice(self, name):
15031506
15041507
This node still needs :func:`set_size_and_deterministic` to be evaluated.
15051508
"""
1506-
1507-
def vars_names(vs):
1508-
return {self.model.rvs_to_values[v].name for v in vs}
1509-
1510-
for vars_, random, ordering in zip(
1511-
self.collect("group"), self.symbolic_randoms, self.collect("ordering")
1512-
):
1513-
if name in vars_names(vars_):
1509+
for random, ordering in zip(self.symbolic_randoms, self.collect("ordering")):
1510+
if name in ordering:
15141511
name_, slc, shape, dtype = ordering[name]
15151512
found = random[..., slc].reshape((random.shape[0], *shape)).astype(dtype)
15161513
found.name = name + "_vi_random_slice"
@@ -1522,7 +1519,7 @@ def vars_names(vs):
15221519
@node_property
15231520
def sample_dict_fn(self):
15241521
s = pt.iscalar()
1525-
names = [self.model.rvs_to_values[v].name for v in self.model.free_RVs]
1522+
names = [name for ordering in self.collect("ordering") for name in ordering]
15261523
sampled = [self.rslice(name) for name in names]
15271524
sampled = self.set_size_and_deterministic(sampled, s, 0)
15281525
sample_fn = compile([s], sampled)
@@ -1567,8 +1564,10 @@ def sample(
15671564
for i in range(draws)
15681565
)
15691566

1567+
model = modelcontext(None)
1568+
15701569
trace = NDArray(
1571-
model=self.model,
1570+
model=model,
15721571
test_point={name: np.asarray(records[0]) for name, records in samples.items()},
15731572
)
15741573
try:
@@ -1582,7 +1581,7 @@ def sample(
15821581
if not return_inferencedata:
15831582
return multi_trace
15841583
else:
1585-
return pm.to_inference_data(multi_trace, model=self.model, **kwargs)
1584+
return pm.to_inference_data(multi_trace, model=model, **kwargs)
15861585

15871586
@property
15881587
def ndim(self):

tests/variational/test_inference.py

Lines changed: 41 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -14,6 +14,7 @@
1414

1515
import io
1616
import operator
17+
import warnings
1718

1819
import cloudpickle
1920
import numpy as np
@@ -24,6 +25,7 @@
2425
import pymc as pm
2526
import pymc.variational.opvi as opvi
2627

28+
from pymc.model.transform.basic import remove_minibatched_nodes
2729
from pymc.variational.inference import ADVI, ASVGD, SVGD, FullRankADVI
2830
from pymc.variational.opvi import NotImplementedInference
2931
from tests import models
@@ -174,8 +176,9 @@ def fit_kwargs(inference, use_minibatch):
174176
return _select[(type(inference), key)]
175177

176178

177-
def test_fit_oo(inference, fit_kwargs, simple_model_data):
178-
trace = inference.fit(**fit_kwargs).sample(10000)
179+
def test_fit_oo(inference, fit_kwargs, simple_model, simple_model_data):
180+
with simple_model:
181+
trace = inference.fit(**fit_kwargs).sample(10000)
179182
mu_post = simple_model_data["mu_post"]
180183
d = simple_model_data["d"]
181184
np.testing.assert_allclose(np.mean(trace.posterior["mu"]), mu_post, rtol=0.05)
@@ -202,7 +205,8 @@ def test_fit_start(inference_spec, simple_model):
202205
inference = inference_spec(**kw)
203206

204207
try:
205-
trace = inference.fit(n=0).sample(10000)
208+
with simple_model:
209+
trace = inference.fit(n=0).sample(10000)
206210
except NotImplementedInference as e:
207211
pytest.skip(str(e))
208212

@@ -444,6 +448,40 @@ def test_fit_data_coords(hierarchical_model, hierarchical_model_data):
444448
assert data["mu"].shape == ()
445449

446450

451+
def test_sample_posterior_after_minibatch():
452+
with pm.Model(coords={"obs_id": [0, 1, 2]}) as model:
453+
x = pm.Data("x", [1.0, 2.0, 3.0], dims="obs_id")
454+
y = pm.Data("y", [1.0, 2.0, 3.0], dims="obs_id")
455+
x_mini, y_mini = pm.Minibatch(x, y, batch_size=2)
456+
beta = pm.Normal("beta", 0, 10.0)
457+
y_hat = pm.Deterministic("y_hat", beta * x_mini, dims="obs_id")
458+
pm.Normal("obs", y_hat, np.sqrt(1e-2), observed=y_mini, total_size=3, dims="obs_id")
459+
approx = pm.fit(
460+
10,
461+
method="advi",
462+
progressbar=False,
463+
)
464+
465+
model_post = remove_minibatched_nodes(model)
466+
with model_post:
467+
trace = approx.sample(500)
468+
469+
assert trace.posterior["beta"].shape == (1, 500)
470+
assert trace.constant_data["x"].shape == (3,)
471+
assert trace.observed_data["obs"].shape == (3,)
472+
473+
with model_post, warnings.catch_warnings():
474+
warnings.filterwarnings("ignore", "Numba will use object mode", UserWarning)
475+
x_test = [5, 6, 9, 12, 15]
476+
pm.set_data(
477+
new_data={"x": x_test, "y": [0.0] * len(x_test)},
478+
coords={"obs_id": list(range(len(x_test)))},
479+
)
480+
y_test = pm.sample_posterior_predictive(trace, predictions=True, progressbar=False)
481+
482+
assert y_test.predictions["obs"].shape == (1, 500, 5)
483+
484+
447485
def test_multiple_minibatch_variables():
448486
"""Regression test for bug reported in
449487
https://discourse.pymc.io/t/verifying-that-minibatch-is-actually-randomly-sampling/14308

tests/variational/test_opvi.py

Lines changed: 9 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -208,24 +208,27 @@ def three_var_approx_single_group_mf(three_var_model):
208208
return MeanField(model=three_var_model)
209209

210210

211-
def test_pickle_approx(three_var_approx):
211+
def test_pickle_approx(three_var_model, three_var_approx):
212212
import cloudpickle
213213

214214
dump = cloudpickle.dumps(three_var_approx)
215215
new = cloudpickle.loads(dump)
216-
assert new.sample(1)
216+
with three_var_model:
217+
assert new.sample(1)
217218

218219

219-
def test_pickle_single_group(three_var_approx_single_group_mf):
220+
def test_pickle_single_group(three_var_model, three_var_approx_single_group_mf):
220221
import cloudpickle
221222

222223
dump = cloudpickle.dumps(three_var_approx_single_group_mf)
223224
new = cloudpickle.loads(dump)
224-
assert new.sample(1)
225+
with three_var_model:
226+
assert new.sample(1)
225227

226228

227-
def test_sample_simple(three_var_approx):
228-
trace = three_var_approx.sample(100, return_inferencedata=False)
229+
def test_sample_simple(three_var_model, three_var_approx):
230+
with three_var_model:
231+
trace = three_var_approx.sample(100, return_inferencedata=False)
229232
assert set(trace.varnames) == {"one", "one_log__", "three", "two"}
230233
assert len(trace) == 100
231234
assert trace[0]["one"].shape == (10, 2)

0 commit comments

Comments
 (0)