# Using Flux to optimize a function of the Singular Values

**URL:** <https://discourse.julialang.org/t/using-flux-to-optimize-a-function-of-the-singular-values/39687>\
**Category:** Optimization (Mathematical)\
**Tags:** differentiation, linearalgebra, flux\
**Created:** [May 18, 2020, 10:53am UTC](https://discourse.julialang.org/t/using-flux-to-optimize-a-function-of-the-singular-values/39687 "2020-05-18T10:53:05Z")\
**Posts on this page:** 7\
**Page:** 1

<div class="post-metadata">

**Author:** ![rajnrao](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/rajnrao/32/20637_2.png) [@rajnrao](https://discourse.julialang.org/u/rajnrao)\
**Post date:** [May 18, 2020, 10:53am UTC](https://discourse.julialang.org/t/using-flux-to-optimize-a-function-of-the-singular-values/39687/1 "2020-05-18T10:53:05Z")

</div>

Update – using latest Flux.

I’d like to optimize a function of the singular values of a matrix. The Flux training loop barfs.

Here’s my code

```julia
using Flux, ForwardDiff, LinearAlgebra, GenericLinearAlgebra
m = 20
X = randn(m,m)
model = Dense(m,m)
loss(X) = sum(GenericLinearAlgebra.svdvals(model(X))) ## nuclear norm
data = [X] 
opt = ADAM()
for i = 1:10
      Flux.train!(loss, params(model), X, opt) ## also tried zip(X)
end

```

This gives a

ERROR: MethodError: no method matching (::Dense{typeof(identity),Array{Float32,2},Array{Float32,1}})(::Float64)  
Closest candidates are:  
Any(::AbstractArray{T,N} where N) where {T\<:Union{Float32, Float64}, W\<:(AbstractArray{T,N} where N)} at C:\Users\rajnr.julia\packages\Flux\Fj3bt\src\layers\basic.jl:133  
Any(::AbstractArray{#s107,N} where N where #s107\<:AbstractFloat) where {T\<:Union{Float32, Float64}, W\<:(AbstractArray{T,N} where N)} at C:\Users\rajnr.julia\packages\Flux\Fj3bt\src\layers\basic.jl:136  
Any(::AbstractArray) at C:\Users\rajnr.julia\packages\Flux\Fj3bt\src\layers\basic.jl:121  
Stacktrace:  
[1] macro expansion at C:\Users\rajnr.julia\packages\Zygote\YeCEW\src\compiler\interface2.jl:0 [inlined]  
[2] \_pullback(::Zygote.Context, ::Dense{typeof(identity),Array{Float32,2},Array{Float32,1}}, ::Float64) at C:\Users\rajnr.julia\packages\Zygote\YeCEW\src\compiler\interface2.jl:7  
[3] loss at .\REPL[65]:1 [inlined]  
[4] \_pullback(::Zygote.Context, ::typeof(loss), ::Float64) at C:\Users\rajnr.julia\packages\Zygote\YeCEW\src\compiler\interface2.jl:0  
[5] adjoint at C:\Users\rajnr.julia\packages\Zygote\YeCEW\src\lib\lib.jl:179 [inlined]  
[6] \_pullback at C:\Users\rajnr.julia\packages\ZygoteRules\6nssF\src\adjoint.jl:47 [inlined]  
[7] #17 at C:\Users\rajnr.julia\packages\Flux\Fj3bt\src\optimise\train.jl:89 [inlined]  
[8] \_pullback(::Zygote.Context, ::Flux.Optimise.var"#17#25"{typeof(loss),Float64}) at C:\Users\rajnr.julia\packages\Zygote\YeCEW\src\compiler\interface2.jl:0  
[9] pullback(::Function, ::Zygote.Params) at C:\Users\rajnr.julia\packages\Zygote\YeCEW\src\compiler\interface.jl:174  
[10] gradient(::Function, ::Zygote.Params) at C:\Users\rajnr.julia\packages\Zygote\YeCEW\src\compiler\interface.jl:54  
[11] macro expansion at C:\Users\rajnr.julia\packages\Flux\Fj3bt\src\optimise\train.jl:88 [inlined]  
[12] macro expansion at C:\Users\rajnr.julia\packages\Juno\f8hj2\src\progress.jl:134 [inlined]  
[13] train!(::typeof(loss), ::Zygote.Params, ::Array{Float64,2}, ::ADAM; cb::Flux.Optimise.var"#18#26") at C:\Users\rajnr.julia\packages\Flux\Fj3bt\src\optimise\train.jl:81  
[14] train!(::Function, ::Zygote.Params, ::Array{Float64,2}, ::ADAM) at C:\Users\rajnr.julia\packages\Flux\Fj3bt\src\optimise\train.jl:79  
[15] top-level scope at .\REPL[66]:2

The AD is able to compute gradients though. For example

```julia
h(W) = ForwardDiff.gradient(W -> loss(W*X),W)

h(randn(20,20))

```

will work. So will

```julia
h(Y) = ForwardDiff.gradient(Y -> loss(Y),Y)
h(randn(20,20))

```

It returns the right answer - it’s the [gradient of the nuclear norm with respect to the matrix](https://math.stackexchange.com/questions/701062/proximal-operator-and-the-derivative-of-the-matrix-nuclear-norm)

How can I close the loop so Flux can optimize to above loss function?

Thanks,

Raj

Thanks @baggepinnen for the tip to update to latest Flux. It fixed the TrackedArrays problem.

---

<div class="post-metadata">

**Author:** ![baggepinnen](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/baggepinnen/32/693_2.png) [@baggepinnen](https://discourse.julialang.org/u/baggepinnen)\
**Post date:** [May 18, 2020, 10:57am UTC](https://discourse.julialang.org/t/using-flux-to-optimize-a-function-of-the-singular-values/39687/2 "2020-05-18T10:57:26Z")

</div>

Welcome to this forum! See this post

> [@Please read: make it easier to help you](https://discourse.julialang.org/t/psa-make-it-easier-to-help-you/14757):
>
> Welcome to the Julia Discourse! We are enthusiastic about helping Julia programmers, both beginner and experienced. This public service announcement (PSA) outlines best practices when asking for help. Following these points makes it easier for us to help you and more likely you’ll get a prompt, useful answer. Keywords are highlighted to make it easier to refer to specific points. Choose a descriptive title that captures the key part of your question, eg “plots with multiple axes” instead of …

to make the code more readable 🙂

---

<div class="post-metadata">

**Author:** ![baggepinnen](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/baggepinnen/32/693_2.png) [@baggepinnen](https://discourse.julialang.org/u/baggepinnen)\
**Post date:** [May 18, 2020, 11:00am UTC](https://discourse.julialang.org/t/using-flux-to-optimize-a-function-of-the-singular-values/39687/3 "2020-05-18T11:00:58Z")

</div>

It appears that you are using a very old version of Flux, later versions do not use `TrackedArrays` anymore. Maybe you could try to update Flux to the latest version?

---

<div class="post-metadata">

**Author:** ![MikeInnes](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/mikeinnes/32/3656_2.png) [@MikeInnes](https://discourse.julialang.org/u/MikeInnes)\
**Post date:** [June 10, 2020, 2:25pm UTC](https://discourse.julialang.org/t/using-flux-to-optimize-a-function-of-the-singular-values/39687/5 "2020-06-10T14:25:46Z")

</div>

Here’s how you can get the gradient working:

```julia
using Flux, Zygote, ForwardDiff
using GenericLinearAlgebra: svdvals

svdvals2(X) = Zygote.forwarddiff(svdvals, X)

m = 20
X = randn(m,m)
model = Dense(m,m)
loss(X) = sum(svdvals2(X)) ## nuclear norm

loss(X)

gradient(loss, X)

```

This plugs in forwarddiff to get the `svdvals` gradient. It’d be nice to have direct support for this, so I opened ChainRules issues for [SVD](https://github.com/JuliaDiff/ChainRules.jl/issues/205) and [svdvals](https://github.com/JuliaDiff/ChainRules.jl/issues/206).

---

<div class="post-metadata">

**Author:** ![rajnrao](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/rajnrao/32/20637_2.png) [@rajnrao](https://discourse.julialang.org/u/rajnrao)\
**Post date:** [June 10, 2020, 6:28pm UTC](https://discourse.julialang.org/t/using-flux-to-optimize-a-function-of-the-singular-values/39687/6 "2020-06-10T18:28:02Z")

</div>

@MikeInnes that works great.

The Zygote.forwarddiff(svdvals,X) doesn’t do what one initially thinks it does (it seems to evaluate svdvals(X) instead of diff(svdvals(X)) ). Perhaps needs a better name – Zygote.evalf(svdvals,X) or Zygote.map(svdvals,X) ? Just a thought 🙂

---

<div class="post-metadata">

**Author:** ![MikeInnes](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/mikeinnes/32/3656_2.png) [@MikeInnes](https://discourse.julialang.org/u/MikeInnes)\
**Post date:** [June 11, 2020, 1:48pm UTC](https://discourse.julialang.org/t/using-flux-to-optimize-a-function-of-the-singular-values/39687/7 "2020-06-11T13:48:46Z")

</div>

`forwarddiff(f, x)` is a little weird because it _is_ the same as `f(x)` for the forward pass – it only affects how gradients are calculated, ie using forward rather than reverse mode.

I’d love to have a name that gets that idea across better, although it might be fundamentally unintuitive until you’re exposed to the idea of functions that (only) affect the backwards pass.

---

<div class="post-metadata">

**Author:** ![rajnrao](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/rajnrao/32/20637_2.png) [@rajnrao](https://discourse.julialang.org/u/rajnrao)\
**Post date:** [June 11, 2020, 3:04pm UTC](https://discourse.julialang.org/t/using-flux-to-optimize-a-function-of-the-singular-values/39687/8 "2020-06-11T15:04:13Z")

</div>

aah I see… that is an idea i haven’t been exposed to but seems like a great name could also expose the user to that idea (if they wanted to know more)

how about `map4autodiff`, `map4forwardpass`, `map4backpass` ?
