# A possible way to improve training in Flux?

**URL:** https://discourse.julialang.org/t/a-possible-way-to-improve-training-in-flux/44629
**Category:** Performance
**Tags:** flux, precision
**Created:** [August 9, 2020, 4:51pm UTC](https://discourse.julialang.org/t/a-possible-way-to-improve-training-in-flux/44629 "2020-08-09T16:51:17Z")
**Posts on this page:** 16
**Page:** 1

<div class="post-metadata">

### Author: ![SambitMishra98](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/sambitmishra98/32/14263_2.png) [@SambitMishra98](https://discourse.julialang.org/u/SambitMishra98)
#### Post date: [August 9, 2020, 4:51pm UTC](https://discourse.julialang.org/t/a-possible-way-to-improve-training-in-flux/44629/1 "2020-08-09T16:51:17Z")

</div>

Let’s say I need my weights to be determined with an accuracy of `10^10` and hence want to train my network using `Float64`. Does it make sense to train my network with `Float32` and then (after I get around 10^-7 accuracy) shift to `Float64`?

Also, I have a feeling that this idea of first training with smaller precision numbers (even `Float16` ) and then moving on to higher precision numbers will speed up the training process a lot. Is there a reason why it isn’t inherently coded yet?

---

<div class="post-metadata">

### Author: ![ToucheSir](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/touchesir/32/14411_2.png) [@ToucheSir](https://discourse.julialang.org/u/ToucheSir)
#### Post date: [August 9, 2020, 5:43pm UTC](https://discourse.julialang.org/t/a-possible-way-to-improve-training-in-flux/44629/2 "2020-08-09T17:43:45Z")

</div>

I can’t comment on your first paragraph, but the truth is most DL applications just don’t need the level of precision afforded by 64-bit floats. Normalization and other forms of regularization generally encourage smaller magnitude weights that can more effectively use the limited precision of float32 or even float16. Models that are sensitive to small weight perturbations are also more likely to be susceptible to adversarial attacks, while training with noisy data can improve network generalization.

With regards to speeding up training, [mixed-precision](https://arxiv.org/abs/1710.03740) has been gaining traction of late (see e.g. [torch.cuda.amp](https://pytorch.org/docs/stable/amp.html)).

---

<div class="post-metadata">

### Author: ![anon92994695](https://avatars.discourse-cdn.com/v4/letter/a/ce7236/32.png) [@anon92994695](https://discourse.julialang.org/u/anon92994695)
#### Post date: [August 9, 2020, 5:53pm UTC](https://discourse.julialang.org/t/a-possible-way-to-improve-training-in-flux/44629/3 "2020-08-09T17:53:49Z")

</div>

I’ll second the usage of mixed precision. But this raises a gotchya that I have been burned with a few times. Flux’s default floating point precision I believe is 32bit. So if you throw in a Float64 it doesn’t mean you will get a Float64 all the way through your chain. Someone more involved in the guts of Flux/Zygote can comment more, but a couple versions back I had serious troubles with this until I converted everything to Float32’s. One thing to check is that your input and reference/groundtruth/whatever values match precision at Float64.

Most applications do not actually require 64 bit floats though so usually not an issue, but consider the error propagation through a deep net - not so sure 1e-10 is easy but I’m too lazy to run the numbers. Even chaining a few simple LAPACK ops can make 1e-11 difficult iirc.

Again I’m not a pro, someone else can probably give more actionable insight.

---

<div class="post-metadata">

### Author: ![SambitMishra98](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/sambitmishra98/32/14263_2.png) [@SambitMishra98](https://discourse.julialang.org/u/SambitMishra98)
#### Post date: [August 9, 2020, 6:13pm UTC](https://discourse.julialang.org/t/a-possible-way-to-improve-training-in-flux/44629/4 "2020-08-09T18:13:56Z")

</div>

> [@ToucheSir](#):
>
> the truth is most DL applications just don’t need the level of precision afforded by 64-bit float

You’re right about this, but I seem to really need that extra precision. I am trying to use Flux to solve a PDE with numerical techniques and need this extra precision to compare between schemes. In fact, I even plan to use NVIDIA Quadro for this, which supports double precision calculations.

> [@ToucheSir](#):
>
> training with noisy data can improve network generalization.

Given that the dataset is created by a code using `Float64` values, I have a solution with noise generated due to machine precision errors only.

> [@anon92994695](#):
>
> if you throw in a Float64 it doesn’t mean you will get a Float64 all the way through your chain

I just implemented Flux with `Float64` successfully(and probably fixed the issues you may be referring to here) so I don’t have any issues till now.

As an example to why I need this `Float64` …  
Consider [solving the Laplace Equation using higher order Stencils](https://www.researchgate.net/publication/324222428_Higher-order_Accurate_Two-step_Finite_Difference_Schemes_for_the_Many-dimensional_Wave_Equation/figures?lo=1).

 ![image](https://global.discourse-cdn.com/julialang/original/3X/a/7/a7b424229ff99b443d99b1d64fb3b906da750def.png)

Now if I want to check if my network is able to capture the properties of the stencils along with their precision, I really need `Float64` to do so. For my research work, I seem to really need a precision of at least `10^-9`. I hope this example is clear enough.

---

<div class="post-metadata">

### Author: ![ToucheSir](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/touchesir/32/14411_2.png) [@ToucheSir](https://discourse.julialang.org/u/ToucheSir)
#### Post date: [August 9, 2020, 6:23pm UTC](https://discourse.julialang.org/t/a-possible-way-to-improve-training-in-flux/44629/5 "2020-08-09T18:23:08Z")

</div>

I suspected there was a SciML component at play here! About as far as you can get from an expert in that area, so I’ll let those who actually know their stuff chime in. The original comment still stands as a response to “Is there a reason why it isn’t inherently coded yet?” however.

---

<div class="post-metadata">

### Author: ![anon92994695](https://avatars.discourse-cdn.com/v4/letter/a/ce7236/32.png) [@anon92994695](https://discourse.julialang.org/u/anon92994695)
#### Post date: [August 9, 2020, 6:38pm UTC](https://discourse.julialang.org/t/a-possible-way-to-improve-training-in-flux/44629/6 "2020-08-09T18:38:23Z")

</div>

Same, time to call in the pros… @ChrisRackauckas @JeffreySarnoff @dhairyagandhi96

---

<div class="post-metadata">

### Author: ![DrChainsaw](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/drchainsaw/32/8497_2.png) [@DrChainsaw](https://discourse.julialang.org/u/DrChainsaw)
#### Post date: [August 9, 2020, 7:12pm UTC](https://discourse.julialang.org/t/a-possible-way-to-improve-training-in-flux/44629/7 "2020-08-09T19:12:47Z")

</div>

Not a pro and I can’t comment on whether this would be useful, but doing it is easy in flux. Just use the same [functor](https://github.com/FluxML/Functors.jl) mapping function as one uses for cpu → gpu conversion:

```julia

julia> m = Chain(Dense(3,4), Dense(4,5));

julia> typeof.(params(m))
4-element Array{DataType,1}:
 Array{Float32,2}
 Array{Float32,1}
 Array{Float32,2}
 Array{Float32,1}

julia> cfun(x::AbstractArray) = Float64.(x); 

julia> cfun(x) = x; Noop for stuff which is not arrays (e.g. activation functions)

julia> m64 = Flux.fmap(cfun, m)
Chain(Dense(3, 4), Dense(4, 5))

julia> typeof.(params(m64))
4-element Array{DataType,1}:
 Array{Float64,2}
 Array{Float64,1}
 Array{Float64,2}
 Array{Float64,1}

```

---

<div class="post-metadata">

### Author: ![SambitMishra98](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/sambitmishra98/32/14263_2.png) [@SambitMishra98](https://discourse.julialang.org/u/SambitMishra98)
#### Post date: [August 10, 2020, 12:55am UTC](https://discourse.julialang.org/t/a-possible-way-to-improve-training-in-flux/44629/8 "2020-08-10T00:55:03Z")

</div>

> [@ToucheSir](#):
>
> I suspected there was a SciML component at play here!

I did consider `SciML` based on suggestions of @ChrisRackauckas (in Slack) and found that the approach taken by the package was not the same as what I had in mind. `SciML` tries to train a network (based on say 5-Point stencil for Poisson Equation) to obtain a solution, but I want to train a numerical scheme based on the solution itself (like what are the coefficients used in the 5-Point stencil scheme).

---

<div class="post-metadata">

### Author: ![dhairyagandhi96](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/dhairyagandhi96/32/7589_2.png) [@dhairyagandhi96](https://discourse.julialang.org/u/dhairyagandhi96)
#### Post date: [August 10, 2020, 2:42am UTC](https://discourse.julialang.org/t/a-possible-way-to-improve-training-in-flux/44629/9 "2020-08-10T02:42:28Z")

</div>

Flux also ships with `f64` and `f32` which act like the `gpu` function. You could use that directly as well.

---

<div class="post-metadata">

### Author: ![SambitMishra98](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/sambitmishra98/32/14263_2.png) [@SambitMishra98](https://discourse.julialang.org/u/SambitMishra98)
#### Post date: [August 10, 2020, 2:56am UTC](https://discourse.julialang.org/t/a-possible-way-to-improve-training-in-flux/44629/10 "2020-08-10T02:56:37Z")

</div>

> [@dhairyagandhi96](#):
>
> Flux also ships with `f64` and `f32` which act like the `gpu` function.

I know that it is possible, I have seen many posts which have explained it.  
My question is whether the performance of the training can be improved by first training the network with say `f32` precision and then shift to `f64` precision (or `f16` → `f32` → `f64`). This way (i think) we can get a high speedup in the initial few epochs of the training.

---

<div class="post-metadata">

### Author: ![Oscar\_Smith](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/oscar_smith/32/25343_2.png) [@Oscar\_Smith](https://discourse.julialang.org/u/Oscar_Smith)
#### Post date: [August 10, 2020, 2:58am UTC](https://discourse.julialang.org/t/a-possible-way-to-improve-training-in-flux/44629/11 "2020-08-10T02:58:50Z")

</div>

This will speed up training, but could make it miss some very narrow optima.

---

<div class="post-metadata">

### Author: ![SambitMishra98](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/sambitmishra98/32/14263_2.png) [@SambitMishra98](https://discourse.julialang.org/u/SambitMishra98)
#### Post date: [August 10, 2020, 9:28am UTC](https://discourse.julialang.org/t/a-possible-way-to-improve-training-in-flux/44629/12 "2020-08-10T09:28:17Z")

</div>

> [@Oscar\_Smith](#):
>
> This will speed up training but could make it miss some very narrow optima.

I think that will happen even otherwise, since training step size at the first few epochs is the same for a given network, irrespective of the datatype (assuming all other things same). The effect of `Float64` only shows up when the accuracy of the weights is 10^-7 and beyond.

---

<div class="post-metadata">

### Author: ![JeffreySarnoff](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/jeffreysarnoff/32/1980_2.png) [@JeffreySarnoff](https://discourse.julialang.org/u/JeffreySarnoff)
#### Post date: [August 11, 2020, 1:01pm UTC](https://discourse.julialang.org/t/a-possible-way-to-improve-training-in-flux/44629/13 "2020-08-11T13:01:16Z")

</div>

Time a few good examples of your computation, or possibly a simplified version of it, coding it each way (use BenchmarkTools.jl). Find out. (That’s what I would do.)

---

<div class="post-metadata">

### Author: ![ChrisRackauckas](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/chrisrackauckas/32/77_2.png) [@ChrisRackauckas](https://discourse.julialang.org/u/ChrisRackauckas)
#### Post date: [August 12, 2020, 5:44am UTC](https://discourse.julialang.org/t/a-possible-way-to-improve-training-in-flux/44629/14 "2020-08-12T05:44:18Z")

</div>

Copying from the Slack, PDE stencils are convolutional layers so if you use them directly you’ll get fast GPU goodness so that’s the way to go. For discovering PDEs from data using the discovery of the discretization, it’s Section 2.3 [https://arxiv.org/pdf/2001.04385.pdf](https://arxiv.org/pdf/2001.04385.pdf)

---

<div class="post-metadata">

### Author: ![SambitMishra98](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/sambitmishra98/32/14263_2.png) [@SambitMishra98](https://discourse.julialang.org/u/SambitMishra98)
#### Post date: [August 12, 2020, 7:33am UTC](https://discourse.julialang.org/t/a-possible-way-to-improve-training-in-flux/44629/15 "2020-08-12T07:33:52Z")

</div>

My question was not related to solving PDE using ML tools here (Slack question was different). Here I only explained that part to show why I need `Float64` precision. My main question here is quite general, about improving the performance of training step of the code by progressively switching from `Float16` to `Float32` and then `Float64`, instead of using `Float64` precision from the start of training till the end.

---

<div class="post-metadata">

### Author: ![ChrisRackauckas](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/chrisrackauckas/32/77_2.png) [@ChrisRackauckas](https://discourse.julialang.org/u/ChrisRackauckas)
#### Post date: [August 12, 2020, 7:44am UTC](https://discourse.julialang.org/t/a-possible-way-to-improve-training-in-flux/44629/16 "2020-08-12T07:44:42Z")

</div>

Oh sorry, forgot why I had this tab open 😆. Yes, changing precision in Julia is a fairly trivial `Float32.(x)` so abuse it to make multi-precision algorithms. Nick Higham gives a great talk on it:

[![](https://global.discourse-cdn.com/julialang/original/3X/7/a/7a8ad96b6ee852ff0e3be3b4c5e2d64c53ae3e99.jpeg "26 Apr 2017, Nick Higham, “The Rise of Multiprecision Computations"") ](https://www.youtube.com/watch?v=SnUKb_w5r9s)
