# Train flux struct with list of models

**URL:** <https://discourse.julialang.org/t/train-flux-struct-with-list-of-models/96874>\
**Category:** General Usage\
**Tags:** flux\
**Created:** [March 31, 2023, 4:59am UTC](https://discourse.julialang.org/t/train-flux-struct-with-list-of-models/96874 "2023-03-31T04:59:32Z")\
**Posts on this page:** 9\
**Page:** 1

<div class="post-metadata">

**Author:** ![CarlosContrerasQ12](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/carloscontrerasq12/32/43287_2.png) [@CarlosContrerasQ12](https://discourse.julialang.org/u/CarlosContrerasQ12)\
**Post date:** [March 31, 2023, 4:59am UTC](https://discourse.julialang.org/t/train-flux-struct-with-list-of-models/96874/1 "2023-03-31T04:59:32Z")

</div>

Hi, I am trying to train a custom model with a list of Flux models as follows

```julia
using Flux,Random,Optimisers,Statistics

mutable struct GlobalModel
    subnets
    function GlobalModel()
        model1=Chain(Dense(3=>3,bias=false,init=rand))
        model2=Chain(Dense(3=>1,bias=false,init=rand))
        subnets=[model1,model2]
        new(subnets)
    end
end

function call_train(glob::GlobalModel,inputs)
    return glob.subnets[2](glob.subnets[1](inputs[1]))
end

(glob::GlobalModel)(inputs) =call_train(glob,inputs)
Flux.@functor GlobalModel 

function loss(y_terminal,inputs)
    delta = (y_terminal.-sum(inputs[1],dims=1)).^2
    return mean(delta)
end

function train!(glob::GlobalModel)
    #opt = Optimisers.Adam(0.01)
    #opt_state = Optimisers.setup(opt, glob)
    optim = Flux.setup(Flux.Adam(0.001), glob)
    
    my_log = []
    for epoch in 1:2000
        input=[randn((3,5)),randn((2,5))]
        val, grads = Flux.withgradient(glob) do m
        result = m(input)
        loss(result, input)
        end

        if epoch%500==0
            inp=[randn((3,5)),randn((2,5))]
            println("Epoch ",epoch, " losses ",loss(glob(inp),inp))
        end
        Flux.update!(optim, Flux.params(glob), grads)
        println(Flux.params(glob))
    end
end

glob=GlobalModel();
train!(glob)

```

However, is seems that the update! function is not updating the parameters in the glob struct, as the parameters remain the same in all epochs. How can I update the glob struct in every step?

---

<div class="post-metadata">

**Author:** ![mcabbott](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/mcabbott/32/6603_2.png) [@mcabbott](https://discourse.julialang.org/u/mcabbott)\
**Post date:** [March 31, 2023, 5:29am UTC](https://discourse.julialang.org/t/train-flux-struct-with-list-of-models/96874/2 "2023-03-31T05:29:24Z")

</div>

> [@CarlosContrerasQ12](#):
>
> `Flux.@functor GlobalModel (model1,model2,)`

This ought to be an error, as it should only accept field names of the struct. But there’s only one field, you want just `@functor GlobalModel` .

---

<div class="post-metadata">

**Author:** ![CarlosContrerasQ12](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/carloscontrerasq12/32/43287_2.png) [@CarlosContrerasQ12](https://discourse.julialang.org/u/CarlosContrerasQ12)\
**Post date:** [March 31, 2023, 11:24am UTC](https://discourse.julialang.org/t/train-flux-struct-with-list-of-models/96874/3 "2023-03-31T11:24:51Z")

</div>

Yes, you’re right, my bd. However, removing it didnt’ do anything.

---

<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:** [March 31, 2023, 2:14pm UTC](https://discourse.julialang.org/t/train-flux-struct-with-list-of-models/96874/4 "2023-03-31T14:14:34Z")

</div>

> [@CarlosContrerasQ12](#):
>
> ```julia
> Flux.update!(optim, Flux.params(glob), grads)
> println(Flux.params(glob))
> 
> ```

You’re mixing different parameter handling styles here, which is why it doesn’t work. If you use `setup` and pass a model to `(with)gradient`, don’t use `params` (and vice versa). Change this to:

```julia
        Flux.update!(optim, glob, grads)

```

And things should work.

---

<div class="post-metadata">

**Author:** ![CarlosContrerasQ12](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/carloscontrerasq12/32/43287_2.png) [@CarlosContrerasQ12](https://discourse.julialang.org/u/CarlosContrerasQ12)\
**Post date:** [March 31, 2023, 2:22pm UTC](https://discourse.julialang.org/t/train-flux-struct-with-list-of-models/96874/5 "2023-03-31T14:22:38Z")

</div>

Thanks! I tried it before, but when I changed `Flux.params(glob)` to just `glob`, I got an error ` type Tuple has no field subnets`

---

<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:** [March 31, 2023, 3:05pm UTC](https://discourse.julialang.org/t/train-flux-struct-with-list-of-models/96874/6 "2023-03-31T15:05:23Z")

</div>

This is when you look at `grads` to make sure it’s what you’d expect. Note that `grads` should be a tuple (one element for each argument to `withgradient`), so you need to extract the first element to get at the gradients of `glob`.

---

<div class="post-metadata">

**Author:** ![CarlosContrerasQ12](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/carloscontrerasq12/32/43287_2.png) [@CarlosContrerasQ12](https://discourse.julialang.org/u/CarlosContrerasQ12)\
**Post date:** [March 31, 2023, 3:30pm UTC](https://discourse.julialang.org/t/train-flux-struct-with-list-of-models/96874/7 "2023-03-31T15:30:51Z")

</div>

Sorry! I also tried it, but the error was the following

`MethodError: no method matching GlobalModel(::Vector{Chain{Tuple{Dense{typeof(identity), Matrix{Float64}, Bool}}}}) Closest candidates are: GlobalModel() at In[4]:5`

---

<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:** [March 31, 2023, 3:47pm UTC](https://discourse.julialang.org/t/train-flux-struct-with-list-of-models/96874/8 "2023-03-31T15:47:59Z")

</div>

> [@CarlosContrerasQ12](#):
>
> ```julia
> function GlobalModel()
> model1=Chain(Dense(3=>3,bias=false,init=rand))
> model2=Chain(Dense(3=>1,bias=false,init=rand))
> subnets=[model1,model2]
> new(subnets)
> end
> 
> ```

Move this inner constructor outside of the definition of `GlobalModel`. Then instead of `new`, just call `GlobalModel(subnets)`. The general recommendation is to avoid inner constructors unless you need them, and in this case you don’t 🙂

---

<div class="post-metadata">

**Author:** ![CarlosContrerasQ12](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/carloscontrerasq12/32/43287_2.png) [@CarlosContrerasQ12](https://discourse.julialang.org/u/CarlosContrerasQ12)\
**Post date:** [March 31, 2023, 4:00pm UTC](https://discourse.julialang.org/t/train-flux-struct-with-list-of-models/96874/9 "2023-03-31T16:00:33Z")

</div>

You’re my hero 😍
