# How to capture model output with loss in Flux.withgradient

**URL:** <https://discourse.julialang.org/t/how-to-capture-model-output-with-loss-in-flux-withgradient/92766>\
**Category:** Machine Learning\
**Created:** [January 10, 2023, 5:54pm UTC](https://discourse.julialang.org/t/how-to-capture-model-output-with-loss-in-flux-withgradient/92766 "2023-01-10T17:54:05Z")\
**Posts on this page:** 4\
**Page:** 1

<div class="post-metadata">

**Author:** ![reachtarunhere](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/reachtarunhere/32/38358_2.png) [@reachtarunhere](https://discourse.julialang.org/u/reachtarunhere)\
**Post date:** [January 10, 2023, 5:54pm UTC](https://discourse.julialang.org/t/how-to-capture-model-output-with-loss-in-flux-withgradient/92766/1 "2023-01-10T17:54:05Z")

</div>

Here is the sample code for my training loop

```julia
function train()
    @showprogress for i in 1:10
        l, grads = Flux.withgradient(m -> lossfn(m(X), Y), model)
        fmap(model, grads[1]) do p, g
            p .= p .- η .* g
        end       
        println(l)
    end
end

```

I am able to capture the loss while computing the gradients. However I also want to say log my predictions generated in the m(X) call. For example I might want to write them to a file. What is the idiomatic way of doing this?

I can modify lossfn to do it inside there but it doesn’t seem clean. Is there a way to return multiple values with Flux.withgradient?

Thanks!

---

<div class="post-metadata">

**Author:** ![CarloLucibello](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/carlolucibello/32/3278_2.png) [@CarloLucibello](https://discourse.julialang.org/u/CarloLucibello)\
**Post date:** [January 11, 2023, 8:03am UTC](https://discourse.julialang.org/t/how-to-capture-model-output-with-loss-in-flux-withgradient/92766/2 "2023-01-11T08:03:09Z")

</div>

You can use `ignore_derivatives` from ChainRulesCore, which can be accessed also from Zygote:

```julia
using Flux, Zygote

function train()
    preds = [] 
    @showprogress for i in 1:10
        l, grads = Flux.withgradient(model) do m 
            Ŷ = m(X)
            Zygote.ignore_derivatives() do 
                push!(preds, Ŷ)
            end
            lossfn(Ŷ, Y)
        end
        fmap(model, grads[1]) do p, g
            p .= p .- η .* g
        end
        println(l)
    end
end

```

---

<div class="post-metadata">

**Author:** ![reachtarunhere](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/reachtarunhere/32/38358_2.png) [@reachtarunhere](https://discourse.julialang.org/u/reachtarunhere)\
**Post date:** [January 11, 2023, 10:33am UTC](https://discourse.julialang.org/t/how-to-capture-model-output-with-loss-in-flux-withgradient/92766/3 "2023-01-11T10:33:51Z")

</div>

Thanks @CarloLucibello I was indeed looking for something like PyTorch’s detach 😃

---

<div class="post-metadata">

**Author:** ![LucasMSpereira](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/lucasmspereira/32/26393_2.png) [@LucasMSpereira](https://discourse.julialang.org/u/LucasMSpereira)\
**Post date:** [January 13, 2023, 2:54pm UTC](https://discourse.julialang.org/t/how-to-capture-model-output-with-loss-in-flux-withgradient/92766/4 "2023-01-13T14:54:59Z")

</div>

[ValueHistories.jl](https://github.com/JuliaML/ValueHistories.jl) worked very nicely for me. Highly recommend taking a look at the [ecosystem](https://fluxml.ai/Flux.jl/stable/ecosystem/) Flux docs page if you’re creating custom training pipelines and models. I didn’t, and realized months later that a bunch of what I did was already in some package 🥲
