From f5d7b082eaf474dddd7a8f42b1d1c759abdc0517 Mon Sep 17 00:00:00 2001 From: WillyHK <80256890+WillyHK@users.noreply.github.com> Date: Thu, 15 Feb 2024 16:04:18 +0100 Subject: [PATCH 1/3] Update optimize_flow.jl there was an false function resulting in an error --- src/optimize_flow.jl | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/src/optimize_flow.jl b/src/optimize_flow.jl index 7ccebd8..c4e273d 100644 --- a/src/optimize_flow.jl +++ b/src/optimize_flow.jl @@ -12,7 +12,7 @@ function negll_flow_loss(flow::F, x::AbstractMatrix{<:Real}, logd_orig::Abstract end function negll_flow(flow::F, x::AbstractMatrix{<:Real}, logd_orig::AbstractVector, logpdf::Tuple{Function, Function}) where F<:AbstractFlow - negll, back = Zygote.pullback(negll_flow, flow, x, logd_orig, logpdf[2]) + negll, back = Zygote.pullback(negll_flow_loss, flow, x, logd_orig, logpdf[2]) d_flow = back(one(eltype(x)))[1] return negll, d_flow end From 6bd67e49213043a26d23a7493be12c022c8a7856 Mon Sep 17 00:00:00 2001 From: WillyHK <80256890+WillyHK@users.noreply.github.com> Date: Tue, 20 Feb 2024 11:19:05 +0100 Subject: [PATCH 2/3] Update optimize_flow.jl --- src/optimize_flow.jl | 13 +++++++------ 1 file changed, 7 insertions(+), 6 deletions(-) diff --git a/src/optimize_flow.jl b/src/optimize_flow.jl index c4e273d..8cb73f0 100644 --- a/src/optimize_flow.jl +++ b/src/optimize_flow.jl @@ -18,22 +18,23 @@ function negll_flow(flow::F, x::AbstractMatrix{<:Real}, logd_orig::AbstractVecto end export negll_flow -function KLDiv_flow_loss(flow::F, x::AbstractMatrix{<:Real}, logd_orig::AbstractVector, logpdfs::Tuple{Function, Function}) where F<:AbstractFlow +function KLDiv_flow_loss(flow::F, x::AbstractMatrix{<:Real}, logd_orig::AbstractVector, logpdf::Function) where F<:AbstractFlow nsamples = size(x, 2) - flow_corr = fchain(flow,logpdfs[2].f) - logpdf_y = logpdfs[2].logdensity + flow_corr = fchain(flow,logpdf.f) + #logpdf_y = logpdfs[2].logdensity y, ladj = with_logabsdet_jacobian(flow_corr, x) - KLDiv = sum(exp.(logd_orig - vec(ladj)) .* (logd_orig - vec(ladj) - logpdf_y(y))) / nsamples + KLDiv = sum(exp.(logd_orig - vec(ladj)) .* (logd_orig - vec(ladj) - logpdf(y))) / nsamples return KLDiv end -function KLDiv_flow(flow::F, x::AbstractMatrix{<:Real}, logd_orig::AbstractVector, logpdfs::Tuple{Function, Function}) where F<:AbstractFlow - KLDiv, back = Zygote.pullback(KLDiv_flow_loss, flow, x, logd_orig, logpdfs) +function KLDiv_flow(flow::F, x::AbstractMatrix{<:Real}, logd_orig::AbstractVector, logpdf::Tuple{Function, Function}) where F<:AbstractFlow + KLDiv, back = Zygote.pullback(KLDiv_flow_loss, flow, x, logd_orig, logpdf[2]) d_flow = back(one(eltype(x)))[1] return KLDiv, d_flow end export KLDiv_flow + function optimize_flow(samples::Union{Matrix, Tuple{Matrix, Matrix}}, initial_flow::F where F<:AbstractFlow, optimizer; From e71696d1a50d4b25544176a8a40bcf06acf6596d Mon Sep 17 00:00:00 2001 From: WillyHK <80256890+WillyHK@users.noreply.github.com> Date: Tue, 20 Feb 2024 11:51:58 +0100 Subject: [PATCH 3/3] Update optimize_flow.jl --- src/optimize_flow.jl | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/src/optimize_flow.jl b/src/optimize_flow.jl index 8cb73f0..063efd5 100644 --- a/src/optimize_flow.jl +++ b/src/optimize_flow.jl @@ -23,7 +23,7 @@ function KLDiv_flow_loss(flow::F, x::AbstractMatrix{<:Real}, logd_orig::Abstract flow_corr = fchain(flow,logpdf.f) #logpdf_y = logpdfs[2].logdensity y, ladj = with_logabsdet_jacobian(flow_corr, x) - KLDiv = sum(exp.(logd_orig - vec(ladj)) .* (logd_orig - vec(ladj) - logpdf(y))) / nsamples + KLDiv = sum(exp.(logd_orig - vec(ladj)) .* (logd_orig - vec(ladj) - logpdf.logdensity(y))) / nsamples return KLDiv end