# Training layers of a Flux model separately

**URL:** <https://discourse.julialang.org/t/training-layers-of-a-flux-model-separately/71342>\
**Category:** Machine Learning\
**Tags:** question\
**Created:** [November 11, 2021, 6:07pm UTC](https://discourse.julialang.org/t/training-layers-of-a-flux-model-separately/71342 "2021-11-11T18:07:50Z")\
**Posts on this page:** 2\
**Page:** 1

<div class="post-metadata">

**Author:** ![RedDocMD](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/reddocmd/32/30099_2.png) [@RedDocMD](https://discourse.julialang.org/u/RedDocMD)\
**Post date:** [November 11, 2021, 6:07pm UTC](https://discourse.julialang.org/t/training-layers-of-a-flux-model-separately/71342/1 "2021-11-11T18:07:50Z")

</div>

Suppose I have a model in Flux, which is a bunch of layers that are stacked using `Flux.chain`. The training algorithm that I have trains layers individually.  
So would the following (simplified) code work?

```julia
layers = [Vector of Layers]
model = Flux.chain(layers)
for layer in layers
    gs = Flux.gradient(Flux.params(layer)) do
             some_loss_fn(some_inp, some_out)
    end
    Flux.update!(optimizer, Flux.params(layer), gs)
end
```

---

<div class="post-metadata">

**Author:** ![Rasmus\_Hoier](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/rasmus_hoier/32/24036_2.png) [@Rasmus\_Hoier](https://discourse.julialang.org/u/Rasmus_Hoier)\
**Post date:** [November 13, 2021, 2:47pm UTC](https://discourse.julialang.org/t/training-layers-of-a-flux-model-separately/71342/2 "2021-11-13T14:47:56Z")

</div>

I think that is pretty close 🙂 But you probably want to work with the model rather than the array of layers. I modified your code to make a working example. Since I don’t know how you perform inference or what your layerwise losses look like I just assumed some dummy targets and losses.

```julia
using Flux

nodes = [16 14 13 15 4]
layers = [Dense(nodes[i], nodes[i+1]) for i=1:length(nodes)-1]
model = Chain(layers...) #splat the layers
opt = Descent(0.01)
batchsize = 64

#= Generate a batch of dummy data. z[0] is the input z[i>0] are some layerwise targets =#
z = [rand(Float32, n, batchsize) for n in nodes]

myloss(x, y) = sum(abs2, x-y)

function train(model, z, opt, loss)
    for (i, layer) in enumerate(model)
        gs = Flux.gradient(Flux.params(layer)) do
            y = layer(z[i])    
            loss(y, z[i+1])
        end
        Flux.update!(opt, Flux.params(layer), gs)
    end
end

train(model, z, opt, myloss)

```
