# With all the changes, what are the options for ADVI training on GPUs?

**URL:** <https://discourse.pymc.io/t/with-all-the-changes-what-are-the-options-for-advi-training-on-gpus/7918>\
**Category:** Questions\
**Tags:** gpu\
**Created:** [August 20, 2021, 12:50pm UTC](https://discourse.pymc.io/t/with-all-the-changes-what-are-the-options-for-advi-training-on-gpus/7918 "2021-08-20T12:50:29Z")\
**Posts on this page:** 7\
**Page:** 1

<div class="post-metadata">

**Author:** ![vitkl](https://yyz2.discourse-cdn.com/flex036/user_avatar/discourse.pymc.io/vitkl/32/3380_2.png) [@vitkl](https://discourse.pymc.io/u/vitkl)\
**Post date:** [August 20, 2021, 12:50pm UTC](https://discourse.pymc.io/t/with-all-the-changes-what-are-the-options-for-advi-training-on-gpus/7918/1 "2021-08-20T12:50:29Z")

</div>

Hi

I am trying to understand what are the current options for using ADVI on GPU. These 3 options come to mind:

1. aesara compilation to C (using pygpu)
2. theano-pymc compilation to C (using pygpu)
3. aesara + JAX as mentioned here [Pymc3-3.11.0 with GPU support - #9 by twiecki](https://discourse.pymc.io/t/pymc3-3-11-0-with-gpu-support/7288/9)

It seems that approach #1 is not yet recommended ([Aesara, theano, theano-pymc - #3 by ricardoV94](https://discourse.pymc.io/t/aesara-theano-theano-pymc/7499/3)) and also it does not work in practice [Moving to pymc3 v4 (replaced theano with aesera) by vitkl · Pull Request #59 · BayraktarLab/cell2location · GitHub\<](https://github.com/BayraktarLab/cell2location/pull/59#issuecomment-902693568).  
Approach #2 does not work for me with the same errors as discussed here [https://discourse.pymc.io/t/pymc3-3-11-0-with-gpu-support/](https://discourse.pymc.io/t/pymc3-3-11-0-with-gpu-support/).  
Approach #3 seems quite experimental. In addition, I found that JAX uses 2x GPU memory compared to pymc3+theano and pyro.

Based on this I can conclude that currently there is no way to use pymc3 ADVI on GPU. Am I wrong or is this a good time to start switching to pymc3 4.0 + aesara?

---

<div class="post-metadata">

**Author:** ![twiecki](https://yyz2.discourse-cdn.com/flex036/user_avatar/discourse.pymc.io/twiecki/32/6930_2.png) [@twiecki](https://discourse.pymc.io/u/twiecki)\
**Post date:** [October 5, 2021, 6:40am UTC](https://discourse.pymc.io/t/with-all-the-changes-what-are-the-options-for-advi-training-on-gpus/7918/2 "2021-10-05T06:40:14Z")

</div>

Yes, support for pygpu is not working and will be dropped. JAX is the way to go but we still have to add VI support for PyMC 4.0 ([https://github.com/pymc-devs/pymc/pull/4582](https://github.com/pymc-devs/pymc/pull/4582)). But then that would be the way to go.

Are you sure that you need it though? Usually slow models can be sped up a lot by better parameterization.

---

<div class="post-metadata">

**Author:** ![la-sekretar](https://yyz2.discourse-cdn.com/flex036/user_avatar/discourse.pymc.io/la-sekretar/32/3589_2.png) [@la-sekretar](https://discourse.pymc.io/u/la-sekretar)\
**Post date:** [July 2, 2022, 11:29am UTC](https://discourse.pymc.io/t/with-all-the-changes-what-are-the-options-for-advi-training-on-gpus/7918/3 "2022-07-02T11:29:32Z")

</div>

Hello,  
Since pymc 4.0 has been released, what’s the update on this? Does ADVI works with aesara + JAX now and how to set it up?

---

<div class="post-metadata">

**Author:** ![twiecki](https://yyz2.discourse-cdn.com/flex036/user_avatar/discourse.pymc.io/twiecki/32/6930_2.png) [@twiecki](https://discourse.pymc.io/u/twiecki)\
**Post date:** [July 2, 2022, 3:21pm UTC](https://discourse.pymc.io/t/with-all-the-changes-what-are-the-options-for-advi-training-on-gpus/7918/4 "2022-07-02T15:21:31Z")

</div>

In principle it should, you can try:

```auto
import aesara
aesara.config["mode"] = "JAX"

```

And run ADVI.

---

<div class="post-metadata">

**Author:** ![la-sekretar](https://yyz2.discourse-cdn.com/flex036/user_avatar/discourse.pymc.io/la-sekretar/32/3589_2.png) [@la-sekretar](https://discourse.pymc.io/u/la-sekretar)\
**Post date:** [July 4, 2022, 12:04pm UTC](https://discourse.pymc.io/t/with-all-the-changes-what-are-the-options-for-advi-training-on-gpus/7918/5 "2022-07-04T12:04:05Z")

</div>

Using pymc version ‘4.0.1’; aesara version ‘2.7.3’

```auto
Input In [21], in <module>
     11 import aesara
---> 12 aesara.config["mode"] = "JAX"

TypeError: 'AesaraConfigParser' object does not support item assignment

```

---

<div class="post-metadata">

**Author:** ![la-sekretar](https://yyz2.discourse-cdn.com/flex036/user_avatar/discourse.pymc.io/la-sekretar/32/3589_2.png) [@la-sekretar](https://discourse.pymc.io/u/la-sekretar)\
**Post date:** [July 4, 2022, 12:35pm UTC](https://discourse.pymc.io/t/with-all-the-changes-what-are-the-options-for-advi-training-on-gpus/7918/6 "2022-07-04T12:35:24Z")

</div>

it seems that

```auto
import aesara
aesara.config.mode = "JAX"

```

works, but somehow it made the inference even a bit slower?

---

<div class="post-metadata">

**Author:** ![twiecki](https://yyz2.discourse-cdn.com/flex036/user_avatar/discourse.pymc.io/twiecki/32/6930_2.png) [@twiecki](https://discourse.pymc.io/u/twiecki)\
**Post date:** [July 4, 2022, 3:38pm UTC](https://discourse.pymc.io/t/with-all-the-changes-what-are-the-options-for-advi-training-on-gpus/7918/7 "2022-07-04T15:38:00Z")

</div>

Yeah, that’s certainly possible. This is still untested with ADVI and as ADVI is implemented in aesara it all gets compiled to C by default already, while our samplers are written in Python, so using JAX samplers removes Python overhead.

I would imagine you can still get speed-ups with JAX if you run on the GPU.
