# Learning rate decay in callback function

**URL:** <https://discourse.julialang.org/t/learning-rate-decay-in-callback-function/107859>\
**Category:** Machine Learning\
**Tags:** question, lux\
**Created:** [December 20, 2023, 1:34pm UTC](https://discourse.julialang.org/t/learning-rate-decay-in-callback-function/107859 "2023-12-20T13:34:23Z")\
**Posts on this page:** 4\
**Page:** 1

<div class="post-metadata">

**Author:** ![KianH](https://avatars.discourse-cdn.com/v4/letter/k/c6cbf5/32.png) [@KianH](https://discourse.julialang.org/u/KianH)\
**Post date:** [December 20, 2023, 1:34pm UTC](https://discourse.julialang.org/t/learning-rate-decay-in-callback-function/107859/1 "2023-12-20T13:34:23Z")

</div>

I was wondering if there is a way to access the learning rate through a callback function when using Lux.jl and Optimization.jl packages. For instance in the line below:

`Optimization.solve(optprob, ADAM(1e-3), callback = callback, progress = true, maxiters = sw)`

I have set the learning rate to `1e-3`, and I would like to access it through the callback function to use learning rate decay. I have found the following relevant resources:

[https://fluxml.ai/Optimisers.jl/dev/#Adjusting-Hyperparameters](https://fluxml.ai/Optimisers.jl/dev/#Adjusting-Hyperparameters)

> [@How to update learning rate during Flux training in a better manner?](https://discourse.julialang.org/t/how-to-update-learning-rate-during-flux-training-in-a-better-manner/53101):
>
> Hi, I am trying to update the training rate during training. I did it with custom training loop like below: using Flux using Flux: @epochs using Flux: Flux.Data.DataLoader M = 10 N = 15 O = 2 X = repeat(1.0:10.0, outer=(1, N)) #input Y = repeat(1.0:2.0, outer=(1, N)) #output data = DataLoader(X,Y, batchsize=5, shuffle=true) dims = [M, O] layers = [Dense(dims[i], dims[i+1]) for i in 1:length(dims)-1] m = Chain(layers...) L(x, y) = Flux.Losses.mse(m(x), y) #cost function ps = Flux.par…

However both resources are intended for training their model using a loop, while the training for my code happens in one line, which is the line of code I pasted above. Therefore I would need to access the learning rate through a callback function. Do you have any suggestions?

---

<div class="post-metadata">

**Author:** ![Vaibhavdixit02](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/vaibhavdixit02/32/2916_2.png) [@Vaibhavdixit02](https://discourse.julialang.org/u/Vaibhavdixit02)\
**Post date:** [December 22, 2023, 8:47pm UTC](https://discourse.julialang.org/t/learning-rate-decay-in-callback-function/107859/2 "2023-12-22T20:47:19Z")

</div>

It should be possible but needs a change in the OptimizationOptimisers wrapper, we currently don’t pass the state as an argument to the callback function but it would be a small change and will then let you do it how it is described in the Flux docs. Can you open an issue in Optimization.jl and I’ll create a PR

---

<div class="post-metadata">

**Author:** ![KianH](https://avatars.discourse-cdn.com/v4/letter/k/c6cbf5/32.png) [@KianH](https://discourse.julialang.org/u/KianH)\
**Post date:** [January 10, 2024, 2:23pm UTC](https://discourse.julialang.org/t/learning-rate-decay-in-callback-function/107859/3 "2024-01-10T14:23:32Z")

</div>

In the documentations of Flux.jl regarding learning rate decay, they create a Flux.setup which is then used to update throughout the training as described in the link below.

[https://fluxml.ai/Flux.jl/stable/training/training/](https://fluxml.ai/Flux.jl/stable/training/training/)

The optimization part of my code is written as the following:

```julia
adtype = Optimization.AutoZygote()

optf = Optimization.OptimizationFunction((x, p) -> loss(x,u0), adtype)

optprob = Optimization.OptimizationProblem(optf, ComponentArray(p))

opt = Adam(1e-2)

res1 = Optimization.solve(optprob, opt, callback = callback)

```

Meaning that I do not have any Flux.setup. Additionally, I am using Lux.jl.

What is equivalent to `opt_state` in the code snippet I have provided above to be able to perform learning rate decay through Lux.adjust!(…)?

---

<div class="post-metadata">

**Author:** ![Vaibhavdixit02](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/vaibhavdixit02/32/2916_2.png) [@Vaibhavdixit02](https://discourse.julialang.org/u/Vaibhavdixit02)\
**Post date:** [January 11, 2024, 4:10pm UTC](https://discourse.julialang.org/t/learning-rate-decay-in-callback-function/107859/4 "2024-01-11T16:10:06Z")

</div>

This works now [Optimization.jl/lib/OptimizationOptimisers/test/runtests.jl at master · SciML/Optimization.jl · GitHub](https://github.com/SciML/Optimization.jl/blob/master/lib/OptimizationOptimisers/test/runtests.jl#L62-L72)
