# Question about writting custom training function using Flux.jl

**URL:** <https://discourse.julialang.org/t/question-about-writting-custom-training-function-using-flux-jl/43743>\
**Category:** Machine Learning\
**Tags:** question\
**Created:** [July 27, 2020, 5:26am UTC](https://discourse.julialang.org/t/question-about-writting-custom-training-function-using-flux-jl/43743 "2020-07-27T05:26:01Z")\
**Posts on this page:** 3\
**Page:** 1

<div class="post-metadata">

**Author:** ![XiaodongMa-MRI](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/xiaodongma-mri/32/16171_2.png) [@XiaodongMa-MRI](https://discourse.julialang.org/u/XiaodongMa-MRI)\
**Post date:** [July 27, 2020, 5:26am UTC](https://discourse.julialang.org/t/question-about-writting-custom-training-function-using-flux-jl/43743/1 "2020-07-27T05:26:01Z")

</div>

Hello,

I am trying to write a custom training function instead of using `Flux!train`, following the documentation [https://fluxml.ai/Flux.jl/stable/training/training/#Model-parameters-1](https://fluxml.ai/Flux.jl/stable/training/training/#Model-parameters-1), but it will pop up error message “ **Only reference types can be differentiated with `Params`** ”.

Here is my code (modified based on trebuchet example in model-zoo):

```julia
using Flux
using Zygote
using Statistics
using Random

function shoot( angle, weight)
  angle/weight*10
end

Random.seed!(0)

model = Chain(Dense(1, 16, σ),
              Dense(16, 64, σ),
              Dense(64, 16, σ),
              Dense(16, 2)) |> f64

θ = params(model)

function loss( target)
    angle, weight = model([target]) 
    angle = σ(angle)*90
    weight = weight + 200
    (shoot( angle, weight) - target)^2
end

DIST = (20, 100)	# Maximum target distance

target() = (rand()*(DIST[2]-DIST[1])+DIST[1])

meanloss() = mean(sqrt(loss(target())) for i = 1:100)

opt = ADAM()
dataset = (target() for i = 1:2000)

@time Flux.train!(loss, θ, dataset, opt, cb = () -> println("meanloss = ",meanloss(),"; W1 = ",θ.order.data[1][1],"; b1 = ",θ.order.data[2][1]))

```

It can work by now. Then I wrote a custom training function and run the training:

```julia

function my_custom_train!(loss, ps, data, opt)
  local training_loss
  ps = Params(ps)
  for d in data
    gs = gradient(ps) do
      training_loss = loss(d...)
      return training_loss
    end
    println("training_loss = ",training_loss,"gradient[1] = ",gs[1],"; W1 = ",ps.order.data[1][1],"; b1 = ",ps.order.data[2][1])
    Flux.update!(opt, ps, gs)
  end
end

@time my_custom_train!(loss, θ, dataset, opt)

```

I got the following error:

```julia
ERROR: Only reference types can be differentiated with `Params`.
Stacktrace:
 [1] error(::String) at .\error.jl:33
 [2] getindex(::Zygote.Grads, ::Int64) at C:\Users\maxiao\.julia\packages\Zygote\YeCEW\src\compiler\interface.jl:142
 [3] my_custom_train!(::typeof(loss), ::Params, ::Base.Generator{UnitRange{Int64},var"#13#14"}, ::ADAM) at .\REPL[33]:21
 [4] top-level scope at .\util.jl:175

```

Do anyone have some clue for this? Thanks a lot!

---

<div class="post-metadata">

**Author:** ![contradict](https://avatars.discourse-cdn.com/v4/letter/c/ac91a4/32.png) [@contradict](https://discourse.julialang.org/u/contradict)\
**Post date:** [July 27, 2020, 10:24pm UTC](https://discourse.julialang.org/t/question-about-writting-custom-training-function-using-flux-jl/43743/2 "2020-07-27T22:24:17Z")

</div>

I think the problem is in your logging statement, `gs[1]` is asking for the gradient with respect to 1, you need to request the gradient with respect to a parameter, for example `gs[ps[1]]`.

---

<div class="post-metadata">

**Author:** ![XiaodongMa-MRI](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/xiaodongma-mri/32/16171_2.png) [@XiaodongMa-MRI](https://discourse.julialang.org/u/XiaodongMa-MRI)\
**Post date:** [July 28, 2020, 1:27am UTC](https://discourse.julialang.org/t/question-about-writting-custom-training-function-using-flux-jl/43743/3 "2020-07-28T01:27:02Z")

</div>

Thank you! Changing `gs[1]` to `gs[ps[1]]` works!
