# Save best model in FluxTraining.jl

**URL:** https://discourse.julialang.org/t/save-best-model-in-fluxtraining-jl/114591
**Category:** General Usage
**Tags:** flux, fluxtraining
**Created:** [May 22, 2024, 7:50pm UTC](https://discourse.julialang.org/t/save-best-model-in-fluxtraining-jl/114591 "2024-05-22T19:50:08Z")
**Posts on this page:** 5
**Page:** 1

<div class="post-metadata">

### Author: ![cirobr](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/cirobr/32/219994_2.png) [@cirobr](https://discourse.julialang.org/u/cirobr)
#### Post date: [May 22, 2024, 7:50pm UTC](https://discourse.julialang.org/t/save-best-model-in-fluxtraining-jl/114591/1 "2024-05-22T19:50:08Z")

</div>

Cheers, is there a built-in way with `fit!` to save the model parameters (or model state) after each epoch? What about saving the best model only?

Thanks in advance.

---

<div class="post-metadata">

### Author: ![Yuan-Ru-Lin](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/yuan-ru-lin/32/46068_2.png) [@Yuan-Ru-Lin](https://discourse.julialang.org/u/Yuan-Ru-Lin)
#### Post date: [May 23, 2024, 2:06am UTC](https://discourse.julialang.org/t/save-best-model-in-fluxtraining-jl/114591/2 "2024-05-23T02:06:47Z")

</div>

An example based on [the document](https://fluxml.ai/Flux.jl/stable/saving/#Checkpointing):

```julia
using Flux, JLD2

x = rand32(10,100)
y = rand32(1,100)
m = Chain(Dense(10 => 5, relu), Dense(5 => 2), Dense(2=>1))
opt_state = Flux.setup(Adam(), m)

jld = jldopen("model-checkpoint.jld2", "w")
jld["loss"] = Flux.Losses.mse(m(x),y)
jld["model_state"] = Flux.state(m)

for epoch in 1:10
    loss, grads = Flux.withgradient(m->Flux.Losses.mse(m(x),y), m)
    if loss < jld["loss"]
        delete!(jld, "loss")
        jld["loss"] = loss
        delete!(jld, "model_state")
        jld["model_state"] = Flux.state(m)
        @info "Better model found; overwrote the model checkpoint"
    end
    Flux.update!(opt_state, m, grads[1])
end

close(jld)

```

PS: it would look nicer if [this issue](https://github.com/JuliaIO/JLD2.jl/issues/124) gets implemented.

---

<div class="post-metadata">

### Author: ![cirobr](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/cirobr/32/219994_2.png) [@cirobr](https://discourse.julialang.org/u/cirobr)
#### Post date: [May 23, 2024, 9:46am UTC](https://discourse.julialang.org/t/save-best-model-in-fluxtraining-jl/114591/3 "2024-05-23T09:46:01Z")

</div>

Thanks for prompt reply, and my apologies for not being clear. I did not mean saving with the Flux package, but with the instruction `fit!` from `FluxTraining.jl`.

Thanks again for the prompt help!

---

<div class="post-metadata">

### Author: ![Yuan-Ru-Lin](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/yuan-ru-lin/32/46068_2.png) [@Yuan-Ru-Lin](https://discourse.julialang.org/u/Yuan-Ru-Lin)
#### Post date: [May 23, 2024, 9:47pm UTC](https://discourse.julialang.org/t/save-best-model-in-fluxtraining-jl/114591/4 "2024-05-23T21:47:51Z")

</div>

Taking `EarlyStopping` as [an example](https://github.com/FluxML/FluxTraining.jl/blob/master/src/callbacks/earlystopping.jl), it seems you need to create `struct YourCallBack <: AbstractCallback ... end` and extend `on(::EpochEnd, phase::Phase, cb::YourCallBack, learner)` etc. I suggest just use `Flux.jl` if you need that much degree of flexibility, though.

---

<div class="post-metadata">

### Author: ![cirobr](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/cirobr/32/219994_2.png) [@cirobr](https://discourse.julialang.org/u/cirobr)
#### Post date: [May 24, 2024, 8:54am UTC](https://discourse.julialang.org/t/save-best-model-in-fluxtraining-jl/114591/5 "2024-05-24T08:54:12Z")

</div>

Thanks for feedback. I’ve adopted a solution where all callbacks but log and metrics are explicit within the loop. That also solved an issue where early stopping is not currently exiting gracefully from FluxTraining [https://github.com/FluxML/FluxTraining.jl/issues/159](https://github.com/FluxML/FluxTraining.jl/issues/159)

```julia
trainlearner = Learner(model, lossfn;
                      optimizer=opt,
                      callbacks=[log_cb], # only log callback
)
validlearner = Learner(model, lossfn;
                      callbacks=[metrics, log_cb] # only log and metrics callbacks
)

for epoch in 1:epochs
        epoch!(trainlearner, TrainingPhase(), trainset)
        epoch!(validlearner, ValidationPhase(), validset)
        # all other callbacks added here, such as save best model, etc
end

```
