# Pymc 4.0 and variational inference

**URL:** <https://discourse.pymc.io/t/pymc-4-0-and-variational-inference/9652>\
**Category:** General\
**Created:** [June 14, 2022, 10:30pm UTC](https://discourse.pymc.io/t/pymc-4-0-and-variational-inference/9652 "2022-06-14T22:30:48Z")\
**Posts on this page:** 10\
**Page:** 1

<div class="post-metadata">

**Author:** ![Erlebach](https://yyz2.discourse-cdn.com/flex036/user_avatar/discourse.pymc.io/erlebach/32/6605_2.png) [@Erlebach](https://discourse.pymc.io/u/Erlebach)\
**Post date:** [June 14, 2022, 10:30pm UTC](https://discourse.pymc.io/t/pymc-4-0-and-variational-inference/9652/1 "2022-06-14T22:30:48Z")

</div>

Is it possible to call `pm.fit('advi')` and specify that I want to use the Adam optimizer for a faster inference? I assume that stochastic gradient descent is used by default. Thanks.

---

<div class="post-metadata">

**Author:** ![cluhmann](https://yyz2.discourse-cdn.com/flex036/user_avatar/discourse.pymc.io/cluhmann/32/3083_2.png) [@cluhmann](https://discourse.pymc.io/u/cluhmann)\
**Post date:** [June 15, 2022, 2:30am UTC](https://discourse.pymc.io/t/pymc-4-0-and-variational-inference/9652/2 "2022-06-15T02:30:20Z")

</div>

I suspect that you are looking for something like this:

```python
pm.fit(n=1000, obj_optimizer=adam())

```

The various update methods can be found [here](https://www.pymc.io/projects/docs/en/latest/api/vi.html#special).

---

<div class="post-metadata">

**Author:** ![Erlebach](https://yyz2.discourse-cdn.com/flex036/user_avatar/discourse.pymc.io/erlebach/32/6605_2.png) [@Erlebach](https://discourse.pymc.io/u/Erlebach)\
**Post date:** [June 15, 2022, 10:32am UTC](https://discourse.pymc.io/t/pymc-4-0-and-variational-inference/9652/3 "2022-06-15T10:32:14Z")

</div>

Thanks.  
I would like to explore my options with pymc version 4. I searched for info on the variational API and found the following for pymc3: [Variational API quickstart — PyMC3 3.11.5 documentation](https://docs.pymc.io/en/v3/pymc-examples/examples/variational_inference/variational_api_quickstart.html)

I have not found an updated version for pymc3 version 4. I am wondering whether the examples demonstrating the use of callbacks will have changed. Thanks.

I tried your suggestion and got an error message stating that `name 'adam' is not defined`. Here are the details.

Code:

```python
def run_and_plot(model, seed, nb_iter=10000):
    # with model:
    # vi_fit2 = pm.fit(method='svgd', n=nb_iter, random_seed=seed)
        
    with model:
        vi_fit2 = pm.fit(method='advi', n=nb_iter, random_seed=seed, obj_optimizer=adam())
        
    trace5 = vi_fit2.sample(10000)
    pm.plot_trace(trace5);
    
    fig, ax = plt.subplots(figsize=(8, 6))
    plot_w = np.arange(K) + 1 
    ax.bar(plot_w - 0.5, trace5.posterior['w'].squeeze().mean(axis=0), width=1., lw=1, ec='w');
    ax.set_xlim(0.5, K);
    ax.set_xlabel('Component');
    ax.set_ylabel('Posterior expected mixture weight');
    plt.savefig("figure.png")
    
    mean_w = np.mean(trace5.posterior['w'].squeeze(), axis=0)
    nonzero_component = np.where(mean_w > 0.05)[0]

    mean_theta = np.mean(trace5.posterior['theta'].squeeze(), axis=0)
    print("mean_theta[nonzero_component]:\n", mean_theta[nonzero_component, :])
    print("theta_actual:\n", theta_actual)

run_and_plot(model, seed=3437, nb_iter=10000)

```

Error:

```python
---------------------------------------------------------------------------
NameError Traceback (most recent call last)
Input In [215], in <cell line: 1>()
----> 1 run_and_plot(model, seed=3437, nb_iter=10000)

Input In [214], in run_and_plot(model, seed, nb_iter)
      1 def run_and_plot(model, seed, nb_iter=10000):
      2 # with model:
      3 # vi_fit2 = pm.fit(method='svgd', n=nb_iter, random_seed=seed)
      5 with model:
----> 6 vi_fit2 = pm.fit(method='advi', n=nb_iter, random_seed=seed, obj_optimizer=adam())
      8 trace5 = vi_fit2.sample(10000)
      9 pm.plot_trace(trace5);

NameError: name 'adam' is not defined

```

A working example here or in the documentation would be nice! Thank you!

---

<div class="post-metadata">

**Author:** ![Erlebach](https://yyz2.discourse-cdn.com/flex036/user_avatar/discourse.pymc.io/erlebach/32/6605_2.png) [@Erlebach](https://discourse.pymc.io/u/Erlebach)\
**Post date:** [June 15, 2022, 11:36am UTC](https://discourse.pymc.io/t/pymc-4-0-and-variational-inference/9652/4 "2022-06-15T11:36:19Z")

</div>

I tried a callback in the simplest way possible and got an error related to the number of argument expected in ` __call__ `. I am at a loss!

Code:

```python
class Callback:
    def __call__ (self):
        raise NotImplementedError
        
class Writeout(Callback):
    def __init__ (self):
        pass
    
    def __call__ (self):
        print("gordon")
        
callback = Writeout()

def run_and_plot(model, seed, nb_iter=10000):
    # with model:
    # vi_fit2 = pm.fit(method='svgd', n=nb_iter, random_seed=seed)
        
    with model:
       # the `callbacks` argument must be iterable (not the case in pymc3)
        vi_fit2 = pm.fit(method='advi', n=nb_iter, random_seed=seed, callbacks=[callback])
                         # obj_optimizer=adam())
        
    trace5 = vi_fit2.sample(10000)
    pm.plot_trace(trace5);
    
    fig, ax = plt.subplots(figsize=(8, 6))
    plot_w = np.arange(K) + 1 
    ax.bar(plot_w - 0.5, trace5.posterior['w'].squeeze().mean(axis=0), width=1., lw=1, ec='w');
    ax.set_xlim(0.5, K);
    ax.set_xlabel('Component');
    ax.set_ylabel('Posterior expected mixture weight');
    plt.savefig("figure.png")
    
    mean_w = np.mean(trace5.posterior['w'].squeeze(), axis=0)
    nonzero_component = np.where(mean_w > 0.05)[0]

    mean_theta = np.mean(trace5.posterior['theta'].squeeze(), axis=0)
    print("mean_theta[nonzero_component]:\n", mean_theta[nonzero_component, :])
    print("theta_actual:\n", theta_actual)

run_and_plot(model, seed=3437, nb_iter=100)

```

And the error:

```auto
---------------------------------------------------------------------------
TypeError Traceback (most recent call last)
Input In [251], in <cell line: 1>()
----> 1 run_and_plot(model, seed=3437, nb_iter=10000)

Input In [250], in run_and_plot(model, seed, nb_iter)
      1 def run_and_plot(model, seed, nb_iter=10000):
      2 # with model:
      3 # vi_fit2 = pm.fit(method='svgd', n=nb_iter, random_seed=seed)
      5 with model:
----> 6 vi_fit2 = pm.fit(method='advi', n=nb_iter, random_seed=seed, callbacks=[callback])
      7 # obj_optimizer=adam())
      9 trace5 = vi_fit2.sample(10000)

File ~/opt/anaconda3/envs/pymc4/lib/python3.9/site-packages/pymc/variational/inference.py:765, in fit(n, local_rv, method, model, random_seed, start, inf_kwargs, **kwargs)
    763 else:
    764 raise TypeError(f"method should be one of {set(_select.keys())} or Inference instance")
--> 765 return inference.fit(n, **kwargs)

File ~/opt/anaconda3/envs/pymc4/lib/python3.9/site-packages/pymc/variational/inference.py:144, in Inference.fit(self, n, score, callbacks, progressbar, **kwargs)
    142 progress = range(n)
    143 if score:
--> 144 state = self._iterate_with_loss(0, n, step_func, progress, callbacks)
    145 else:
    146 state = self._iterate_without_loss(0, n, step_func, progress, callbacks)

File ~/opt/anaconda3/envs/pymc4/lib/python3.9/site-packages/pymc/variational/inference.py:240, in Inference._iterate_with_loss(self, s, n, step_func, progress, callbacks)
    238 progress.comment = f"Average Loss = {avg_loss:,.5g}"
    239 for callback in callbacks:
--> 240 callback(self.approx, scores[: i + 1], i + s + 1)
    241 except (KeyboardInterrupt, StopIteration) as e: # pragma: no cover
    242 # do not print log on the same line
    243 scores = scores[:i]

TypeError: __call__ () takes 1 positional argument but 4 were given

```

---

<div class="post-metadata">

**Author:** ![cluhmann](https://yyz2.discourse-cdn.com/flex036/user_avatar/discourse.pymc.io/cluhmann/32/3083_2.png) [@cluhmann](https://discourse.pymc.io/u/cluhmann)\
**Post date:** [June 15, 2022, 12:52pm UTC](https://discourse.pymc.io/t/pymc-4-0-and-variational-inference/9652/5 "2022-06-15T12:52:55Z")

</div>

Apologies, my code snippet was written quickly and should have been this:

```python
pm.fit(n=1000, obj_optimizer=pm.adam())

```

The signature for the callback functions is this:

```python
callbacks: list[function: (Approximation, losses, i) -> None]

```

The source of the callbacks can be found [here](https://github.com/pymc-devs/pymc/blob/37ba9a3e3a19b738f48cb30007f4d70c33bdd0f6/pymc/variational/callbacks.py), but it’s likely you just want `pm.callbacks.CheckParametersConvergence()`. So something like this:

```python
with pm.Model() as model:
    x = pm.Normal('x', mu=0, sigma=1)
    y = pm.Normal('y', mu=x, sigma=1)

    vi_fit2 = pm.fit(method='advi',
                     n=1000,
                     callbacks=[pm.callbacks.CheckParametersConvergence()],
                     obj_optimizer=pm.adam()
                    )
        
    trace = vi_fit2.sample(10000)

```

---

<div class="post-metadata">

**Author:** ![Erlebach](https://yyz2.discourse-cdn.com/flex036/user_avatar/discourse.pymc.io/erlebach/32/6605_2.png) [@Erlebach](https://discourse.pymc.io/u/Erlebach)\
**Post date:** [June 15, 2022, 2:02pm UTC](https://discourse.pymc.io/t/pymc-4-0-and-variational-inference/9652/6 "2022-06-15T14:02:07Z")

</div>

I got it to work. However, now I have to learn to use the output of the callbacks. Will report back later. 🙂

---

<div class="post-metadata">

**Author:** ![Erlebach](https://yyz2.discourse-cdn.com/flex036/user_avatar/discourse.pymc.io/erlebach/32/6605_2.png) [@Erlebach](https://discourse.pymc.io/u/Erlebach)\
**Post date:** [June 16, 2022, 2:14am UTC](https://discourse.pymc.io/t/pymc-4-0-and-variational-inference/9652/7 "2022-06-16T02:14:54Z")

</div>

Now that I have callbacks working, I would like to access its output. I cannot figure out the logic. Is there information out there explaining how it works (aside from looking at the source code?) Thanks.

---

<div class="post-metadata">

**Author:** ![Erlebach](https://yyz2.discourse-cdn.com/flex036/user_avatar/discourse.pymc.io/erlebach/32/6605_2.png) [@Erlebach](https://discourse.pymc.io/u/Erlebach)\
**Post date:** [June 16, 2022, 3:37pm UTC](https://discourse.pymc.io/t/pymc-4-0-and-variational-inference/9652/8 "2022-06-16T15:37:28Z")

</div>

I cannot figure out how to access the data stored in the callback `CheckParametersConvergence()` . Any help would be appreciated! Thanks.

---

<div class="post-metadata">

**Author:** ![cluhmann](https://yyz2.discourse-cdn.com/flex036/user_avatar/discourse.pymc.io/cluhmann/32/3083_2.png) [@cluhmann](https://discourse.pymc.io/u/cluhmann)\
**Post date:** [June 16, 2022, 3:59pm UTC](https://discourse.pymc.io/t/pymc-4-0-and-variational-inference/9652/9 "2022-06-16T15:59:09Z")

</div>

I think [the quickstart](https://www.pymc.io/projects/examples/en/latest/variational_inference/variational_api_quickstart.html) illustrates how to interrogate the convergence.

---

<div class="post-metadata">

**Author:** ![Erlebach](https://yyz2.discourse-cdn.com/flex036/user_avatar/discourse.pymc.io/erlebach/32/6605_2.png) [@Erlebach](https://discourse.pymc.io/u/Erlebach)\
**Post date:** [June 16, 2022, 6:05pm UTC](https://discourse.pymc.io/t/pymc-4-0-and-variational-inference/9652/10 "2022-06-16T18:05:32Z")

</div>

I went through the Quickstart notebook in detail and tried things out before writing on discourse. Consider the following line:

```python
with model:
    mean_field = pm.fit(method="advi", callbacks=[CheckParametersConvergence()])
plt.plot(mean_field.hist);

```

I tried this (with pymc 4). I would like to do the following: update an array with my own diagnostics at a specified frequency. So the question is how to retrieve my array from the callback? I guess I must create my own callback? I might experiment with that.

It would be nice to check results every n iterations? Is that done by running `pm.sample` at a lower frequency, and simply restarting the simulation ?I have not tried this yet. Thanks!

I created my own class, copying `CheckParametersConvergence`. I was able to get the results allowed. I will now experiment with `Tracker`.
