# Adam+Half Precision = NaNs?

**URL:** <https://discuss.pytorch.org/t/adam-half-precision-nans/1765>\
**Category:** Uncategorized\
**Created:** [April 9, 2017, 8:37pm UTC](https://discuss.pytorch.org/t/adam-half-precision-nans/1765 "2017-04-09T20:37:22Z")\
**Posts on this page:** 16\
**Page:** 1

<div class="post-metadata">

**Author:** ![ajbrock](https://discuss.pytorch.org/user_avatar/discuss.pytorch.org/ajbrock/32/139_2.png) [@ajbrock](https://discuss.pytorch.org/u/ajbrock)\
**Post date:** [April 9, 2017, 8:37pm UTC](https://discuss.pytorch.org/t/adam-half-precision-nans/1765/1 "2017-04-09T20:37:22Z")

</div>

Hi guys,

I’ve been running into the sudden appearance of NaNs when I attempt to train using Adam and Half (float16) precision; my nets train just fine on half precision with SGD+nesterov momentum, and they train just fine with single precision (float32) and Adam, but switching them over to half seems to cause numerical instability. I’ve fiddled with the hyperparams a bit; upping epsilon helps a _tiny_ bit but doesn’t fix the issue.

Is this something anyone else has info on? If not I can throw together a reproduction script and dig into the issue.

Thanks again! Been a good while since I’ve had to post on account of hitting no issues otherwise.

---

<div class="post-metadata">

**Author:** ![smth](https://discuss.pytorch.org/user_avatar/discuss.pytorch.org/smth/32/13_2.png) [@smth](https://discuss.pytorch.org/u/smth)\
**Post date:** [April 9, 2017, 10:08pm UTC](https://discuss.pytorch.org/t/adam-half-precision-nans/1765/2 "2017-04-09T22:08:42Z")

</div>

half precision is super finicky during training, so I’m not surprised.

One thing I recommend trying is to do the forward + backward in half precision, but the optimizer step in float precision.  
To do this, you might have to clone your parameters, and cast them to float32 and once forward+backward is over, you copy over the param .data and .grad into this float32 copy (and call optimizer.step on this float32 copy) and then copy back…

Other than that, I dont have a good idea of why adam + half is giving NaNs.

---

<div class="post-metadata">

**Author:** ![ajbrock](https://discuss.pytorch.org/user_avatar/discuss.pytorch.org/ajbrock/32/139_2.png) [@ajbrock](https://discuss.pytorch.org/u/ajbrock)\
**Post date:** [April 11, 2017, 8:43pm UTC](https://discuss.pytorch.org/t/adam-half-precision-nans/1765/3 "2017-04-11T20:43:22Z")

</div>

Thanks, that’s just the answer I was looking for–will try out the precision swaps and report back.

---

<div class="post-metadata">

**Author:** ![ajbrock](https://discuss.pytorch.org/user_avatar/discuss.pytorch.org/ajbrock/32/139_2.png) [@ajbrock](https://discuss.pytorch.org/u/ajbrock)\
**Post date:** [April 29, 2017, 2:08pm UTC](https://discuss.pytorch.org/t/adam-half-precision-nans/1765/4 "2017-04-29T14:08:15Z")

</div>

EDIT: see below, it looks like eps was the culprit after all, no need for this solution.

Alright, got this working by just hanging onto fp32 copies of the parameters and keeping all of the Adam values in fp32 as well, as shown in [This Gist](https://gist.github.com/ajbrock/075c0ca4036dc4d8581990a6e76e07a3). I suspect you could get the desired stability and do this even more efficiently by just keeping the Adam values in fp32 (I think there’s a divide-by-0 happening somewhere) but this gives the desired memory reduction without any loss of speed over fp32.

---

<div class="post-metadata">

**Author:** ![apaszke](https://discuss.pytorch.org/user_avatar/discuss.pytorch.org/apaszke/32/21_2.png) [@apaszke](https://discuss.pytorch.org/u/apaszke)\
**Post date:** [April 29, 2017, 2:38pm UTC](https://discuss.pytorch.org/t/adam-half-precision-nans/1765/5 "2017-04-29T14:38:26Z")

</div>

It’s probably a 0 division somewhere. Have you tried using a much larger eps (say 1e-4)? The default 1e-8 is rounded to 0 in half precision.

---

<div class="post-metadata">

**Author:** ![ajbrock](https://discuss.pytorch.org/user_avatar/discuss.pytorch.org/ajbrock/32/139_2.png) [@ajbrock](https://discuss.pytorch.org/u/ajbrock)\
**Post date:** [April 29, 2017, 3:00pm UTC](https://discuss.pytorch.org/t/adam-half-precision-nans/1765/6 "2017-04-29T15:00:42Z")

</div>

I had previously tried upping epsilon after tagging it as the culprit, but I can’t recall to exactly what values–as of right now I’m training with eps=1e-4 and it’s working just fine. Guess I should have dug into that further, thanks!

---

<div class="post-metadata">

**Author:** ![michaelklachko](https://discuss.pytorch.org/user_avatar/discuss.pytorch.org/michaelklachko/32/541_2.png) [@michaelklachko](https://discuss.pytorch.org/u/michaelklachko)\
**Post date:** [April 30, 2017, 5:28pm UTC](https://discuss.pytorch.org/t/adam-half-precision-nans/1765/7 "2017-04-30T17:28:20Z")

</div>

You might want to look at this paper:  
[https://arxiv.org/abs/1609.07061](https://arxiv.org/abs/1609.07061)  
if you’re willing to keep a copy of weights/gradients in FP32 you might be able to reduce the precision of forward/backward step much further than FP16.

---

<div class="post-metadata">

**Author:** ![iidsample](https://discuss.pytorch.org/letter_avatar_proxy/v4/letter/i/3d9bf3/32.png) [@iidsample](https://discuss.pytorch.org/u/iidsample)\
**Post date:** [December 15, 2018, 9:17pm UTC](https://discuss.pytorch.org/t/adam-half-precision-nans/1765/8 "2018-12-15T21:17:00Z")

</div>

For anybody who arrives here through a google search -  
This is a paper by Nvidia which sheds more light on training in FP16.  
[https://arxiv.org/abs/1710.03740](https://arxiv.org/abs/1710.03740)

Also they have been nice to provide implemented code -

> **[NVIDIA/apex](https://github.com/NVIDIA/apex)**
>
> A PyTorch Extension: Tools for easy mixed precision and distributed training in Pytorch - NVIDIA/apex

---

<div class="post-metadata">

**Author:** ![Lin\_Jia](https://discuss.pytorch.org/user_avatar/discuss.pytorch.org/lin_jia/32/22799_2.png) [@Lin\_Jia](https://discuss.pytorch.org/u/Lin_Jia)\
**Post date:** [October 14, 2020, 3:50am UTC](https://discuss.pytorch.org/t/adam-half-precision-nans/1765/9 "2020-10-14T03:50:58Z")

</div>

This is a very old post, and google search took me here. For future references, it seems Adam is now adapted to do half-precision training with tuning on hyperparameters:

> **[pytorch1.1 half-precision training Adam RMSprop optimizer Nan problem -...](https://www.programmersought.com/article/22245145270/)**

---

<div class="post-metadata">

**Author:** ![alexmath](https://discuss.pytorch.org/user_avatar/discuss.pytorch.org/alexmath/32/16393_2.png) [@alexmath](https://discuss.pytorch.org/u/alexmath)\
**Post date:** [December 1, 2020, 5:39pm UTC](https://discuss.pytorch.org/t/adam-half-precision-nans/1765/10 "2020-12-01T17:39:40Z")

</div>

I had a similar issue. Changing BCELoss to BCEWithLogitsLoss and Adam epsilon from 10\*\*(-8) to 10\*\*(-4) worked for me.

I found it useful to inspect the computed gradients when debugging.

```auto
print([p.grad for p in model.parameters()])

```

---

<div class="post-metadata">

**Author:** ![gingsi](https://discuss.pytorch.org/user_avatar/discuss.pytorch.org/gingsi/32/34033_2.png) [@gingsi](https://discuss.pytorch.org/u/gingsi)\
**Post date:** [January 30, 2021, 5:54pm UTC](https://discuss.pytorch.org/t/adam-half-precision-nans/1765/11 "2021-01-30T17:54:16Z")

</div>

For me the backward function died before the optimizer was even called. My loss was very high and dividing the loss by 1e5 as last step before the backward pass helped to get rid of the NaN. Of course, performance could suffer, learning rate may need to be tuned etc.

---

<div class="post-metadata">

**Author:** ![m4ttr4ymond](https://discuss.pytorch.org/user_avatar/discuss.pytorch.org/m4ttr4ymond/32/40764_2.png) [@m4ttr4ymond](https://discuss.pytorch.org/u/m4ttr4ymond)\
**Post date:** [July 29, 2021, 7:31pm UTC](https://discuss.pytorch.org/t/adam-half-precision-nans/1765/12 "2021-07-29T19:31:11Z")

</div>

I had this problem with TensorFlow, but this is the only post I found discussing it, so I might as well share my solution. Float16s can only represent numbers as small as 10e-5, but the default adam epsilons for TensorFlow and PyTorch are lower than this (10e-7 and 10e-8 respectively). This seems to cause underflow errors when using float16s. Changing the epsilon to 10e-4 solved the problem for me.

---

<div class="post-metadata">

**Author:** ![Zeratul](https://discuss.pytorch.org/user_avatar/discuss.pytorch.org/zeratul/32/44616_2.png) [@Zeratul](https://discuss.pytorch.org/u/Zeratul)\
**Post date:** [August 29, 2023, 6:40am UTC](https://discuss.pytorch.org/t/adam-half-precision-nans/1765/13 "2023-08-29T06:40:13Z")

</div>

Wow, thank you so much !!🙂 Your solution works really well !!!

---

<div class="post-metadata">

**Author:** ![Mark\_Hamazaspyan](https://discuss.pytorch.org/user_avatar/discuss.pytorch.org/mark_hamazaspyan/32/65173_2.png) [@Mark\_Hamazaspyan](https://discuss.pytorch.org/u/Mark_Hamazaspyan)\
**Post date:** [November 2, 2023, 6:48am UTC](https://discuss.pytorch.org/t/adam-half-precision-nans/1765/14 "2023-11-02T06:48:12Z")

</div>

Another solution, use bitsandbytes. It has Adam8bit optimizer.

```auto
import bitsandbytes as bnb

# adam = torch.optim.Adam(model.parameters(), lr=0.001, betas=(0.9, 0.995)) # comment out old optimizer
adam = bnb.optim.Adam8bit(model.parameters(), lr=0.001, betas=(0.9, 0.995)) # add bnb optimizer

```

> **[GitHub - TimDettmers/bitsandbytes: 8-bit CUDA functions for PyTorch](https://github.com/TimDettmers/bitsandbytes)**
>
> 8-bit CUDA functions for PyTorch. Contribute to TimDettmers/bitsandbytes development by creating an account on GitHub.

---

<div class="post-metadata">

**Author:** ![Nth1](https://discuss.pytorch.org/user_avatar/discuss.pytorch.org/nth1/32/63959_2.png) [@Nth1](https://discuss.pytorch.org/u/Nth1)\
**Post date:** [August 10, 2024, 8:26pm UTC](https://discuss.pytorch.org/t/adam-half-precision-nans/1765/15 "2024-08-10T20:26:34Z")

</div>

Thank you very much 🙂

---

<div class="post-metadata">

**Author:** ![rianhuimo](https://discuss.pytorch.org/user_avatar/discuss.pytorch.org/rianhuimo/32/79521_2.png) [@rianhuimo](https://discuss.pytorch.org/u/rianhuimo)\
**Post date:** [April 11, 2026, 10:51am UTC](https://discuss.pytorch.org/t/adam-half-precision-nans/1765/16 "2026-04-11T10:51:53Z")

</div>

oh my god this worked. thank u so much
