# MLJ on GPU?

**URL:** <https://discourse.julialang.org/t/mlj-on-gpu/55047>\
**Category:** General Usage\
**Tags:** question, gpu, mlj\
**Created:** [February 11, 2021, 8:32am UTC](https://discourse.julialang.org/t/mlj-on-gpu/55047 "2021-02-11T08:32:40Z")\
**Posts on this page:** 5\
**Page:** 1

<div class="post-metadata">

**Author:** ![mdsa3d](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/mdsa3d/32/20389_2.png) [@mdsa3d](https://discourse.julialang.org/u/mdsa3d)\
**Post date:** [February 11, 2021, 8:32am UTC](https://discourse.julialang.org/t/mlj-on-gpu/55047/1 "2021-02-11T08:32:40Z")

</div>

I would like to know, how to train `MLJ` model on GPU.

# Example

```julia
using MLJ
X = MLJ.table(rand(100, 10));
y = 2X.x1 - X.x2 + 0.05*rand(100);

@load LinearRegressor pkg=MLJLinearModels verbosity=0;
model = LinearRegressor()

mach = machine(model, X, y);
fit!(mach)

params = fitted_params(mach)
params.coefs # coefficient of the regression with names
params.intercept # intercept

Xnew = MLJ.table(rand(3, 10));
ypred = predict(mach, Xnew)

```

What will be the correct approach to implement `gpu` support on this model?

---

<div class="post-metadata">

**Author:** ![tlienart](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/tlienart/32/7640_2.png) [@tlienart](https://discourse.julialang.org/u/tlienart)\
**Post date:** [February 11, 2021, 8:53am UTC](https://discourse.julialang.org/t/mlj-on-gpu/55047/2 "2021-02-11T08:53:30Z")

</div>

Hmm so I’m the current (lazy) maintainer of MLJLinearModels and there’s indeed no explicit support for GPU. Note that using GPU for regression seems a bit overkill but maybe you have a use case that requires it with giant data or something.

Some of the package relies on IterativeSolvers.jl which does support GPU (eg [Conjugate Gradients · IterativeSolvers.jl](https://julialinearalgebra.github.io/IterativeSolvers.jl/dev/linear_systems/cg/#On-the-GPU)) but it seems to require some care to ensure all vectors are on the GPU; I’ve not tried this; though if someone has a clear idea of what’s required, I can try help expose this.

**Edit** : note that some models that can be called via MLJ do have GPU support (e.g. Flux) but I don’t know whether this “just works” or requires careful data handling by MLJ, @ablaom or @samuel_okon should be able to give clearer explanations on that front

---

<div class="post-metadata">

**Author:** ![ablaom](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/ablaom/32/4889_2.png) [@ablaom](https://discourse.julialang.org/u/ablaom)\
**Post date:** [February 11, 2021, 7:46pm UTC](https://discourse.julialang.org/t/mlj-on-gpu/55047/3 "2021-02-11T19:46:34Z")

</div>

Yes GPU support in MLJ is model-specific. (Meta algorithms, such as hyper-parameter tuning, do support multi-processor and multi-threading but there’s no real sense in supporting GPU for this, I’d say.)

The MLJFlux models support training on a GPU. You present your data as normal and the transfer to the GPU is handled under the hood. You enable the GPU for training by setting the hyperparameter `acceleration=CUDALibs()`.

---

<div class="post-metadata">

**Author:** ![mdsa3d](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/mdsa3d/32/20389_2.png) [@mdsa3d](https://discourse.julialang.org/u/mdsa3d)\
**Post date:** [February 17, 2021, 1:37am UTC](https://discourse.julialang.org/t/mlj-on-gpu/55047/4 "2021-02-17T01:37:25Z")

</div>

> [@tlienart](#):
>
> ta

Thank you @tlienart for the response and explanation. Yeah, I am working with large datasets and was trying to reduce the time. And, I will have a look into `IterativeSolvers.jl` and update on my findings on GPU support.

Thanks for maintaining the package, amazing work !!!

---

<div class="post-metadata">

**Author:** ![mdsa3d](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/mdsa3d/32/20389_2.png) [@mdsa3d](https://discourse.julialang.org/u/mdsa3d)\
**Post date:** [February 17, 2021, 1:46am UTC](https://discourse.julialang.org/t/mlj-on-gpu/55047/5 "2021-02-17T01:46:44Z")

</div>

Thanks @ablaom for the response and clearing the doubts regarding gpu support. I did manage to run my `Linear model` on multithreading that really reduced almost 50% of the time.  
I will look into `MLJFlux` for more understanding on gpu implementation.  
I will try changing my data to gpu compatible i.e. `CuArray` and see if it works, will update my findings!
