1414
1515import io
1616import operator
17+ import warnings
1718
1819import cloudpickle
1920import numpy as np
2425import pymc as pm
2526import pymc .variational .opvi as opvi
2627
28+ from pymc .model .transform .basic import remove_minibatched_nodes
2729from pymc .variational .inference import ADVI , ASVGD , SVGD , FullRankADVI
2830from pymc .variational .opvi import NotImplementedInference
2931from 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+
447485def 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
0 commit comments