# Automatic differentiation of complex matrix fails

**URL:** <https://discourse.julialang.org/t/automatic-differentiation-of-complex-matrix-fails/75733>\
**Category:** General Usage\
**Tags:** zygote\
**Created:** [February 3, 2022, 3:13pm UTC](https://discourse.julialang.org/t/automatic-differentiation-of-complex-matrix-fails/75733 "2022-02-03T15:13:20Z")\
**Posts on this page:** 12\
**Page:** 1

<div class="post-metadata">

**Author:** ![Eriklw](https://avatars.discourse-cdn.com/v4/letter/e/ea666f/32.png) [@Eriklw](https://discourse.julialang.org/u/Eriklw)\
**Post date:** [February 3, 2022, 3:13pm UTC](https://discourse.julialang.org/t/automatic-differentiation-of-complex-matrix-fails/75733/1 "2022-02-03T15:13:21Z")

</div>

While automatic differentiation with Zygote works fine for matrices with element in Float64 it fails for elements in ComplexF64:  
‘’’  
using Zygote  
function svdtest(A)  
U,S,V = svd(A)  
a = S[1]  
return a  
end

A = rand(ComplexF64, 10,10)  
svdtest(A)  
gradient(A → svdtest(A), A)‘’’  
Is there a way around this?

---

<div class="post-metadata">

**Author:** ![goerch](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/goerch/32/29122_2.png) [@goerch](https://discourse.julialang.org/u/goerch)\
**Post date:** [February 3, 2022, 3:49pm UTC](https://discourse.julialang.org/t/automatic-differentiation-of-complex-matrix-fails/75733/2 "2022-02-03T15:49:25Z")

</div>

Here is the error message I see:

```julia
ERROR: LoadError: Can't differentiate foreigncall expression

```

Interesting: it looks like `LAPACK`’s `cgesvd` isn’t referenced in svd.jl.

---

<div class="post-metadata">

**Author:** ![Eriklw](https://avatars.discourse-cdn.com/v4/letter/e/ea666f/32.png) [@Eriklw](https://discourse.julialang.org/u/Eriklw)\
**Post date:** [February 3, 2022, 3:51pm UTC](https://discourse.julialang.org/t/automatic-differentiation-of-complex-matrix-fails/75733/3 "2022-02-03T15:51:36Z")

</div>

I see the same error message. Do you know how to interpret this? Is an SVD for a complex matrix simply not supported?

---

<div class="post-metadata">

**Author:** ![goerch](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/goerch/32/29122_2.png) [@goerch](https://discourse.julialang.org/u/goerch)\
**Post date:** [February 3, 2022, 3:54pm UTC](https://discourse.julialang.org/t/automatic-differentiation-of-complex-matrix-fails/75733/4 "2022-02-03T15:54:55Z")

</div>

Complex SVD seems to be supported (because `svdtest(A)` works), but it maybe is differently implemented than real SVD (because your MWE works for Float64 indeed). And I would have expected to find a reference to `cgesvd`.

Edit: ah, I think I understand: that is maybe due to different storage formats of Julia and BLAS?  
Edit: sorry, my misunderstanding, if I replace `svd` with a direct call to `LAPACK`

```julia
    # U,S,V = svd(A)
    U, S, V = LAPACK.gesvd!('A', 'A', A)

```

I run into the same error. Looks like the SVD for reals is somehow smarter?

Edit: Zygote has special rules for [SVD](https://github.com/JuliaDiff/ChainRules.jl/blob/3590f9421950508a97d5a9dbc207208e331c8b75/src/rulesets/LinearAlgebra/factorization.jl#L212-L276) and they are only implemented for reals IIUC.

---

<div class="post-metadata">

**Author:** ![trahflow](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/trahflow/32/30585_2.png) [@trahflow](https://discourse.julialang.org/u/trahflow)\
**Post date:** [February 3, 2022, 4:19pm UTC](https://discourse.julialang.org/t/automatic-differentiation-of-complex-matrix-fails/75733/5 "2022-02-03T16:19:40Z")

</div>

Isn’t that simply the error you get if no ad-rules are defined (e.g. via ChainRules.jl) for foreign (i.e. non-julia) functions?  
I didn’t check, but possibly the rules are defined for `Real`s only?

---

<div class="post-metadata">

**Author:** ![goerch](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/goerch/32/29122_2.png) [@goerch](https://discourse.julialang.org/u/goerch)\
**Post date:** [February 3, 2022, 4:32pm UTC](https://discourse.julialang.org/t/automatic-differentiation-of-complex-matrix-fails/75733/6 "2022-02-03T16:32:05Z")

</div>

A simple hack seems to work around the problem:

```julia
using LinearAlgebra, Zygote, ChainRules

function ChainRules.rrule(::typeof(svd), X::AbstractMatrix{<:Complex})
    F = svd(X)
    svd_pullback(ȳ) = ChainRules._svd_pullback(ȳ, F)
    return F, svd_pullback
end

function svdtest(A)
    U,S,V = svd(A)
    a = S[1]
    return a
end

A = rand(ComplexF64, 10,10)
svdtest(A)
gradient(A -> svdtest(A), A)

```

@Eriklw would you care to check and file an issue against `ChainRules.jl`?

---

<div class="post-metadata">

**Author:** ![mtfishman](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/mtfishman/32/30755_2.png) [@mtfishman](https://discourse.julialang.org/u/mtfishman)\
**Post date:** [February 3, 2022, 8:15pm UTC](https://discourse.julialang.org/t/automatic-differentiation-of-complex-matrix-fails/75733/7 "2022-02-03T20:15:42Z")

</div>

Indeed, it is currently only defined for Real matrices in ChainRules right now:

> <https://github.com/JuliaDiff/ChainRules.jl/blob/8108a77a96af5d4b0c460aac393e44f8943f3c5e/src/rulesets/LinearAlgebra/factorization.jl#L221-L225>

This package extends it to complex:

> <https://github.com/GiggleLiu/BackwardsLinalg.jl/blob/master/src/svd.jl>

Note that it is more subtle than the rule defined by @goerch, see the references:

[https://giggleliu.github.io/2019/04/02/einsumbp.html](https://giggleliu.github.io/2019/04/02/einsumbp.html)

> **[Automatic Differentiation for Complex Valued SVD](https://arxiv.org/abs/1909.02659)**
>
> In this note, we report the back propagation formula for complex valued singular value decompositions (SVD). This formula is an important ingredient for a complete automatic differentiation(AD) infrastructure in terms of complex numbers, and it is...

Probably it would be good for someone to make a proper PR of the complex AD rule from `BackwardsLinalg.jl` to get it into ChainRules.

---

<div class="post-metadata">

**Author:** ![goerch](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/goerch/32/29122_2.png) [@goerch](https://discourse.julialang.org/u/goerch)\
**Post date:** [February 3, 2022, 9:31pm UTC](https://discourse.julialang.org/t/automatic-differentiation-of-complex-matrix-fails/75733/8 "2022-02-03T21:31:59Z")

</div>

> [@mtfishman](#):
>
> Note that it is more subtle than the rule defined by @goerch

Sorry, I naively didn’t check if this was a recent research question.

Edit: nevertheless I did a quick gradient check

```julia
using FiniteDifferences, LinearAlgebra, Zygote, ChainRules

function ChainRules.rrule(::typeof(svd), X::AbstractMatrix{<:Complex})
    F = svd(X)
    svd_pullback(ȳ) = ChainRules._svd_pullback(ȳ, F)
    return F, svd_pullback
end

function svdtest(A)
    U,S,V = svd(A)
    a = S[1]
    return a
end

A = rand(ComplexF64, 10,10)
svdtest(A)
a1 = gradient(A -> svdtest(A), A)
a2 = grad(central_fdm(5, 1), A -> svdtest(A), A)

norm.(a2 .- a1)

```

yielding

```julia
(4.390263449925908e-12,)

```

Coincidence?

---

<div class="post-metadata">

**Author:** ![Eriklw](https://avatars.discourse-cdn.com/v4/letter/e/ea666f/32.png) [@Eriklw](https://discourse.julialang.org/u/Eriklw)\
**Post date:** [February 4, 2022, 8:56am UTC](https://discourse.julialang.org/t/automatic-differentiation-of-complex-matrix-fails/75733/9 "2022-02-04T08:56:00Z")

</div>

Does the BackwardsLinalg.jl package work for you?

---

<div class="post-metadata">

**Author:** ![goerch](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/goerch/32/29122_2.png) [@goerch](https://discourse.julialang.org/u/goerch)\
**Post date:** [February 4, 2022, 11:30am UTC](https://discourse.julialang.org/t/automatic-differentiation-of-complex-matrix-fails/75733/10 "2022-02-04T11:30:35Z")

</div>

Didn’t get it to work with current `Zygote`, older versions error`d out.

---

<div class="post-metadata">

**Author:** ![mtfishman](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/mtfishman/32/30755_2.png) [@mtfishman](https://discourse.julialang.org/u/mtfishman)\
**Post date:** [February 7, 2022, 1:05am UTC](https://discourse.julialang.org/t/automatic-differentiation-of-complex-matrix-fails/75733/11 "2022-02-07T01:05:08Z")

</div>

Please refer to the end of: [[1909.02659] Automatic Differentiation for Complex Valued SVD](https://arxiv.org/abs/1909.02659)

It says that the part of the complex SVD back propagation formula that is unique to the complex case is zero if your loss function only depends on `S` (in fact it has to depend on both `U` and `V` to test the complex back propagation formula). Perhaps try a loss function that depends on `U` and `V` like the one they suggest in that paper.

---

<div class="post-metadata">

**Author:** ![mtfishman](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/mtfishman/32/30755_2.png) [@mtfishman](https://discourse.julialang.org/u/mtfishman)\
**Post date:** [February 7, 2022, 1:06am UTC](https://discourse.julialang.org/t/automatic-differentiation-of-complex-matrix-fails/75733/12 "2022-02-07T01:06:55Z")

</div>

I haven’t tested it in a long time, so I wouldn’t be surprised. Ideally the rule should be ported to `ChainRules`, but that shouldn’t be difficult to write in terms of `svd_back` in [https://github.com/GiggleLiu/BackwardsLinalg.jl/blob/master/src/svd.jl](https://github.com/GiggleLiu/BackwardsLinalg.jl/blob/master/src/svd.jl).
