# Training a simple linear model in Flux

**URL:** <https://discourse.julialang.org/t/training-a-simple-linear-model-in-flux/24741>\
**Category:** Machine Learning\
**Created:** [May 29, 2019, 4:41pm UTC](https://discourse.julialang.org/t/training-a-simple-linear-model-in-flux/24741 "2019-05-29T16:41:53Z")\
**Posts on this page:** 5\
**Page:** 1

<div class="post-metadata">

**Author:** ![robsmith11](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/robsmith11/32/29641_2.png) [@robsmith11](https://discourse.julialang.org/u/robsmith11)\
**Post date:** [May 29, 2019, 4:41pm UTC](https://discourse.julialang.org/t/training-a-simple-linear-model-in-flux/24741/1 "2019-05-29T16:41:54Z")

</div>

I’m trying to get started with Flux, but finding the documentation a bit lacking in regards to explaining the basic functionality.

For example, how would I train the following linear model with a single scalar parameter `b` using ADAM?

```julia
julia> X = randn(100);

julia> Y = 0.5 .* X .+ randn(100);

julia> loss(xs, ys, b) = sum((ys .- b .* xs).^2)
loss (generic function with 1 method)

julia> Flux.Tracker.update!(Flux.ADAM(), Params(0.0), b -> gradient(b0 -> loss(X,Y,b0), b))
ERROR: MethodError: no method matching getindex(::getfield(Main, Symbol("##183#185")), ::Float64)
Closest candidates are:
  getindex(::Any, ::AbstractTrees.ImplicitRootState) at /home/me/.julia/packages/AbstractTrees/z1wBY/src/AbstractTrees.jl:344
Stacktrace:
 [1] update!(::ADAM, ::Params, ::Function) at /home/me/.julia/packages/Flux/qXNjB/src/optimise/train.jl:11
 [2] top-level scope at REPL[245]:1

```

---

<div class="post-metadata">

**Author:** ![BLI](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/bli/32/37206_2.png) [@BLI](https://discourse.julialang.org/u/BLI)\
**Post date:** [May 29, 2019, 10:35pm UTC](https://discourse.julialang.org/t/training-a-simple-linear-model-in-flux/24741/2 "2019-05-29T22:35:33Z")

</div>

The following should work.

First, introduce some packages…

```plaintext
# Packages
using Flux
using Plots; pyplot()
using LaTeXStrings
using Statistics

```

Next, your data that you wish to fit a model to:

```plaintext
X = rand(100)
Y = 0.5X + rand(100)
plot(X,Y,st=:scatter,label=L"y")
plot!(xlabel=L"x",ylabel=L"y",title="Data for linear model")

```

The data may look as follows:

 ![image](https://global.discourse-cdn.com/julialang/original/3X/f/2/f2314bc83de2e7073117841bad56a5b8a1aedef7.png)

Next, put your data in data arrays of a shape that Flux can use:

```plaintext
# Preparing data in correct data structure
Xd = reduce(hcat,X)
Yd = reduce(hcat,Y)
data = [(Xd,Yd)]

```

Next, set up the Flux model:

```plaintext
# Set up Flux problem
#
# Model 
mod = Dense(1,1)
# Initial mapping
Yd_0 = Tracker.data(mod(Xd))
# Setting up loss/cost function
loss(x, y) = mean((mod(x).-y).^2)
# Selecting parameter optimization method
opt = ADAM(0.01, (0.99, 0.999))
# Extracting parameters from model
par = params(mod);

```

Comments:

- Since you have a monovariable, linear (affine) mapping, a single layer (no hidden layers) with a linear activation function (default) is sufficient.
- `Dense` is the Flux name for the standard Feedforward Neural Net (FNN) block.
- `Yd_0` is the mapping from x to y with the initial (randomly generated) set of model parameters in model `mod`.
- In the last line, I name the parameters by `par` so that I can refer to par in the next code block where I train the “network” (the linear model).

Next, you need to train the model against the data – a major iteration is denoted an “epoch”:

```plaintext
# Training over nE epochs
nE = 1_000
for i in 1:nE
    Flux.train!(loss,par,data,opt)
end
# Final mapping
Yd_nE = Tracker.data(mod(Xd));

```

Here, `Yd_nE` is the mapping from x to y with the model parameters as they are after `nE` epochs:

```plaintext
plot(X,Y,st=:scatter,label=L"y")
plot!(Xd',Yd_0',lc=:green,label=L"y_0")
plot!(Xd',Yd_nE',lc=:red,label=L"y_{n_\mathrm{E}}")
plot!(xlabel=L"x",ylabel=L"y",title="Data for linear model")

```

… and then the result:

 ![image](https://global.discourse-cdn.com/julialang/original/3X/1/4/143cb29e2908b87c2313e409421da0a2e1dcfb0c.png)

Of course, in this case, it would be much simpler and faster to solve the model using linear algebra.

---

<div class="post-metadata">

**Author:** ![alejandromerchan](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/alejandromerchan/32/10500_2.png) [@alejandromerchan](https://discourse.julialang.org/u/alejandromerchan)\
**Post date:** [May 29, 2019, 11:32pm UTC](https://discourse.julialang.org/t/training-a-simple-linear-model-in-flux/24741/3 "2019-05-29T23:32:48Z")

</div>

This should be on a blog or something. Nice and simple!

---

<div class="post-metadata">

**Author:** ![BLI](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/bli/32/37206_2.png) [@BLI](https://discourse.julialang.org/u/BLI)\
**Post date:** [June 1, 2019, 7:36pm UTC](https://discourse.julialang.org/t/training-a-simple-linear-model-in-flux/24741/4 "2019-06-01T19:36:25Z")

</div>

Yes, I think it is a simple introduction. One should just add a nonlinear case + a multi input case, plus a few things more, and that would be a simple way to make people start using Flux.

I’m not really a blogger; I would have to figure out how to set up a blog myself, were I to do it.

---

<div class="post-metadata">

**Author:** ![williamfgc](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/williamfgc/32/15445_2.png) [@williamfgc](https://discourse.julialang.org/u/williamfgc)\
**Post date:** [September 11, 2020, 3:20pm UTC](https://discourse.julialang.org/t/training-a-simple-linear-model-in-flux/24741/5 "2020-09-11T15:20:10Z")

</div>

Since Tracker is not working with Julia 1.4, I just pull the data using `Yd_0 = mod(Xd)` and `Yd_nE = mod(Xd)`. Thanks to @BLI for the nice answer.
