# How to move optimiser from gpu to cpu?

**URL:** https://discourse.julialang.org/t/how-to-move-optimiser-from-gpu-to-cpu/69622
**Category:** Machine Learning
**Tags:** flux
**Created:** [October 12, 2021, 1:38pm UTC](https://discourse.julialang.org/t/how-to-move-optimiser-from-gpu-to-cpu/69622 "2021-10-12T13:38:42Z")
**Posts on this page:** 11
**Page:** 1

<div class="post-metadata">

### Author: ![weiqi](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/weiqi/32/22148_2.png) [@weiqi](https://discourse.julialang.org/u/weiqi)
#### Post date: [October 12, 2021, 1:38pm UTC](https://discourse.julialang.org/t/how-to-move-optimiser-from-gpu-to-cpu/69622/1 "2021-10-12T13:38:42Z")

</div>

Currently, cpu(optimiser) won’t move it. For instance, the state still consists of variables in GPU.

---

<div class="post-metadata">

### Author: ![CarloLucibello](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/carlolucibello/32/3278_2.png) [@CarloLucibello](https://discourse.julialang.org/u/CarloLucibello)
#### Post date: [October 12, 2021, 2:40pm UTC](https://discourse.julialang.org/t/how-to-move-optimiser-from-gpu-to-cpu/69622/2 "2021-10-12T14:40:17Z")

</div>

It should work adding `Flux.@functor ADAM` (or whatever optimizer you are using) to your code. You should open an issue in Flux.jl stating your use case to see if it is worth adding this feature

---

<div class="post-metadata">

### Author: ![weiqi](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/weiqi/32/22148_2.png) [@weiqi](https://discourse.julialang.org/u/weiqi)
#### Post date: [October 12, 2021, 4:27pm UTC](https://discourse.julialang.org/t/how-to-move-optimiser-from-gpu-to-cpu/69622/3 "2021-10-12T16:27:47Z")

</div>

Flux.@functor ADAM doesn’t work.  
I thought this is a quite basic functionality as we need to restart the training for large-scale learning. Unless we stay with small problems that can easily finished within couple of hours.  
And without loading previously saved optimizer, the training can not be restarted properly.

---

<div class="post-metadata">

### Author: ![findmyway](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/findmyway/32/4946_2.png) [@findmyway](https://discourse.julialang.org/u/findmyway)
#### Post date: [October 12, 2021, 4:45pm UTC](https://discourse.julialang.org/t/how-to-move-optimiser-from-gpu-to-cpu/69622/4 "2021-10-12T16:45:50Z")

</div>

Like Dhairya answered on [slack](https://julialang.slack.com/archives/C7LFJTXV5/p1633972616204100) have you tried Optimizers.jl ? Check out the [tests](https://github.com/FluxML/Optimisers.jl/blob/master/test/runtests.jl) for usages.

---

<div class="post-metadata">

### Author: ![weiqi](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/weiqi/32/22148_2.png) [@weiqi](https://discourse.julialang.org/u/weiqi)
#### Post date: [October 12, 2021, 5:31pm UTC](https://discourse.julialang.org/t/how-to-move-optimiser-from-gpu-to-cpu/69622/5 "2021-10-12T17:31:26Z")

</div>

I tried but didn’t figure out how to use it with Flux models? The tested examples are not for neural network models.

---

<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: [October 12, 2021, 11:23pm UTC](https://discourse.julialang.org/t/how-to-move-optimiser-from-gpu-to-cpu/69622/6 "2021-10-12T23:23:22Z")

</div>

Getting what you want here might require a bit of extra effort.

Flux’s current optimizers use `IdDict`s to map weights to optimizer state and when you move parameters to and from the gpu you create new copies. Result is that the `IdDict` will not recoginize them as the same weighs and instead you have a memory-leak like (weight-leak?) situation.

Depending on what method you use for storing the models and state you might end up with the same problem here before even moving anything (e.g. that weights in optimizer are no longer the same objects as the weights in the model).

I haven’t followed the development of the new optimizers very carefully, but I suppose both new and current optimizers would require you to manually compare weight values (hoping that there are no duplicates) or use some other way to identify the weights and then remap weigths to optimizer state.

---

<div class="post-metadata">

### Author: ![weiqi](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/weiqi/32/22148_2.png) [@weiqi](https://discourse.julialang.org/u/weiqi)
#### Post date: [October 12, 2021, 11:38pm UTC](https://discourse.julialang.org/t/how-to-move-optimiser-from-gpu-to-cpu/69622/7 "2021-10-12T23:38:37Z")

</div>

@DrChainsaw Wonderful insights. Indeed, manually re-mapping the weights and optimizer state is just too much work. The time spent will enable me to re-implement everything in PyTorch 🙂

To get around the issue, I guess the Flux optimizer has to store the optimizer state in a string key rather than use the CUDA matrix as a key, which is so unreliable. For example

`gpu(cpu(a))` will not be `a` anymore when a is a CUDA array.

---

<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: [October 12, 2021, 11:47pm UTC](https://discourse.julialang.org/t/how-to-move-optimiser-from-gpu-to-cpu/69622/8 "2021-10-12T23:47:09Z")

</div>

> [@DrChainsaw](#):
>
> I haven’t followed the development of the new optimizers very carefully, but I suppose both new and current optimizers would require you to manually compare weight values (hoping that there are no duplicates) or use some other way to identify the weights and then remap weigths to optimizer state.

The new optimizers work off of Zygote’s support for structural gradients. That is, you get a nested (named)tuple back which has the same structure as your model. For those who’ve used JAX-based libraries recently, this may look familiar to you (likewise for `state_dict` in PyTorch). You can try out Optimisers.jl today, and there should be 0 IdDicts stored anywhere when you use it 🙂

---

<div class="post-metadata">

### Author: ![weiqi](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/weiqi/32/22148_2.png) [@weiqi](https://discourse.julialang.org/u/weiqi)
#### Post date: [October 12, 2021, 11:54pm UTC](https://discourse.julialang.org/t/how-to-move-optimiser-from-gpu-to-cpu/69622/9 "2021-10-12T23:54:34Z")

</div>

How to work with Flux models, like Dense, in optimiser.jl? Any example code will be appreciated.

---

<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: [October 13, 2021, 12:37am UTC](https://discourse.julialang.org/t/how-to-move-optimiser-from-gpu-to-cpu/69622/10 "2021-10-13T00:37:11Z")

</div>

Currently we need a bit more internal plumbing (Optimisers.jl is still experimental) to get most Flux layers working OOTB. [https://github.com/FluxML/Optimisers.jl/issues/26](https://github.com/FluxML/Optimisers.jl/issues/26) has a good summary there. In the meantime, you can try something like these (warning: untested!) functions:

```julia
# change opt type and IdDict field name for whatever you're using
function extract_opt_state(opt::ADAM, model)
    func = Flux.Functors.children(model)
    map(func) do child
        if Flux.isleaf(child)
            get(opt.state, child, nothing)
        else
            extract_opt_state(opt, child)
        end
    end
end

function restore_opt_state!(opt::ADAM, model, state)
    func = Flux.Functors.children(model)
    map(func, state) do child, st
        if Flux.isleaf(child) && st !== nothing
            opt.state[child] = st
        else
            restore_opt_state!(opt, child, st)
        end
    end
end

```

---

<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: [December 12, 2021, 6:29pm UTC](https://discourse.julialang.org/t/how-to-move-optimiser-from-gpu-to-cpu/69622/11 "2021-12-12T18:29:18Z")

</div>

So it turns out there is an easier way to go about this, see [Deepcopy Flux Model - #9 by ToucheSir](https://discourse.julialang.org/t/deepcopy-flux-model/72930/9).
