# Flux: Custom Training + Logging

**URL:** <https://discourse.julialang.org/t/flux-custom-training-logging/41688>\
**Category:** General Usage\
**Tags:** flux\
**Created:** [June 18, 2020, 6:53pm UTC](https://discourse.julialang.org/t/flux-custom-training-logging/41688 "2020-06-18T18:53:33Z")\
**Posts on this page:** 8\
**Page:** 1

<div class="post-metadata">

**Author:** ![jmurray](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/jmurray/32/15806_2.png) [@jmurray](https://discourse.julialang.org/u/jmurray)\
**Post date:** [June 18, 2020, 6:53pm UTC](https://discourse.julialang.org/t/flux-custom-training-logging/41688/1 "2020-06-18T18:53:33Z")

</div>

In the Flux documentation, they give an [example](https://fluxml.ai/Flux.jl/v0.10/training/training/#Custom-Training-loops-1) of how to do a custom training routine and they indicate where you would place code to do logging (see below). I want to log the loss by initializing `LossLog = Float64[]` outside the function, then calling `push!(LossLog, training_loss)`. I’ve tried both placing this snippet just above the `return training_loss` statement and also placing it before `update!`, but both yield errors saying `Mutating arrays not supported` that appears to be coming from `Zygote`.

```julia
# Unchanged code from documentation. I got errors when I tried 
# adding a `push!` statement to log the loss as described above.
function my_custom_train!(loss, ps, data, opt)
  ps = Params(ps)
  for d in data
    gs = gradient(ps) do
      training_loss = loss(d...)
      # Insert what ever code you want here that needs Training loss, e.g. logging
      return training_loss
    end
    # insert what ever code you want here that needs gradient
    # E.g. logging with TensorBoardLogger.jl as histogram so you can see if it is becoming huge
    update!(opt, ps, gs)
    # Here you might like to check validation set accuracy, and break out to do early stopping
  end
end

```

What is the problem, and what do I need to do to make this work?

Also, I want to understand what the code is doing. I understand (I hope!) the `do` block syntax as described [in the docs](https://docs.julialang.org/en/v1/manual/functions/#Do-Block-Syntax-for-Function-Arguments-1). But I’m not sure how the assignment `gs =` factors in. ~~Naively, I’d guess that the code is applying the function `loss(d...)` to each element of the collection `gradient(ps)` and the output after all this is then saved to a variable `gs`.~~

~~However, `gradient(ps)` is, I think, [∇W1, ∇b1, ∇W2, ∇b2, …]. And there has to be a way to get the loss to `update!`; there’s no explicit input of the loss, so it must be included in `gs`. But if that’s true, the `do` block is making a tuple of the gradients and the loss in a manner I’m not familiar with (which isn’t saying much). Point being, I’m confused and could use some help.~~

Edit: it’s of course the gradient of the _loss function_ that needs to be passed to `update!`. I also better understand `do` blocks now and I see that Flux’s example code is assigning to `gs` the output of `gradient(loss(d...), ps)`. And so, for a loss L, I think we have

> gs = [∇W1L, ∇b1L, ∇W2L, ∇b2L, …]

---

<div class="post-metadata">

**Author:** ![jmurray](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/jmurray/32/15806_2.png) [@jmurray](https://discourse.julialang.org/u/jmurray)\
**Post date:** [June 18, 2020, 9:13pm UTC](https://discourse.julialang.org/t/flux-custom-training-logging/41688/2 "2020-06-18T21:13:17Z")

</div>

For clarity, I place a call `evalcb()` after `training_loss = loss(d...)`. I simplified the callback to eliminate the `push!` command and I pre-allocate an array to store values:

```julia
N = 100000
idx = 1
LossLog = Array{Float64, 1}(undef, N)
function evalcb()
    global idx
    global LossLog
    @show typeof(LossLog), size(LossLog) # verifies we are accessing global
    if idx < N
        # FAILS if line below is uncommented.
        LossLog[idx] = 0.1 # dummy value
        
        if true
            idx += 1
            println("Next idx = "*string(idx)*"; LossLog[1]="*string(LossLog[1])) #) idx, LossLog[1] # prints once before failure
        end
    end
end

```

If I comment out the array assignment, it runs: it’s able to increment `idx` as I can see from the print statement. But with the array assignment in, running `my_custom_train` gives the error:

```julia
(typeof(LossLog), size(LossLog)) = (Array{Float64,1}, (100000,))
Next idx = 2; LossLog[1]=0.1

Mutating arrays is not supported

Stacktrace:
 [1] error(::String) at .\error.jl:33
 [2] (::Zygote.var"#1048#1049")(::Nothing) at C:\Users\username\.julia\packages\Zygote\YeCEW\src\lib\array.jl:61
 [3] (::Zygote.var"#2775#back#1050"{Zygote.var"#1048#1049"})(::Nothing) at C:\Users\username\.julia\packages\ZygoteRules\6nssF\src\adjoint.jl:49
 [4] evalcb at .\In[127]:39 [inlined]
 [5] (::typeof(∂(evalcb)))(::Nothing) at C:\Users\username\.julia\packages\Zygote\YeCEW\src\compiler\interface2.jl:0
 [6] #113 at .\In[127]:60 [inlined]
 [7] (::typeof(∂(λ)))(::Float64) at C:\Users\username\.julia\packages\Zygote\YeCEW\src\compiler\interface2.jl:0
 [8] (::Zygote.var"#49#50"{Params,Zygote.Context,typeof(∂(λ))})(::Float64) at C:\Users\username\.julia\packages\Zygote\YeCEW\src\compiler\interface.jl:179
 [9] gradient(::Function, ::Params) at C:\Users\username\.julia\packages\Zygote\YeCEW\src\compiler\interface.jl:55
 [10] my_custom_train!(::Function, ::Params, ::DataLoader, ::ADAM) at .\In[127]:57
 [11] top-level scope at .\In[127]:73

```

The full code to reproduce the error is below. Thanks in advance for any help or insights you can provide!

```julia
using Distributions
using Plots
using Flux
using Flux: param, mse
using Flux.Data: DataLoader
using Zygote
using Zygote: Params

##### DATA #####################
num_samples = 50
x_noise_std = 0.01
y_noise_std = 0.25
function generate_linear_data()
    x = reshape(range(-1, stop=1, length=num_samples), num_samples, 1)
    x_noise = rand(Normal(0,x_noise_std), num_samples)
    y_noise = rand(Normal(0,y_noise_std), num_samples)
    
    y = 3 .* x .+ y_noise
    
    x = transpose(x)
    y = transpose(y)
    
    return x, y
end
X, Y = generate_linear_data() # Training data of shape (1,num_samples)

train_loader = DataLoader(X, Y, batchsize=10, shuffle=true) 

##### CALLBACK #################
N = 100000
idx = 1
LossLog = Array{Float64, 1}(undef, N)
function evalcb()
    global idx
    global LossLog
    @show typeof(LossLog), size(LossLog) # verifies we are accessing global
    if idx < N
        # FAILS if line below is uncommented.
        LossLog[idx] = 0.1 # dummy value
        
        if true
            idx += 1
            println("Next idx = "*string(idx)*"; LossLog[1]="*string(LossLog[1])) #) idx, LossLog[1] # prints once before failure
        end
    end
end

##### MODEL & TRAINING #####################
m = Chain(Dense(size(X, 1), 10, tanh), Dense(10, 10, tanh), Dense(10, size(Y,1), tanh))
opt = ADAM()
loss(x, y) = mse(m(x), y)

# From https://fluxml.ai/Flux.jl/v0.10/training/training/#Custom-Training-loops-1
function my_custom_train!(loss, ps, data, opt)
  ps = Params(ps)
  for d in data
    gs = gradient(ps) do
      training_loss = loss(d...)
      # Insert what ever code you want here that needs Training loss, e.g. logging
      evalcb() # eventually want to pass out training_loss...
      return training_loss
    end

    # insert what ever code you want here that needs gradient
    # E.g. logging with TensorBoardLogger.jl as histogram so you can see if it is becoming huge
    Flux.update!(opt, ps, gs)
    # Here you might like to check validation set accuracy, and break out to do early stopping
  end
end

for epoch in 1:100
    for (x, y) in train_loader
        my_custom_train!(loss, Flux.params(m), train_loader, opt)
    end
end

```

---

<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:** [June 18, 2020, 10:47pm UTC](https://discourse.julialang.org/t/flux-custom-training-logging/41688/3 "2020-06-18T22:47:41Z")

</div>

Zygote has two tools for stopping gradients, I think `ignore` is the one you want.

[`dropgrad`](https://fluxml.ai/Zygote.jl/latest/utils/#Zygote.dropgrad) can be used in an expression to remove the enclosed variable from the backwards pass.

[`ignore`](https://fluxml.ai/Zygote.jl/latest/utils/#Zygote.ignore) can be used to remove a whole block from the backward pass.

`dropgrad` does not seem to apply to assignment, so I think you can use `ignore` like this:

```julia
Zygote.ignore() do
    evalcb(training_loss)
end

```

---

<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:** [June 19, 2020, 4:48pm UTC](https://discourse.julialang.org/t/flux-custom-training-logging/41688/4 "2020-06-19T16:48:44Z")

</div>

I just saw an even better answer. [DiffEqFlux.jl](https://github.com/SciML/DiffEqFlux.jl/blob/master/src/train.jl#L70) uses a similar pattern in its custom training loop, but moves the logging call outside the gradient loop by declaring the loss to be `local`. This works and is a bit easier to understand if you ask me:

```julia
function my_custom_train!(loss, ps, data, opt)
  # declare training loss local so we can use it outside gradient calculation
  local training_loss                                                            
  ps = Params(ps)                                                                   
  for d in data                                                                     
    gs = gradient(ps) do                                                            
      training_loss = loss(d...)
    end                                                                             
    # Insert what ever code you want here that needs Training loss, e.g. logging
    evalcb(training_loss) # eventually want to pass out training_loss...            
                                                                                    
    # insert what ever code you want here that needs gradient                       
    # E.g. logging with TensorBoardLogger.jl as histogram so you can see if it is becoming huge
    Flux.update!(opt, ps, gs)                                                       
    # Here you might like to check validation set accuracy, and break out to do early stopping
  end                                                                               
end                                                                                 

```

---

<div class="post-metadata">

**Author:** ![jmurray](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/jmurray/32/15806_2.png) [@jmurray](https://discourse.julialang.org/u/jmurray)\
**Post date:** [June 19, 2020, 6:00pm UTC](https://discourse.julialang.org/t/flux-custom-training-logging/41688/5 "2020-06-19T18:00:30Z")

</div>

Thanks. That’s actually what I ended up doing late yesterday – just pulling the callback out of the `do` block.

Having read much more on `do` blocks (so I mostly understand them), I now see that Flux’s example code is assigning to `gs` the output of `gradient(loss(d...), ps)`. By placing the logging in the `do` block (as the example comment had indicated), the callback is also passed into `gradient`! That’s not really what we want, and that’s why it might require using Zygote’s `dropgrad` or `ignore` as you indicated in your original reply.

Placing the logging _after_ the `do` block is certainly easiest, and it is probably just as efficient as logging inside the `do` block (if not more so).

---

<div class="post-metadata">

**Author:** ![oxinabox](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/oxinabox/32/206603_2.png) [@oxinabox](https://discourse.julialang.org/u/oxinabox)\
**Post date:** [June 19, 2020, 6:12pm UTC](https://discourse.julialang.org/t/flux-custom-training-logging/41688/6 "2020-06-19T18:12:26Z")

</div>

Would be good to update the docs for this

---

<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:** [June 19, 2020, 6:53pm UTC](https://discourse.julialang.org/t/flux-custom-training-logging/41688/7 "2020-06-19T18:53:03Z")

</div>

How about [this](https://github.com/FluxML/Flux.jl/pull/1240)?

---

<div class="post-metadata">

**Author:** ![jmurray](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/jmurray/32/15806_2.png) [@jmurray](https://discourse.julialang.org/u/jmurray)\
**Post date:** [June 19, 2020, 7:15pm UTC](https://discourse.julialang.org/t/flux-custom-training-logging/41688/8 "2020-06-19T19:15:58Z")

</div>

That’s great! Thanks for doing that.
