diff --git a/src/layers/invertible_layer_actnorm.jl b/src/layers/invertible_layer_actnorm.jl index c1c348b8..01e6c073 100644 --- a/src/layers/invertible_layer_actnorm.jl +++ b/src/layers/invertible_layer_actnorm.jl @@ -57,8 +57,7 @@ function ActNorm(k; logdet=false) end # 2-3D Foward pass: Input X, Output Y -function forward(X::AbstractArray{T, N}, AN::ActNorm; logdet=nothing) where {T, N} - isnothing(logdet) ? logdet = (AN.logdet && ~AN.is_reversed) : logdet = logdet +function forward(X::AbstractArray{T, N}, AN::ActNorm;) where {T, N} inds = [i!=(N-1) ? 1 : Colon() for i=1:N] dims = collect(1:N-1); dims[end] +=1 @@ -73,12 +72,11 @@ function forward(X::AbstractArray{T, N}, AN::ActNorm; logdet=nothing) where {T, Y = X .* reshape(AN.s.data, inds...) .+ reshape(AN.b.data, inds...) # If logdet true, return as second ouput argument - logdet ? (return Y, logdet_forward(size(X)[1:N-2]..., AN.s)) : (return Y) + AN.logdet ? (return Y, logdet_forward(size(X)[1:N-2]..., AN.s)) : (return Y) end # 2-3D Inverse pass: Input Y, Output X -function inverse(Y::AbstractArray{T, N}, AN::ActNorm; logdet=nothing) where {T, N} - isnothing(logdet) ? logdet = (AN.logdet && AN.is_reversed) : logdet = logdet +function inverse(Y::AbstractArray{T, N}, AN::ActNorm;) where {T, N} inds = [i!=(N-1) ? 1 : Colon() for i=1:N] dims = collect(1:N-1); dims[end] +=1 @@ -93,7 +91,7 @@ function inverse(Y::AbstractArray{T, N}, AN::ActNorm; logdet=nothing) where {T, X = (Y .- reshape(AN.b.data, inds...)) ./ reshape(AN.s.data, inds...) # If logdet true, return as second ouput argument - logdet ? (return X, -logdet_forward(size(Y)[1:N-2]..., AN.s)) : (return X) + AN.logdet ? (return X, -logdet_forward(size(Y)[1:N-2]..., AN.s)) : (return X) end # 2-3D Backward pass: Input (ΔY, Y), Output (ΔY, Y) @@ -102,7 +100,7 @@ function backward(ΔY::AbstractArray{T, N}, Y::AbstractArray{T, N}, AN::ActNorm; dims = collect(1:N-1); dims[end] +=1 nn = size(ΔY)[1:N-2] - X = inverse(Y, AN; logdet=false) + AN.logdet ? (X, logdet_i) = inverse(Y, AN;) : X = inverse(Y, AN;) ΔX = ΔY .* reshape(AN.s.data, inds...) Δs = sum(ΔY .* X, dims=dims)[inds...] if AN.logdet @@ -129,7 +127,7 @@ function backward_inv(ΔX::AbstractArray{T, N}, X::AbstractArray{T, N}, AN::ActN dims = collect(1:N-1); dims[end] +=1 nn = size(ΔX)[1:N-2] - Y = forward(X, AN; logdet=false) + AN.logdet ? (Y, logdet_i) = forward(X, AN;) : Y = forward(X, AN;) ΔY = ΔX ./ reshape(AN.s.data, inds...) Δs = -sum(ΔX .* X ./ reshape(AN.s.data, inds...), dims=dims)[inds...] if AN.logdet @@ -152,20 +150,19 @@ end ## Jacobian-related functions # 2-£D function jacobian(ΔX::AbstractArray{T, N}, Δθ::AbstractArray{Parameter, 1}, X::AbstractArray{T, N}, AN::ActNorm; logdet=nothing) where {T, N} - isnothing(logdet) ? logdet = (AN.logdet && ~AN.is_reversed) : logdet = logdet inds = [i!=(N-1) ? 1 : Colon() for i=1:N] nn = size(ΔX)[1:N-2] Δs = Δθ[1].data Δb = Δθ[2].data # Forward evaluation - logdet ? (Y, lgdet) = forward(X, AN; logdet=logdet) : Y = forward(X, AN; logdet=logdet) + AN.logdet ? (Y, lgdet) = forward(X, AN;) : Y = forward(X, AN;) # Jacobian evaluation ΔY = ΔX .* reshape(AN.s.data, inds...) .+ X .* reshape(Δs, inds...) .+ reshape(Δb, inds...) # Hessian evaluation of logdet terms - if logdet + if AN.logdet nx, ny, _, _ = size(X) HlogΔθ = [Parameter(logdet_hessian(nn..., AN.s).*Δs), Parameter(zeros(Float32, size(Δb)))] return ΔY, Y, lgdet, HlogΔθ diff --git a/src/layers/invertible_layer_glow.jl b/src/layers/invertible_layer_glow.jl index 39d8e1cf..1c37fff8 100644 --- a/src/layers/invertible_layer_glow.jl +++ b/src/layers/invertible_layer_glow.jl @@ -91,20 +91,24 @@ end CouplingLayerGlow3D(args...;kw...) = CouplingLayerGlow(args...; kw..., ndims=3) # Forward pass: Input X, Output Y -function forward(X::AbstractArray{T, N}, L::CouplingLayerGlow) where {T,N} +function forward(X::AbstractArray{T, N}, L::CouplingLayerGlow; save=false) where {T,N} X_ = L.C.forward(X) X1, X2 = tensor_split(X_) - Y2 = copy(X2) logS_T = L.RB.forward(X2) logSm, Tm = tensor_split(logS_T) Sm = L.activation.forward(logSm) Y1 = Sm.*X1 + Tm - Y = tensor_cat(Y1, Y2) + Y = tensor_cat(Y1, X2) + + if L.logdet + save ? (return Y, Y1, X2, Sm, glow_logdet_forward(Sm)) : (return Y, glow_logdet_forward(Sm)) + else + save ? (return Y, Y1, X2, Sm) : (return Y) + end - L.logdet == true ? (return Y, glow_logdet_forward(Sm)) : (return Y) end # Inverse pass: Input Y, Output X @@ -120,7 +124,11 @@ function inverse(Y::AbstractArray{T, N}, L::CouplingLayerGlow; save=false) where X_ = tensor_cat(X1, X2) X = L.C.inverse(X_) - save == true ? (return X, X1, X2, Sm) : (return X) + if L.logdet + save ? (return X, X1, X2, Sm, -glow_logdet_forward(Sm)) : (return X, -glow_logdet_forward(Sm)) + else + save ? (return X, X1, X2, Sm) : (return X) + end end # Backward pass: Input (ΔY, Y), Output (ΔX, X) @@ -160,13 +168,37 @@ function backward(ΔY::AbstractArray{T, N}, Y::AbstractArray{T, N}, L::CouplingL end end +# 2D/3D Reverse backward pass: Input (ΔX, X), Output (ΔY, Y) +function backward_inv(ΔX::AbstractArray{T, N}, X::AbstractArray{T, N}, L::CouplingLayerGlow; set_grad::Bool=true) where {T, N} -## Jacobian-related functions + ΔX, X = L.C.forward((ΔX, X)) + X1, X2 = tensor_split(X) + ΔX1, ΔX2 = tensor_split(ΔX) -function jacobian(ΔX::AbstractArray{T, N}, Δθ::Array{Parameter, 1}, X, L::CouplingLayerGlow) where {T,N} + # Recompute forward state + logS_T = L.RB.forward(X2) + logSm, Tm = tensor_split(logS_T) + Sm = L.activation.forward(logSm) + Y1 = Sm.*X1 + Tm + + # Backpropagate residual + ΔT = -ΔX1 ./ Sm + ΔS = X1 .* ΔT + if L.logdet == true + ΔS += coupling_logdet_backward(Sm) + end + + ΔY2 = L.RB.backward(tensor_cat(L.activation.backward(ΔS, Sm), ΔT), X2) + ΔX2 + ΔY1 = -ΔT + + ΔY = tensor_cat(ΔY1, ΔY2) + Y = tensor_cat(Y1, X2) - # Get dimensions - k = Int(L.C.k/2) + return ΔY, Y +end + +## Jacobian-related functions +function jacobian(ΔX::AbstractArray{T, N}, Δθ::Array{Parameter, 1}, X, L::CouplingLayerGlow) where {T,N} ΔX_, X_ = L.C.jacobian(ΔX, Δθ[1:3], X) X1, X2 = tensor_split(X_) @@ -175,17 +207,19 @@ function jacobian(ΔX::AbstractArray{T, N}, Δθ::Array{Parameter, 1}, X, L::Cou Y2 = copy(X2) ΔY2 = copy(ΔX2) ΔlogS_T, logS_T = L.RB.jacobian(ΔX2, Δθ[4:end], X2) - Sm = L.activation.forward(logS_T[:,:,1:k,:]) - ΔS = L.activation.backward(ΔlogS_T[:,:,1:k,:], nothing;x=logS_T[:,:,1:k,:]) - Tm = logS_T[:, :, k+1:end, :] - ΔT = ΔlogS_T[:, :, k+1:end, :] + logSm, Tm = tensor_split(logS_T) + ΔlogSm, ΔT = tensor_split(ΔlogS_T) + + Sm = L.activation.forward(logSm) + ΔS = L.activation.backward(ΔlogSm, nothing;x=logSm) Y1 = Sm.*X1 + Tm ΔY1 = ΔS.*X1 + Sm.*ΔX1 + ΔT Y = tensor_cat(Y1, Y2) ΔY = tensor_cat(ΔY1, ΔY2) # Gauss-Newton approximation of logdet terms - JΔθ = L.RB.jacobian(cuzeros(ΔX2, size(ΔX2)), Δθ[4:end], X2)[1][:, :, 1:k, :] + JΔθ = L.RB.jacobian(cuzeros(ΔX2, size(ΔX2)), Δθ[4:end], X2)[1] + JΔθ = tensor_split(JΔθ)[1] GNΔθ = cat(0f0*Δθ[1:3], -L.RB.adjointJacobian(tensor_cat(L.activation.backward(JΔθ, Sm), zeros(Float32, size(Sm))), X2)[2]; dims=1) L.logdet ? (return ΔY, Y, glow_logdet_forward(Sm), GNΔθ) : (return ΔY, Y) @@ -195,6 +229,6 @@ function adjointJacobian(ΔY::AbstractArray{T, N}, Y::AbstractArray{T, N}, L::Co return backward(ΔY, Y, L; set_grad=false) end -# Logdet (correct?) +# Logdet glow_logdet_forward(S) = sum(log.(abs.(S))) / size(S)[end] glow_logdet_backward(S) = 1f0./ S / size(S)[end] diff --git a/src/networks/invertible_network_glow.jl b/src/networks/invertible_network_glow.jl index 105b3c75..2583362d 100644 --- a/src/networks/invertible_network_glow.jl +++ b/src/networks/invertible_network_glow.jl @@ -67,12 +67,13 @@ struct NetworkGlow <: InvertibleNetwork K::Int64 squeezer::Squeezer split_scales::Bool + logdet::Bool end @Flux.functor NetworkGlow # Constructor -function NetworkGlow(n_in, n_hidden, L, K; freeze_conv=false, split_scales=false, k1=3, k2=1, p1=1, p2=0, s1=1, s2=1, ndims=2, squeezer::Squeezer=ShuffleLayer(), activation::ActivationFunction=SigmoidLayer()) +function NetworkGlow(n_in, n_hidden, L, K; logdet=true,freeze_conv=false, split_scales=false, k1=3, k2=1, p1=1, p2=0, s1=1, s2=1, ndims=2, squeezer::Squeezer=ShuffleLayer(), activation::ActivationFunction=SigmoidLayer()) AN = Array{ActNorm}(undef, L, K) # activation normalization CL = Array{CouplingLayerGlow}(undef, L, K) # coupling layers w/ 1x1 convolution and residual block @@ -87,13 +88,13 @@ function NetworkGlow(n_in, n_hidden, L, K; freeze_conv=false, split_scales=false for i=1:L n_in *= channel_factor # squeeze if split_scales is turned on for j=1:K - AN[i, j] = ActNorm(n_in; logdet=true) - CL[i, j] = CouplingLayerGlow(n_in, n_hidden; freeze_conv=freeze_conv, k1=k1, k2=k2, p1=p1, p2=p2, s1=s1, s2=s2, logdet=true, activation=activation, ndims=ndims) + AN[i, j] = ActNorm(n_in; logdet=logdet) + CL[i, j] = CouplingLayerGlow(n_in, n_hidden; logdet=logdet, freeze_conv=freeze_conv, k1=k1, k2=k2, p1=p1, p2=p2, s1=s1, s2=s2, activation=activation, ndims=ndims) end (i < L && split_scales) && (n_in = Int64(n_in/2)) # split end - return NetworkGlow(AN, CL, Z_dims, L, K, squeezer, split_scales) + return NetworkGlow(AN, CL, Z_dims, L, K, squeezer, split_scales, logdet) end NetworkGlow3D(args; kw...) = NetworkGlow(args...; kw..., ndims=3) @@ -101,14 +102,14 @@ NetworkGlow3D(args; kw...) = NetworkGlow(args...; kw..., ndims=3) # Forward pass and compute logdet function forward(X::AbstractArray{T, N}, G::NetworkGlow) where {T, N} G.split_scales && (Z_save = array_of_array(X, G.L-1)) - + orig_shape = size(X) logdet = 0 for i=1:G.L (G.split_scales) && (X = G.squeezer.forward(X)) for j=1:G.K - X, logdet1 = G.AN[i, j].forward(X) - X, logdet2 = G.CL[i, j].forward(X) - logdet += (logdet1 + logdet2) + G.logdet ? (X, logdet1) = G.AN[i, j].forward(X) : X = G.AN[i, j].forward(X) + G.logdet ? (X, logdet2) = G.CL[i, j].forward(X) : X = G.CL[i, j].forward(X) + G.logdet && (logdet += (logdet1 + logdet2)) end if G.split_scales && i < G.L # don't split after last iteration X, Z = tensor_split(X) @@ -116,25 +117,27 @@ function forward(X::AbstractArray{T, N}, G::NetworkGlow) where {T, N} G.Z_dims[i] = collect(size(Z)) end end - G.split_scales && (X = cat_states(Z_save, X)) - return X, logdet + G.split_scales && (X = reshape(cat_states(Z_save, X),orig_shape)) + G.logdet ? (return X, logdet) : (return X) end # Inverse pass function inverse(X::AbstractArray{T, N}, G::NetworkGlow) where {T, N} - G.split_scales && ((Z_save, X) = split_states(X, G.Z_dims)) + G.split_scales && ((Z_save, X) = split_states(X[:], G.Z_dims)) + logdet = 0 for i=G.L:-1:1 if G.split_scales && i < G.L X = tensor_cat(X, Z_save[i]) end for j=G.K:-1:1 - X = G.CL[i, j].inverse(X) - X = G.AN[i, j].inverse(X) + G.logdet ? (X, logdet1) = G.CL[i, j].inverse(X) : X = G.CL[i, j].inverse(X) + G.logdet ? (X, logdet2) = G.AN[i, j].inverse(X) : X = G.AN[i, j].inverse(X) + G.logdet && (logdet += (logdet1 + logdet2)) end (G.split_scales) && (X = G.squeezer.inverse(X)) end - return X + G.logdet ? (return X, logdet) : (return X) end # Backward pass and compute gradients @@ -142,8 +145,8 @@ function backward(ΔX::AbstractArray{T, N}, X::AbstractArray{T, N}, G::NetworkGl # Split data and gradients if G.split_scales - ΔZ_save, ΔX = split_states(ΔX, G.Z_dims) - Z_save, X = split_states(X, G.Z_dims) + ΔZ_save, ΔX = split_states(ΔX[:], G.Z_dims) + Z_save, X = split_states(X[:], G.Z_dims) end if ~set_grad @@ -180,14 +183,44 @@ function backward(ΔX::AbstractArray{T, N}, X::AbstractArray{T, N}, G::NetworkGl set_grad ? (return ΔX, X) : (return ΔX, vcat(ΔθAN, ΔθCL), X, vcat(∇logdetAN, ∇logdetCL)) end +# Backward reverse pass and compute gradients +function backward_inv(ΔX::AbstractArray{T, N}, X::AbstractArray{T, N}, G::NetworkGlow) where {T, N} + G.split_scales && (X_save = array_of_array(X, G.L-1)) + G.split_scales && (ΔX_save = array_of_array(ΔX, G.L-1)) + orig_shape = size(X) + + for i=1:G.L + G.split_scales && (ΔX = G.squeezer.forward(ΔX)) + G.split_scales && (X = G.squeezer.forward(X)) + for j=1:G.K + ΔX_, X_ = backward_inv(ΔX, X, G.AN[i, j]) + ΔX, X = backward_inv(ΔX_, X_, G.CL[i, j]) + end + + if G.split_scales && i < G.L # don't split after last iteration + X, Z = tensor_split(X) + ΔX, ΔZx = tensor_split(ΔX) + + X_save[i] = Z + ΔX_save[i] = ΔZx + + G.Z_dims[i] = collect(size(X)) + end + end + + G.split_scales && (X = reshape(cat_states(X_save, X), orig_shape)) + G.split_scales && (ΔX = reshape(cat_states(ΔX_save, ΔX), orig_shape)) + return ΔX, X +end ## Jacobian-related utils function jacobian(ΔX::AbstractArray{T, N}, Δθ::Vector{Parameter}, X, G::NetworkGlow) where {T, N} - if G.split_scales Z_save = array_of_array(ΔX, G.L-1) ΔZ_save = array_of_array(ΔX, G.L-1) end + orig_shape = size(X) + logdet = 0 cls = 2*G.K*G.L ΔθAN = Vector{Parameter}(undef, 0) @@ -217,8 +250,8 @@ function jacobian(ΔX::AbstractArray{T, N}, Δθ::Vector{Parameter}, X, G::Netwo end end if G.split_scales - X = cat_states(Z_save, X) - ΔX = cat_states(ΔZ_save, ΔX) + X = reshape(cat_states(Z_save, X), orig_shape) + ΔX = reshape(cat_states(ΔZ_save, ΔX), orig_shape) end return ΔX, X, logdet, vcat(ΔθAN, ΔθCL) diff --git a/test/test_layers/test_actnorm.jl b/test/test_layers/test_actnorm.jl index 852549df..c066ffcb 100644 --- a/test/test_layers/test_actnorm.jl +++ b/test/test_layers/test_actnorm.jl @@ -75,26 +75,27 @@ AN_rev = reverse(AN) # Test with logdet enabled AN = ActNorm(nc; logdet=true) Y, lgdt = AN.forward(X) +#X_, lgdt = AN.inverse(X) # Test initialization @test isapprox(mean(Y), 0f0; atol=1f-6) @test isapprox(var(Y), 1f0; atol=1f-3) # Test invertibility -@test isapprox(norm(X - AN.inverse(AN.forward(X)[1]))/norm(X), 0f0, atol=1f-6) -@test isapprox(norm(X - AN.forward(AN.inverse(X))[1])/norm(X), 0f0, atol=1f-6) +@test isapprox(norm(X - AN.inverse(AN.forward(X)[1])[1])/norm(X), 0f0, atol=1f-6) +@test isapprox(norm(X - AN.forward(AN.inverse(X)[1])[1])/norm(X), 0f0, atol=1f-6) # Reversed layer (all combinations) AN_rev = reverse(AN) -@test isapprox(norm(X - AN_rev.inverse(AN_rev.forward(X)[1]))/norm(X), 0f0, atol=1f-6) -@test isapprox(norm(X - AN_rev.forward(AN_rev.inverse(X))[1])/norm(X), 0f0, atol=1f-6) +@test isapprox(norm(X - AN_rev.inverse(AN_rev.forward(X)[1])[1])/norm(X), 0f0, atol=1f-6) +@test isapprox(norm(X - AN_rev.forward(AN_rev.inverse(X)[1])[1])/norm(X), 0f0, atol=1f-6) @test isapprox(norm(X - AN_rev.forward(AN.forward(X)[1])[1])/norm(X), 0f0, atol=1f-6) -@test isapprox(norm(X - AN_rev.inverse(AN.inverse(X)))/norm(X), 0f0, atol=1f-6) +@test isapprox(norm(X - AN_rev.inverse(AN.inverse(X)[1])[1])/norm(X), 0f0, atol=1f-6) @test isapprox(norm(X - AN.forward(AN_rev.forward(X)[1])[1])/norm(X), 0f0, atol=1f-6) -@test isapprox(norm(X - AN.inverse(AN_rev.inverse(X)))/norm(X), 0f0, atol=1f-6) +@test isapprox(norm(X - AN.inverse(AN_rev.inverse(X)[1])[1])/norm(X), 0f0, atol=1f-6) ############################################################################### diff --git a/test/test_layers/test_coupling_layer_glow.jl b/test/test_layers/test_coupling_layer_glow.jl index 34135b87..a5b9a927 100644 --- a/test/test_layers/test_coupling_layer_glow.jl +++ b/test/test_layers/test_coupling_layer_glow.jl @@ -1,7 +1,6 @@ # Invertible CNN layer from Dinh et al. (2017)/Kingma and Dhariwal (2018) # Author: Philipp Witte, pwitte3@gatech.edu # Date: January 2020 - using InvertibleNetworks, LinearAlgebra, Test, Random # Random seed @@ -22,115 +21,126 @@ X = randn(Float32, nx, ny, k, batchsize) X0 = randn(Float32, nx, ny, k, batchsize) dX = X - X0 -# 1x1 convolution and residual blocks -C = Conv1x1(k) -RB = ResidualBlock(div(k,2), n_hidden; n_out=k, k1=3, k2=3, p1=1, p2=1, fan=true) -L = CouplingLayerGlow(C, RB; logdet=true) - -X_ = L.inverse(L.forward(X)[1]) -@test isapprox(norm(X - X_)/norm(X), 0f0; atol=1e-2) - -X_ = L.forward(L.inverse(X))[1] -@test isapprox(norm(X - X_)/norm(X), 0f0; atol=1e-2) - -################################################################################################### -# Gradient tests - -# Loss Function -function loss(L, X, Y) - Y_, logdet = L.forward(X) - f = mse(Y_, Y) - logdet - ΔY = ∇mse(Y_, Y) - ΔX = L.backward(ΔY, Y_)[1] - - # Pass back gradients w.r.t. input X and from the residual block and 1x1 conv. layer - return f, ΔX, L.C.v1.grad, L.C.v2.grad, L.C.v3.grad, L.RB.W1.grad, L.RB.W2.grad, L.RB.W3.grad +for (logdet,rev) in [(true, false),(false, true)] + # 1x1 convolution and residual blocks + C = Conv1x1(k) + RB = ResidualBlock(div(k,2), n_hidden; n_out=k, k1=3, k2=3, p1=1, p2=1, fan=true) + + L = CouplingLayerGlow(C, RB; logdet=logdet) + rev && (L = reverse(L)) + + L.logdet ? (Y, logdet_i) = L.forward(X) : Y = L.forward(X) + L.logdet ? (X_, logdet_i) = L.inverse(Y) : X_ = L.inverse(Y) + @test isapprox(norm(X - X_)/norm(X), 0f0; atol=1e-2) + + L.logdet ? (Y, logdet_i) = L.forward(X) : Y = L.forward(X) + X_ = L.backward(Y.*0f0, Y)[2] + @test isapprox(norm(X_-X)/norm(X), 0f0; atol=1e-2) + + + ################################################################################################### + # Gradient tests + + # Loss Function + function loss(L, X, Y) + logdet_i = 0 + L.logdet ? (Y_, logdet_i) = L.forward(X) : Y_ = L.forward(X) + f = mse(Y_, Y) - logdet_i + ΔY = ∇mse(Y_, Y) + ΔX = L.backward(ΔY, Y_)[1] + + # Pass back gradients w.r.t. input X and from the residual block and 1x1 conv. layer + return f, ΔX, L.C.v1.grad, L.C.v2.grad, L.C.v3.grad, L.RB.W1.grad, L.RB.W2.grad, L.RB.W3.grad + end + + # Gradient test w.r.t. input X0 + L.logdet ? (Y, logdet_i) = L.forward(X) : Y = L.forward(X) + f0, ΔX = loss(L, X0, Y)[1:2] + h = 0.1f0 + maxiter = 6 + err1 = zeros(Float32, maxiter) + err2 = zeros(Float32, maxiter) + + print("\nGradient wrt input of glow coupling layer\n") + for j=1:maxiter + f = loss(L, X0 + h*dX, Y)[1] + err1[j] = abs(f - f0) + err2[j] = abs(f - f0 - h*dot(dX, ΔX)) + print(err1[j], "; ", err2[j], "\n") + h = h/2f0 + end + + @test isapprox(err1[end] / (err1[1]/2^(maxiter-1)), 1f0; atol=1f0) + @test isapprox(err2[end] / (err2[1]/4^(maxiter-1)), 1f0; atol=1f0) + + + # Gradient test w.r.t. weights of residual block + # Invertible layers + C0 = Conv1x1(k) + RB0 = ResidualBlock(div(k,2), n_hidden;n_out=k, k1=3, k2=3, p1=1, p2=1, fan=true) + L01 = CouplingLayerGlow(C0, RB; logdet=logdet) + L02 = CouplingLayerGlow(C, RB0; logdet=logdet) + + rev && (L01 = reverse(L01)) + rev && (L02 = reverse(L02)) + + L.logdet ? (Y, logdet_i) = L.forward(X) : Y = L.forward(X) + Lini = deepcopy(L02) + dW1 = L.RB.W1.data - L02.RB.W1.data + dW2 = L.RB.W2.data - L02.RB.W2.data + dW3 = L.RB.W3.data - L02.RB.W3.data + + f0, ΔX, Δv1, Δv2, Δv3, ΔW1, ΔW2, ΔW3 = loss(L02, X, Y) + h = 0.1f0 + maxiter = 4 + err3 = zeros(Float32, maxiter) + err4 = zeros(Float32, maxiter) + + print("\nGradient wrt RB weights glow coupling layer\n") + for j=1:maxiter + L02.RB.W1.data = Lini.RB.W1.data + h*dW1 + L02.RB.W2.data = Lini.RB.W2.data + h*dW2 + L02.RB.W3.data = Lini.RB.W3.data + h*dW3 + f = loss(L02, X, Y)[1] + err3[j] = abs(f - f0) + err4[j] = abs(f - f0 - h*dot(dW1, ΔW1) - h*dot(dW2, ΔW2) - h*dot(dW3, ΔW3)) + print(err3[j], "; ", err4[j], "\n") + h = h/2f0 + end + + @test isapprox(err3[end] / (err3[1]/2^(maxiter-1)), 1f0; atol=1f0) + @test isapprox(err4[end] / (err4[1]/4^(maxiter-1)), 1f0; atol=1f0) + + + # Gradient test w.r.t. 1x1 conv weights + L.logdet ? (Y, logdet_i) = L.forward(X) : Y = L.forward(X) + Lini = deepcopy(L01) + dv1 = C.v1.data - C0.v1.data + dv2 = C.v2.data - C0.v2.data + dv3 = C.v3.data - C0.v3.data + + f0, ΔX, Δv1, Δv2, Δv3, ΔW1, ΔW2, ΔW3 = loss(L01, X, Y) + h = 0.1f0 + maxiter = 4 + err5 = zeros(Float32, maxiter) + err6 = zeros(Float32, maxiter) + + print("\nGradient test wrt 1x1 conv weights glow coupling layer\n") + for j=1:maxiter + L01.C.v1.data = Lini.C.v1.data + h*dv1 + L01.C.v2.data = Lini.C.v2.data + h*dv2 + L01.C.v3.data = Lini.C.v3.data + h*dv3 + f = loss(L01, X, Y)[1] + err5[j] = abs(f - f0) + err6[j] = abs(f - f0 - h*dot(dv1, Δv1) - h*dot(dv2, Δv2) - h*dot(dv3, Δv3)) + print(err5[j], "; ", err6[j], "\n") + h = h/2f0 + end + + @test isapprox(err5[end] / (err5[1]/2^(maxiter-1)), 1f0; atol=1f0) + @test isapprox(err6[end] / (err6[1]/4^(maxiter-1)), 1f0; atol=1f0) end -# Invertible layers -C0 = Conv1x1(k) -RB0 = ResidualBlock(div(k,2), n_hidden;n_out=k, k1=3, k2=3, p1=1, p2=1, fan=true) -L01 = CouplingLayerGlow(C0, RB; logdet=true) -L02 = CouplingLayerGlow(C, RB0; logdet=true) - -# Gradient test w.r.t. input X0 -Y = L.forward(X)[1] -f0, ΔX = loss(L, X0, Y)[1:2] -h = 0.1f0 -maxiter = 6 -err1 = zeros(Float32, maxiter) -err2 = zeros(Float32, maxiter) - -print("\nGradient test coupling layer\n") -for j=1:maxiter - f = loss(L, X0 + h*dX, Y)[1] - err1[j] = abs(f - f0) - err2[j] = abs(f - f0 - h*dot(dX, ΔX)) - print(err1[j], "; ", err2[j], "\n") - global h = h/2f0 -end - -@test isapprox(err1[end] / (err1[1]/2^(maxiter-1)), 1f0; atol=1f0) -@test isapprox(err2[end] / (err2[1]/4^(maxiter-1)), 1f0; atol=1f0) - - -# Gradient test w.r.t. weights of residual block -Y = L.forward(X)[1] -Lini = deepcopy(L02) -dW1 = L.RB.W1.data - L02.RB.W1.data -dW2 = L.RB.W2.data - L02.RB.W2.data -dW3 = L.RB.W3.data - L02.RB.W3.data - -f0, ΔX, Δv1, Δv2, Δv3, ΔW1, ΔW2, ΔW3 = loss(L02, X, Y) -h = 0.1f0 -maxiter = 4 -err3 = zeros(Float32, maxiter) -err4 = zeros(Float32, maxiter) - -print("\nGradient test coupling layer\n") -for j=1:maxiter - L02.RB.W1.data = Lini.RB.W1.data + h*dW1 - L02.RB.W2.data = Lini.RB.W2.data + h*dW2 - L02.RB.W3.data = Lini.RB.W3.data + h*dW3 - f = loss(L02, X, Y)[1] - err3[j] = abs(f - f0) - err4[j] = abs(f - f0 - h*dot(dW1, ΔW1) - h*dot(dW2, ΔW2) - h*dot(dW3, ΔW3)) - print(err3[j], "; ", err4[j], "\n") - global h = h/2f0 -end - -@test isapprox(err3[end] / (err3[1]/2^(maxiter-1)), 1f0; atol=1f0) -@test isapprox(err4[end] / (err4[1]/4^(maxiter-1)), 1f0; atol=1f0) - -# Gradient test w.r.t. 1x1 conv weights -Y = L.forward(X)[1] -Lini = deepcopy(L01) -dv1 = C.v1.data - C0.v1.data -dv2 = C.v2.data - C0.v2.data -dv3 = C.v3.data - C0.v3.data - -f0, ΔX, Δv1, Δv2, Δv3, ΔW1, ΔW2, ΔW3 = loss(L01, X, Y) -h = 0.1f0 -maxiter = 4 -err5 = zeros(Float32, maxiter) -err6 = zeros(Float32, maxiter) - -print("\nGradient test coupling layer\n") -for j=1:maxiter - L01.C.v1.data = Lini.C.v1.data + h*dv1 - L01.C.v2.data = Lini.C.v2.data + h*dv2 - L01.C.v3.data = Lini.C.v3.data + h*dv3 - f = loss(L01, X, Y)[1] - err5[j] = abs(f - f0) - err6[j] = abs(f - f0 - h*dot(dv1, Δv1) - h*dot(dv2, Δv2) - h*dot(dv3, Δv3)) - print(err5[j], "; ", err6[j], "\n") - global h = h/2f0 -end - -@test isapprox(err5[end] / (err5[1]/2^(maxiter-1)), 1f0; atol=1f0) -@test isapprox(err6[end] / (err6[1]/4^(maxiter-1)), 1f0; atol=1f0) - - ################################################################################################### # Jacobian-related tests diff --git a/test/test_layers/test_layer_conv1x1.jl b/test/test_layers/test_layer_conv1x1.jl index 9a2e7ea9..7697071c 100644 --- a/test/test_layers/test_layer_conv1x1.jl +++ b/test/test_layers/test_layer_conv1x1.jl @@ -74,7 +74,7 @@ Y_ = C0.forward(X) # Test gradients are zero in inverse pass if freeze =true -C_frozen = Conv1x1(v10, v20, v30;freeze=true) |> device +C_frozen = Conv1x1(v10, v20, v30; freeze=true) |> device # Predicted data and misfit C_frozen.v1.grad = nothing diff --git a/test/test_networks/test_glow.jl b/test/test_networks/test_glow.jl index 9bbd0542..640f2f55 100644 --- a/test/test_networks/test_glow.jl +++ b/test/test_networks/test_glow.jl @@ -19,15 +19,118 @@ K = 2 for split_scales = [true,false] for N in [(nx, ny), (nx, ny, nz)] - ###########################################Test with split_scales = false ######################### - # Invertibility + ############################### Test reverse####################################### + println("testing reverse glow with split_scales=$(split_scales) ndims=$(length(N))") + # Network and input + G = NetworkGlow(n_in, n_hidden, L, K; logdet = true, split_scales=split_scales, ndims=length(N)) + G = reverse(G) + + X = rand(Float32, N..., n_in, batchsize) + + #Y = G.inverse(X) + #X_ = G.forward(Y) + + G.logdet ? (Y, logdet_i) = G.inverse(X) : X_ = G.inverse(X) + G.logdet ? (X_, logdet_i) = G.forward(Y) : X_ = G.forward(Y) + + + @test isapprox(norm(X - X_)/norm(X), 0f0; atol=1f-5) + + ################################################################################################### + # Test gradients are set and cleared + G.backward(Y, Y) + + P = get_params(G) + gsum = 0 + for p in P + ~isnothing(p.grad) && (gsum += 1) + end + @test isequal(gsum, L*K*10) + + clear_grad!(G) + gsum = 0 + for p in P + ~isnothing(p.grad) && (gsum += 1) + end + @test isequal(gsum, 0) + + + ################################################################################################### + # Gradient test + function loss_rev(G, X) + #Y = G.forward(X) + logdet_i = 0 + G.logdet ? (Y, logdet_i) = G.forward(X) : Y = G.forward(X) + + f = -log_likelihood(Y) - logdet_i + ΔY = -∇log_likelihood(Y) + ΔX, X_ = G.backward(ΔY, Y) + return f, ΔX, G.CL[1,1].RB.W1.grad, G.CL[1,1].C.v1.grad + end + + # Gradient test w.r.t. input + X = rand(Float32, N..., n_in, batchsize) + X0 = rand(Float32, N..., n_in, batchsize) + dX = X - X0 + + f0, ΔX = loss_rev(G, X0)[1:2] + h = 0.1f0 + maxiter = 4 + err1 = zeros(Float32, maxiter) + err2 = zeros(Float32, maxiter) + + print("\nGradient test glow: input\n") + for j=1:maxiter + f = loss_rev(G, X0 + h*dX,)[1] + err1[j] = abs(f - f0) + err2[j] = abs(f - f0 - h*dot(dX, ΔX)) + print(err1[j], "; ", err2[j], "\n") + h = h/2f0 + end + + @test isapprox(err1[end] / (err1[1]/2^(maxiter-1)), 1f0; atol=1f1) + @test isapprox(err2[end] / (err2[1]/4^(maxiter-1)), 1f0; atol=1f1) + + # Test one parameter from residual block and 1x1 conv + G0 = NetworkGlow(n_in, n_hidden, L, K; logdet = false, split_scales=split_scales, ndims=length(N)) + G0.forward(X) + G0 = reverse(G0) + Gini = deepcopy(G0) + + dW = G.CL[1,1].RB.W1.data - G0.CL[1,1].RB.W1.data + dv = G.CL[1,1].C.v1.data - G0.CL[1,1].C.v1.data + + f0, ΔX, ΔW, Δv = loss_rev(G0, X) + h = 0.1f0 + maxiter = 4 + err3 = zeros(Float32, maxiter) + err4 = zeros(Float32, maxiter) + + print("\nGradient test glow: input\n") + for j=1:maxiter + G0.CL[1,1].RB.W1.data = Gini.CL[1,1].RB.W1.data + h*dW + G0.CL[1,1].C.v1.data = Gini.CL[1,1].C.v1.data + h*dv + + f = loss_rev(G0, X)[1] + err3[j] = abs(f - f0) + err4[j] = abs(f - f0 - h*dot(dW, ΔW) - h*dot(dv, Δv)) + print(err3[j], "; ", err4[j], "\n") + h = h/2f0 + end + + @test isapprox(err3[end] / (err3[1]/2^(maxiter-1)), 1f0; atol=1f1) + @test isapprox(err4[end] / (err4[1]/4^(maxiter-1)), 1f0; atol=1f1) + + + ########################################## Not reversed ####################### + println("testing glow with split_scales=$(split_scales) ndims=$(length(N))") # Network and input G = NetworkGlow(n_in, n_hidden, L, K; split_scales=split_scales, ndims=length(N)) X = rand(Float32, N..., n_in, batchsize) - Y = G.forward(X)[1] - X_ = G.inverse(Y) + G.logdet ? (Y, logdet_i) = G.forward(X) : Y = G.forward(X) + G.logdet ? (X_, logdet_i) = G.inverse(Y) : X_ = G.inverse(Y) @test isapprox(norm(X - X_)/norm(X), 0f0; atol=1f-5) @@ -62,9 +165,7 @@ for split_scales = [true,false] end # Gradient test w.r.t. input - G = NetworkGlow(n_in, n_hidden, L, K) - X = rand(Float32, nx, ny, n_in, batchsize) - X0 = rand(Float32, nx, ny, n_in, batchsize) + X0 = rand(Float32, N..., n_in, batchsize) dX = X - X0 f0, ΔX = loss(G, X0)[1:2] @@ -86,9 +187,8 @@ for split_scales = [true,false] @test isapprox(err2[end] / (err2[1]/4^(maxiter-1)), 1f0; atol=1f1) # Gradient test w.r.t. parameters - X = rand(Float32, nx, ny, n_in, batchsize) - G = NetworkGlow(n_in, n_hidden, L, K) - G0 = NetworkGlow(n_in, n_hidden, L, K) + G = NetworkGlow(n_in, n_hidden, L, K; split_scales=split_scales, ndims=length(N)) + G0 = NetworkGlow(n_in, n_hidden, L, K; split_scales=split_scales, ndims=length(N)) Gini = deepcopy(G0) # Test one parameter from residual block and 1x1 conv @@ -122,16 +222,16 @@ for split_scales = [true,false] # Gradient test # Initialization - G = NetworkGlow(n_in, n_hidden, L, K); G.forward(randn(Float32, nx, ny, n_in, batchsize)) + G.forward(randn(Float32, N..., n_in, batchsize)) θ = deepcopy(get_params(G)) - G0 = NetworkGlow(n_in, n_hidden, L, K); G0.forward(randn(Float32, nx, ny, n_in, batchsize)) + G0.forward(randn(Float32, N..., n_in, batchsize)) θ0 = deepcopy(get_params(G0)) - X = randn(Float32, nx, ny, n_in, batchsize) + X = randn(Float32, N..., n_in, batchsize) # Perturbation (normalized) dθ = θ-θ0 dθ .*= norm.(θ0)./(norm.(dθ).+1f-10) - dX = randn(Float32, nx, ny, n_in, batchsize); dX *= norm(X)/norm(dX) + dX = randn(Float32, N..., n_in, batchsize); dX *= norm(X)/norm(dX) # Jacobian eval dY, Y, _, _ = G.jacobian(dX, dθ, X) @@ -165,4 +265,4 @@ for split_scales = [true,false] @test isapprox(a, b; rtol=1f-3) end -end \ No newline at end of file +end