Skip to content

Commit 848a3b4

Browse files
authored
A couple missing zero-derivs and an inplace rrule test for svd_trunc (#448)
* A couple missing zero-derivs and an inplace rrule for svd_trunc * Fixes and add test * Mark init output as zero derivative * rrule liposuction * Enable svd_full and svd_compact tests too
1 parent 20ec628 commit 848a3b4

4 files changed

Lines changed: 45 additions & 74 deletions

File tree

Lines changed: 2 additions & 63 deletions
Original file line numberDiff line numberDiff line change
@@ -1,63 +1,2 @@
1-
for f in (:svd_compact, :svd_full)
2-
f_pullback = Symbol(f, :_pullback)
3-
@eval begin
4-
@is_primitive DefaultCtx ReverseMode Tuple{typeof($f), AbstractTensorMap, MatrixAlgebraKit.AbstractAlgorithm}
5-
function Mooncake.rrule!!(::CoDual{typeof($f)}, A_dA::CoDual{<:AbstractTensorMap}, alg_dalg::CoDual)
6-
A, dA = arrayify(A_dA)
7-
alg = primal(alg_dalg)
8-
9-
USVᴴ = $f(A, primal(alg_dalg))
10-
USVᴴ_dUSVᴴ = Mooncake.zero_fcodual(USVᴴ)
11-
dUSVᴴ = last.(arrayify.(USVᴴ, tangent(USVᴴ_dUSVᴴ)))
12-
13-
function $f_pullback(::NoRData)
14-
MatrixAlgebraKit.svd_pullback!(dA, A, USVᴴ, dUSVᴴ)
15-
MatrixAlgebraKit.zero!.(dUSVᴴ)
16-
return ntuple(Returns(NoRData()), 3)
17-
end
18-
19-
return USVᴴ_dUSVᴴ, $f_pullback
20-
end
21-
end
22-
23-
# mutating version is not guaranteed to actually mutate
24-
# so we can simply use the non-mutating version instead and avoid having to worry about
25-
# storing copies and restoring state
26-
f! = Symbol(f, :!)
27-
f!_pullback = Symbol(f!, :_pullback)
28-
@eval begin
29-
@is_primitive DefaultCtx ReverseMode Tuple{typeof($f!), AbstractTensorMap, Any, MatrixAlgebraKit.AbstractAlgorithm}
30-
Mooncake.rrule!!(::CoDual{typeof($f!)}, A_dA::CoDual{<:AbstractTensorMap}, USVᴴ_dUSVᴴ::CoDual, alg_dalg::CoDual) =
31-
Mooncake.rrule!!(Mooncake.zero_fcodual($f), A_dA, alg_dalg)
32-
end
33-
end
34-
35-
@is_primitive DefaultCtx ReverseMode Tuple{typeof(svd_trunc), AbstractTensorMap, MatrixAlgebraKit.AbstractAlgorithm}
36-
function Mooncake.rrule!!(
37-
::CoDual{typeof(svd_trunc)},
38-
A_dA::CoDual{<:AbstractTensorMap},
39-
alg_dalg::CoDual{<:MatrixAlgebraKit.TruncatedAlgorithm}
40-
)
41-
A, dA = arrayify(A_dA)
42-
alg = primal(alg_dalg)
43-
44-
USVᴴ = svd_compact(A, alg.alg)
45-
USVᴴtrunc, ind = MatrixAlgebraKit.truncate(svd_trunc!, USVᴴ, alg.trunc)
46-
ϵ = MatrixAlgebraKit.truncation_error(diagview(USVᴴ[2]), ind)
47-
48-
USVᴴtrunc_dUSVᴴtrunc = Mooncake.zero_fcodual((USVᴴtrunc..., ϵ))
49-
dUSVᴴtrunc = last.(arrayify.(USVᴴtrunc, Base.front(tangent(USVᴴtrunc_dUSVᴴtrunc))))
50-
51-
function svd_trunc_pullback((_, _, _, dϵ)::Tuple{NoRData, NoRData, NoRData, Real})
52-
abs(dϵ) MatrixAlgebraKit.defaulttol(dϵ) ||
53-
@warn "Gradient for `svd_trunc` ignores non-zero tangents for truncation error"
54-
MatrixAlgebraKit.svd_pullback!(dA, A, USVᴴ, dUSVᴴtrunc, ind)
55-
return ntuple(Returns(NoRData()), 3)
56-
end
57-
58-
return USVᴴtrunc_dUSVᴴtrunc, svd_trunc_pullback
59-
end
60-
61-
@is_primitive DefaultCtx ReverseMode Tuple{typeof(svd_trunc!), AbstractTensorMap, Any, MatrixAlgebraKit.AbstractAlgorithm}
62-
Mooncake.rrule!!(::CoDual{typeof(svd_trunc!)}, A_dA::CoDual{<:AbstractTensorMap}, USVᴴ_dUSVᴴ::CoDual, alg_dalg::CoDual) =
63-
Mooncake.rrule!!(Mooncake.zero_fcodual(svd_trunc), A_dA, alg_dalg)
1+
# needed for the ising bimodule case
2+
@zero_derivative DefaultCtx Tuple{typeof(MatrixAlgebraKit.initialize_output), Any, AbstractTensorMap, MatrixAlgebraKit.AbstractAlgorithm}

ext/TensorKitMooncakeExt/tensoroperations.jl

Lines changed: 17 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -250,3 +250,20 @@ function trace_permute_pullback_ΔA!(
250250
)
251251
return NoRData()
252252
end
253+
254+
@is_primitive(
255+
DefaultCtx,
256+
Tuple{
257+
typeof(TensorKit.scalar),
258+
AbstractTensorMap,
259+
}
260+
)
261+
function Mooncake.rrule!!(::CoDual{typeof(TensorKit.scalar)}, t_dt::CoDual{<:AbstractTensorMap})
262+
t, dt = arrayify(t_dt)
263+
val = scalar(t)
264+
function scalar_pullback(Δval)
265+
first(blocks(dt))[2][1] = Δval
266+
return NoRData(), NoRData()
267+
end
268+
return Mooncake.zero_fcodual(val), scalar_pullback
269+
end

ext/TensorKitMooncakeExt/utility.jl

Lines changed: 6 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -65,6 +65,12 @@ Mooncake.tangent_type(::Type{<:HomSpace}) = Mooncake.NoTangent
6565
@zero_derivative DefaultCtx Tuple{typeof(TensorKit.sectorstructure), Any}
6666
@zero_derivative DefaultCtx Tuple{typeof(TensorKit.degeneracystructure), Any}
6767

68+
@zero_derivative DefaultCtx Tuple{typeof(TensorOperations.tensorstructure), AbstractTensorMap}
69+
@zero_derivative DefaultCtx Tuple{typeof(TensorOperations.tensorstructure), AbstractTensorMap, Int, Bool}
70+
71+
@zero_derivative DefaultCtx Tuple{typeof(TensorOperations.tensorcontract_structure), AbstractTensorMap, Index2Tuple, Bool, AbstractTensorMap, Index2Tuple, Bool, Index2Tuple}
72+
73+
@zero_derivative DefaultCtx Tuple{typeof(TensorKit.has_shared_permute), AbstractTensorMap, Index2Tuple}
6874
@zero_derivative DefaultCtx Tuple{typeof(TensorKit.select), HomSpace, Index2Tuple}
6975
@zero_derivative DefaultCtx Tuple{typeof(TensorKit.flip), HomSpace, Any}
7076
@zero_derivative DefaultCtx Tuple{typeof(TensorKit.permute), HomSpace, Index2Tuple}

test/mooncake/factorizations.jl

Lines changed: 20 additions & 11 deletions
Original file line numberDiff line numberDiff line change
@@ -8,6 +8,11 @@ using MatrixAlgebraKit: remove_qr_gauge_dependence!, remove_lq_gauge_dependence!
88
using Mooncake
99
using Random
1010

11+
function call_and_zero!(f!, A, alg)
12+
F′ = f!(A, alg)
13+
MatrixAlgebraKit.zero!(A)
14+
return F′
15+
end
1116

1217
mode = Mooncake.ReverseMode
1318
rng = Random.default_rng()
@@ -18,7 +23,6 @@ eltypes = (Float64, ComplexF64)
1823
@timedtestset "Mooncake - Factorizations: $(TensorKit.type_repr(sectortype(eltype(V)))) ($T)" for V in spacelist, T in eltypes
1924
atol = default_tol(T)
2025
rtol = default_tol(T)
21-
2226
@timedtestset "QR" begin
2327
A = randn(T, V[1] V[2] V[1] V[2])
2428

@@ -29,8 +33,7 @@ eltypes = (Float64, ComplexF64)
2933
ΔQR = Mooncake.randn_tangent(rng, QR)
3034
remove_qr_gauge_dependence!(ΔQR..., A, QR...)
3135
Mooncake.TestUtils.test_rule(rng, qr_full, A; output_tangent = ΔQR, atol, rtol, mode, is_primitive = false)
32-
# TODO:
33-
# Mooncake.TestUtils.test_rule(rng, qr_null, A; atol, rtol, mode, is_primitive = false)
36+
#Mooncake.TestUtils.test_rule(rng, qr_null, A; atol, rtol, mode, is_primitive = false)
3437

3538
A = randn(T, V[1] V[2] V[3] (V[4] V[5])')
3639

@@ -41,34 +44,31 @@ eltypes = (Float64, ComplexF64)
4144
ΔQR = Mooncake.randn_tangent(rng, QR)
4245
remove_qr_gauge_dependence!(ΔQR..., A, QR...)
4346
Mooncake.TestUtils.test_rule(rng, qr_full, A; output_tangent = ΔQR, atol, rtol, mode, is_primitive = false)
44-
# TODO:
45-
# Mooncake.TestUtils.test_rule(rng, qr_null, A; atol, rtol, mode, is_primitive = false)
47+
#Mooncake.TestUtils.test_rule(rng, qr_null, A; atol, rtol, mode, is_primitive = false)
4648
end
4749

4850
@timedtestset "LQ" begin
4951
A = randn(T, V[1] V[2] V[1] V[2])
5052

5153
Mooncake.TestUtils.test_rule(rng, lq_compact, A; atol, rtol, mode, is_primitive = false)
5254

53-
# qr_full/qr_null requires being careful with gauges
55+
# lq_full/lq_null requires being careful with gauges
5456
LQ = lq_full(A)
5557
ΔLQ = Mooncake.randn_tangent(rng, LQ)
5658
remove_lq_gauge_dependence!(ΔLQ..., A, LQ...)
5759
Mooncake.TestUtils.test_rule(rng, lq_full, A; output_tangent = ΔLQ, atol, rtol, mode, is_primitive = false)
58-
# TODO:
59-
# Mooncake.TestUtils.test_rule(rng, lq_null, A; atol, rtol, mode, is_primitive = false)
60+
#Mooncake.TestUtils.test_rule(rng, lq_null, A; atol, rtol, mode, is_primitive = false)
6061

6162
A = randn(T, V[1] V[2] (V[3] V[4] V[5])')
6263

6364
Mooncake.TestUtils.test_rule(rng, lq_compact, A; atol, rtol, mode, is_primitive = false)
6465

65-
# qr_full/qr_null requires being careful with gauges
66+
# lq_full/lq_null requires being careful with gauges
6667
LQ = lq_full(A)
6768
ΔLQ = Mooncake.randn_tangent(rng, LQ)
6869
remove_lq_gauge_dependence!(ΔLQ..., A, LQ...)
6970
Mooncake.TestUtils.test_rule(rng, lq_full, A; output_tangent = ΔLQ, atol, rtol, mode, is_primitive = false)
70-
# TODO:
71-
# Mooncake.TestUtils.test_rule(rng, lq_null, A; atol, rtol, mode, is_primitive = false)
71+
#Mooncake.TestUtils.test_rule(rng, lq_null, A; atol, rtol, mode, is_primitive = false)
7272
end
7373

7474
@timedtestset "Eigenvalue decomposition" begin
@@ -105,6 +105,15 @@ eltypes = (Float64, ComplexF64)
105105
ΔUSVᴴtrunc = (Mooncake.randn_tangent(rng, Base.front(USVᴴtrunc))..., zero(last(USVᴴtrunc)))
106106
remove_svd_gauge_dependence!(ΔUSVᴴtrunc[1], ΔUSVᴴtrunc[3], Base.front(USVᴴtrunc)...)
107107
Mooncake.TestUtils.test_rule(rng, svd_trunc, t, alg; output_tangent = ΔUSVᴴtrunc, atol, rtol, mode)
108+
109+
V_trunc = spacetype(t)(c => min(size(b)...) ÷ 2 for (c, b) in blocks(t))
110+
trunc = truncspace(V_trunc)
111+
USVᴴ = svd_compact(t)
112+
alg = MatrixAlgebraKit.select_algorithm(svd_trunc, t, nothing; trunc)
113+
USVᴴtrunc = svd_trunc(t, alg)
114+
ΔUSVᴴtrunc = (Mooncake.randn_tangent(rng, Base.front(USVᴴtrunc))..., zero(last(USVᴴtrunc)))
115+
remove_svd_gauge_dependence!(ΔUSVᴴtrunc[1], ΔUSVᴴtrunc[3], Base.front(USVᴴtrunc)...)
116+
Mooncake.TestUtils.test_rule(rng, call_and_zero!, svd_trunc!, t, alg; output_tangent = ΔUSVᴴtrunc, atol, rtol, mode, is_primitive = false)
108117
end
109118
end
110119
end

0 commit comments

Comments
 (0)