# How to make parameters of function within Flux.Chain trainable?

**URL:** <https://discourse.julialang.org/t/how-to-make-parameters-of-function-within-flux-chain-trainable/66983>\
**Category:** Machine Learning\
**Created:** [August 25, 2021, 2:57pm UTC](https://discourse.julialang.org/t/how-to-make-parameters-of-function-within-flux-chain-trainable/66983 "2021-08-25T14:57:06Z")\
**Posts on this page:** 3\
**Page:** 1

<div class="post-metadata">

**Author:** ![DiracFermion1](https://avatars.discourse-cdn.com/v4/letter/d/2acd7d/32.png) [@DiracFermion1](https://discourse.julialang.org/u/DiracFermion1)\
**Post date:** [August 25, 2021, 2:57pm UTC](https://discourse.julialang.org/t/how-to-make-parameters-of-function-within-flux-chain-trainable/66983/1 "2021-08-25T14:57:07Z")

</div>

Hi, I have a `NeuralODE` defined as follows:

```julia
dudt = Chain(x -> Dense(10, 20, tanh)(x[1]) + Dense(10, 20, tanh)(x[2]))
n_ode = NeuralODE(dudt, (0., 1.), Tsit5(), saveat=range(0., 1., length=100), reltol=1e-7, abstol=1e-9)

```

which takes a tuple `x` as input and pass its two components through two branches separately. When I check the trainable parameters with:

```julia
Flux.params(n_ode)

```

I got:

`Params([])`

it looks like parameters for two `Dense` layers are not detected as trainable parameters. I am wondering if there is a way to make them trainable?

Thank you!

---

<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 25, 2021, 3:05pm UTC](https://discourse.julialang.org/t/how-to-make-parameters-of-function-within-flux-chain-trainable/66983/2 "2021-08-25T15:05:06Z")

</div>

Related to your [other question](https://discourse.julialang.org/t/flux-chain-vs-expand-everything-in-a-function/66982/2), Flux layers are _stateful_ and need to be kept around instead of re-created on every forward pass. If you want to run something through two branches without writing a custom layer/function, use [`Parallel`](https://fluxml.ai/Flux.jl/stable/models/layers/#Flux.Parallel):

```julia
dudt = Parallel(+, Dense(10, 20, tanh), Dense(10, 20, tanh))

```

For documentation on how Flux layers work and what it takes to make a model trainable, see [Basics · Flux](https://fluxml.ai/Flux.jl/stable/models/basics/#Building-Layers) and [Advanced Model Building · Flux](https://fluxml.ai/Flux.jl/stable/models/advanced/).

---

<div class="post-metadata">

**Author:** ![DiracFermion1](https://avatars.discourse-cdn.com/v4/letter/d/2acd7d/32.png) [@DiracFermion1](https://discourse.julialang.org/u/DiracFermion1)\
**Post date:** [August 25, 2021, 3:36pm UTC](https://discourse.julialang.org/t/how-to-make-parameters-of-function-within-flux-chain-trainable/66983/3 "2021-08-25T15:36:13Z")

</div>

Thanks for the explanation, this is really helpful!

I created a new struct for the same functionality:

```julia
struct TwoInputsLayer
    layer1
    layer2
    op # operation to aggregate them
end

(m::TwoInputsLayer)(x) = m.op(m.layer1(x[1]), m.layer2(x[2]))
Flux.@functor TwoInputsLayer
tmp_two_inputs = TwoInputsLayer(Dense(10, 20, tanh), Dense(10, 20, tanh), +)
dudt = Chain(tmp_two_inputs)
n_ode = NeuralODE(dudt, (0., 1.), Tsit5(), saveat=range(0., 1., length=100), reltol=1e-7, abstol=1e-9)
Flux.params(n_ode)

```

now it seems to be able to track trainable parameters, just wonder if this definition looks good to you or there is anything else I need to be aware of when training models with this struct?

Thanks!
