@@ -8,6 +8,11 @@ using MatrixAlgebraKit: remove_qr_gauge_dependence!, remove_lq_gauge_dependence!
88using Mooncake
99using 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
1217mode = Mooncake. ReverseMode
1318rng = 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
110119end
0 commit comments