# Flux model on GPU?

**URL:** <https://discourse.julialang.org/t/flux-model-on-gpu/102595>\
**Category:** Machine Learning\
**Tags:** flux\
**Created:** [August 8, 2023, 10:23am UTC](https://discourse.julialang.org/t/flux-model-on-gpu/102595 "2023-08-08T10:23:51Z")\
**Posts on this page:** 5\
**Page:** 1

<div class="post-metadata">

**Author:** ![johnbb](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/johnbb/32/34233_2.png) [@johnbb](https://discourse.julialang.org/u/johnbb)\
**Post date:** [August 8, 2023, 10:23am UTC](https://discourse.julialang.org/t/flux-model-on-gpu/102595/1 "2023-08-08T10:23:51Z")

</div>

How can I determine whether a Flux model is on a GPU (of any kind)? For example, if `m = Dense(1, 1) |> gpu/cpu` should it somehow be based on `typeof(m)` or are there other alternatives? My use case is that I have a (training) function with the model (amongst other) as input and if it is on a GPU I need to copy some other variables in the function to the GPU as well.

---

<div class="post-metadata">

**Author:** ![HenriDeh](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/henrideh/32/8316_2.png) [@HenriDeh](https://discourse.julialang.org/u/HenriDeh)\
**Post date:** [August 8, 2023, 10:31am UTC](https://discourse.julialang.org/t/flux-model-on-gpu/102595/2 "2023-08-08T10:31:25Z")

</div>

You can use the [GPUArrays.device](https://juliagpu.github.io/GPUArrays.jl/stable/#GPUArrays.device-Tuple%7BAbstractArray%7D) function.

---

<div class="post-metadata">

**Author:** ![johnbb](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/johnbb/32/34233_2.png) [@johnbb](https://discourse.julialang.org/u/johnbb)\
**Post date:** [August 8, 2023, 2:42pm UTC](https://discourse.julialang.org/t/flux-model-on-gpu/102595/3 "2023-08-08T14:42:06Z")

</div>

Hmm. How would you apply this? I get

```julia
julia> methods(GPUArrays.device)
# 0 methods for generic function "device" from GPUArrays

```

From the link a variable of type AbstractArray is assumed.

---

<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 8, 2023, 2:44pm UTC](https://discourse.julialang.org/t/flux-model-on-gpu/102595/4 "2023-08-08T14:44:04Z")

</div>

Use the mapping and traversal functions in the [Functors API](https://fluxml.ai/Flux.jl/stable/models/functors/) to apply it over a model. The reason it’s not trivial to determine if an entire model is “on GPU” is that one could construct e.g. a hybrid model that has some params on CPU and others on GPU.

---

<div class="post-metadata">

**Author:** ![HenriDeh](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/henrideh/32/8316_2.png) [@HenriDeh](https://discourse.julialang.org/u/HenriDeh)\
**Post date:** [August 9, 2023, 12:29pm UTC](https://discourse.julialang.org/t/flux-model-on-gpu/102595/5 "2023-08-09T12:29:49Z")

</div>

Yes, a model can live on several devices (multiple GPUs typically). If you know all the parameters of a model are on the same device, then you can do `device(m.layers[1].weight)`, which will tell you the GPU of that specific matrix. Weirdly, I couldn’t make this work with GPUArrays, so I directly used the device function of CUDA, but it does not work on normal matrices.

So I guess to answer your question, you could do `model.layers[1].weight isa GPUArrays.AnyGPUArray`, which will be true for any brand of GPU. If you only care about CUDA, you can do `isa CuArray` and drop the GPUArrays interface.
