From 80b99c6d2a414fd13c61e1104f23f67b4108c251 Mon Sep 17 00:00:00 2001 From: akim126 Date: Fri, 26 Apr 2024 16:37:08 -0400 Subject: [PATCH 1/5] correct version --- Manifest.toml | 36 ++++++++++++------------------------ 1 file changed, 12 insertions(+), 24 deletions(-) diff --git a/Manifest.toml b/Manifest.toml index 57d6a35..14e6177 100644 --- a/Manifest.toml +++ b/Manifest.toml @@ -159,11 +159,11 @@ git-tree-sha1 = "aebf55e6d7795e02ca500a689d326ac979aaf89e" uuid = "9718e550-a3fa-408a-8086-8db961cd8217" version = "0.1.1" -[[deps.BioCore]] -deps = ["Automa", "BufferedStreams", "YAML"] -git-tree-sha1 = "476edbf4ef94594fff430a84ca96f86cb2327a71" -uuid = "37cfa864-2cd6-5c12-ad9e-b6597d696c81" -version = "2.0.5" +[[deps.BioGenerics]] +deps = ["TranscodingStreams"] +git-tree-sha1 = "7bbc085aebc6faa615740b63756e4986c9e85a70" +uuid = "47718e42-2ac5-11e9-14af-e5595289c2ea" +version = "0.1.4" [[deps.BitFlags]] git-tree-sha1 = "2dc09997850d68179b69dafb58ae806167a32b1b" @@ -2008,12 +2008,6 @@ weakdeps = ["ChainRulesCore", "InverseFunctions"] StatsFunsChainRulesCoreExt = "ChainRulesCore" StatsFunsInverseFunctionsExt = "InverseFunctions" -[[deps.StringEncodings]] -deps = ["Libiconv_jll"] -git-tree-sha1 = "b765e46ba27ecf6b44faf70df40c57aa3a547dcb" -uuid = "69024149-9ee7-55f6-a4c4-859efe599b68" -version = "0.3.7" - [[deps.StringManipulation]] deps = ["PrecompileTools"] git-tree-sha1 = "a04cabe79c5f01f4d723cc6704070ada0b9d46d5" @@ -2177,10 +2171,10 @@ uuid = "d80eeb9a-aca5-4d75-85e5-170c8b632249" version = "0.1.3" [[deps.VariantCallFormat]] -deps = ["Automa", "BGZFStreams", "BioCore", "BufferedStreams"] -git-tree-sha1 = "f73ea34d3085cdbf6a18fa4c4b690e0f4a147730" +deps = ["Automa", "BGZFStreams", "BioGenerics", "BufferedStreams"] +git-tree-sha1 = "96fbe09c9e3b488666c883772fed8f6c1256c714" uuid = "28eba6e3-a997-4ad9-87c6-d933b8bca6c1" -version = "0.5.5" +version = "0.5.6" [[deps.VectorizationBase]] deps = ["ArrayInterface", "CPUSummary", "HostCPUFeatures", "IfElse", "LayoutPointers", "Libdl", "LinearAlgebra", "SIMDTypes", "Static", "StaticArrayInterface"] @@ -2219,9 +2213,9 @@ version = "1.1.34+0" [[deps.XZ_jll]] deps = ["Artifacts", "JLLWrappers", "Libdl"] -git-tree-sha1 = "31c421e5516a6248dfb22c194519e37effbf1f30" +git-tree-sha1 = "ac88fb95ae6447c8dda6a5503f3bafd496ae8632" uuid = "ffd25f8a-64ca-5728-b0f7-c24cf3aae800" -version = "5.6.1+0" +version = "5.4.6+0" [[deps.Xorg_libX11_jll]] deps = ["Artifacts", "JLLWrappers", "Libdl", "Xorg_libxcb_jll", "Xorg_xtrans_jll"] @@ -2271,12 +2265,6 @@ git-tree-sha1 = "e92a1a012a10506618f10b7047e478403a046c77" uuid = "c5fb5394-a638-5e4d-96e5-b29de1b5cf10" version = "1.5.0+0" -[[deps.YAML]] -deps = ["Base64", "Dates", "Printf", "StringEncodings"] -git-tree-sha1 = "e6330e4b731a6af7959673621e91645eb1356884" -uuid = "ddb6d928-2868-570f-bddf-ab3f9cf99eb6" -version = "0.4.9" - [[deps.Zlib_jll]] deps = ["Libdl"] uuid = "83775a58-1f1d-513f-b197-d71354ab007a" @@ -2284,9 +2272,9 @@ version = "1.2.13+0" [[deps.Zstd_jll]] deps = ["Artifacts", "JLLWrappers", "Libdl"] -git-tree-sha1 = "49ce682769cd5de6c72dcf1b94ed7790cd08974c" +git-tree-sha1 = "e678132f07ddb5bfa46857f0d7620fb9be675d3b" uuid = "3161d3a3-bdf6-5164-811a-617609db77b4" -version = "1.5.5+0" +version = "1.5.6+0" [[deps.Zygote]] deps = ["AbstractFFTs", "ChainRules", "ChainRulesCore", "DiffRules", "Distributed", "FillArrays", "ForwardDiff", "GPUArrays", "GPUArraysCore", "IRTools", "InteractiveUtils", "LinearAlgebra", "LogExpFunctions", "MacroTools", "NaNMath", "PrecompileTools", "Random", "Requires", "SparseArrays", "SpecialFunctions", "Statistics", "ZygoteRules"] From 02a71d59afe10d2b4a8dd1058bfc3bd5bd0025d6 Mon Sep 17 00:00:00 2001 From: akim126 Date: Fri, 26 Apr 2024 16:37:35 -0400 Subject: [PATCH 2/5] use CI for CAVI improvement check --- src/loss.jl | 56 +++++++++++++++++------ src/train.jl | 126 +++++++++++++++++++++++++++++++++------------------ 2 files changed, 122 insertions(+), 60 deletions(-) diff --git a/src/loss.jl b/src/loss.jl index e6fa422..d72715b 100644 --- a/src/loss.jl +++ b/src/loss.jl @@ -14,7 +14,6 @@ end Calculates the log density of β based on a spiek and slab prior """ function log_prior(β::Vector, σ2_β::Vector, p_causal::Vector) - P = length(β) # prob_slab = 0.10 # L = prob_slab * 1_000 @@ -26,6 +25,23 @@ function log_prior(β::Vector, σ2_β::Vector, p_causal::Vector) return sum(logprobs) end +## TO COMPLETE ## +function log_prior_lse(β::Vector, σ2_β::Vector, p_causal::Vector) + + P = length(β) + # prob_slab = 0.10 + # L = prob_slab * 1_000 + # h2 = 0.10 + spike_σ2 = 1e-6 + slab_dist = Normal.(0, sqrt.(σ2_β .+ spike_σ2)) + #slab_dist = Normal.(0, sqrt.(σ2_β)) + spike_dist = Normal.(0, sqrt(spike_σ2)) + densities = pdf.(slab_dist, β) .* p_causal .+ pdf.(spike_dist, β) .* (1 .- p_causal) + logprobs = log.(densities) + ##logprobs = log.(pdf.(slab_dist, β) .* p_causal .+ pdf.(spike_dist, β) .* (1 .- p_causal)) + return sum(logprobs) +end + """ rss(β, coef, SE, R) Calculate the summary statistic RSS likelihood @@ -133,21 +149,31 @@ elbo( ) ``` """ -function elbo(z::Vector, q_μ::Vector, log_q_var::Vector, coef::Vector, SE::Vector, R::AbstractArray, σ2_β::Vector, p_causal::Vector, to) - q_var = @timeit to "q_var" exp.(log_q_var) - q = @timeit to "q" MvNormal(q_μ, Diagonal(q_var)) - q_sd = @timeit to "q_sd" sqrt.(q_var) - ϕ = @timeit to "ϕ" q_μ .+ q_sd .* z - # γ = compute_γ(q_μ, q_var) - # jl = joint_log_prob(γ .* ϕ, coef, SE, R) - jl = @timeit to "joint_log_prob" joint_log_prob(ϕ, coef, SE, R, σ2_β, p_causal, to) - q = @timeit to "logpd" logpdf(q, ϕ) - # jac = prod(z) - return (jl - q) -end +#function elbo(z::Vector, q_μ::Vector, log_q_var::Vector, coef::Vector, SE::Vector, R::AbstractArray, σ2_β::Vector, p_causal::Vector, to) +# q_var = @timeit to "q_var" exp.(log_q_var) +# q = @timeit to "q" MvNormal(q_μ, Diagonal(q_var)) +# q_sd = @timeit to "q_sd" sqrt.(q_var) +# ϕ = @timeit to "ϕ" q_μ .+ q_sd .* z +# # γ = compute_γ(q_μ, q_var) +# # jl = joint_log_prob(γ .* ϕ, coef, SE, R) +# jl = @timeit to "joint_log_prob" joint_log_prob(ϕ, coef, SE, R, σ2_β, p_causal, to) +# q = @timeit to "logpd" logpdf(q, ϕ) +# # jac = prod(z) +# return (jl - q) +#end + +#function elbo(z::Vector, q_μ::Vector, log_q_var::Vector, coef::Vector, Σ::AbstractPDMat, SRSinv::Matrix, σ2_β::Vector, p_causal::Vector, to) -function elbo(z::Vector, q_μ::Vector, log_q_var::Vector, coef::Vector, Σ::AbstractPDMat, SRSinv::Matrix, σ2_β::Vector, p_causal::Vector, to) - q_var = @timeit to "q_var" exp.(log_q_var) +""" + elbo(z, q_μ, q_var, coef, Σ, SRSinv, σ2_β, p_causal, to) + + Calculate ELBO sampled from MC sampling +""" +function elbo(z::Vector, q_μ::Vector, q_var::Vector, coef::Vector, Σ::AbstractPDMat, SRSinv::Matrix, σ2_β::Vector, p_causal::Vector, to) + #q_var = @timeit to "q_var" exp.(log_q_var) + #if (any(x->x<0, q_var)) + # println("Negative q_var: ", q_var[findall(x->x<0, q_var)]) # Debug output + #end q = @timeit to "q" MvNormal(q_μ, Diagonal(q_var)) q_sd = @timeit to "q_sd" sqrt.(q_var) ϕ = @timeit to "ϕ" q_μ .+ q_sd .* z diff --git a/src/train.jl b/src/train.jl index 5d65023..e66db05 100644 --- a/src/train.jl +++ b/src/train.jl @@ -1,3 +1,5 @@ +#using Plots + function check_no_nan(data) if sum(isnan.(data[1])) > 0 @@ -13,6 +15,7 @@ function check_no_nan(data) end end + """ fit_heritability_nn(model, q_μ, q_var, q_alpha, G, i) @@ -46,7 +49,7 @@ end yhat[:, 2] .= 1.0 ./ (1.0 .+ exp.(-yhat[:, 2])) ``` """ -function fit_heritability_nn(model, q_var, q_α, G, i=1; max_epochs=50, patience=30, mse_improvement_threshold=0.01, test_ratio=0.2, num_splits=5, weight_slab=1, weight_causal=1) +function fit_heritability_nn(model, q_var, q_α, G, i=1; max_epochs=50, patience=10, mse_improvement_threshold=0.1, test_ratio=0.2, num_splits=5, weight_slab=0.2, weight_causal=0.8) # RMSE function loss(model, x, y_slab, y_causal) ## ak: need two losses for slab variance and percent causal @@ -57,8 +60,6 @@ function fit_heritability_nn(model, q_var, q_α, G, i=1; max_epochs=50, patience weighted_loss_slab = weight_slab * loss_slab loss_causal = Flux.mse(yhat[2, :], y_causal) weighted_loss_causal = weight_causal * loss_causal - #println("loss_slab = $weighted_loss_slab") - #println("loss_causal = $weighted_loss_causal") total_loss = weighted_loss_slab + weighted_loss_causal ## ak: losses summed to form the total loss for training # if !isfinite(total_loss) # println("loss_slab = $loss_slab") @@ -132,28 +133,21 @@ function fit_heritability_nn(model, q_var, q_α, G, i=1; max_epochs=50, patience best_model_epoch = 1 train_losses = Float64[] test_losses = Float64[] + scaling_factor = 0.001 for epoch in 1:max_epochs check_no_nan(data[1]) train!(loss, model, data, opt) - #println("just trained") - #yhat_just_trained = model(transpose(G)) - #println("yhat just trained; slab, causal") - #println(yhat_just_trained[1, 1:5], yhat_just_trained[2, 1:5]) - train_loss = loss(model, best_train_data[1], log.(best_train_data[2]), logit.(best_train_data[3])) + train_loss = loss(model, best_train_data[1], log.(best_train_data[2]), logit.(best_train_data[3])) push!(train_losses, train_loss) - #println("computed train loss") # ak: validation loss test_loss = loss(model, best_test_data[1], log.(best_test_data[2]), logit.(best_test_data[3])) - #println("computed test loss") - #println("Test loss = $test_loss") - #println("Train loss = $train_loss") push!(test_losses, test_loss) # check for improvement in loss mse_improvement = (test_loss - best_loss) / test_loss + #println("MSE IMPROVEMENT: $mse_improvement") - #println("MSE improvement = $mse_improvement") # if improvement from prev iteration is greater than threshold if mse_improvement < -mse_improvement_threshold best_loss = test_loss @@ -162,7 +156,7 @@ function fit_heritability_nn(model, q_var, q_α, G, i=1; max_epochs=50, patience best_model_epoch = epoch # and save current model as best model best_model = deepcopy(model) - #println("in best model") + println("in best model") #yhat_current_best_model = best_model(transpose(G)) #println("yhat current best model; slab, causal") #println(yhat_current_best_model[1, 1:5], yhat_current_best_model[2, 1:5]) @@ -176,15 +170,17 @@ function fit_heritability_nn(model, q_var, q_α, G, i=1; max_epochs=50, patience @debug "$(ltime()) Early stopping after $epoch epochs." break end - + end + #plot_loss_vs_epochs(train_losses, test_losses, i, best_model_epoch) + #savefig("nn_split_iter$i.png") + @info "$(ltime()) Best model taken from epoch $best_model_epoch." return best_model end -#function train_cavi(p_causal, σ2_β, X_sd, i_iter, coef, SE, R, D, to; P = 1_000, n_elbo = 10, max_iter = 2, N = 10_000, σ2 = 1.0) function train_cavi(p_causal, σ2_β, X_sd, i_iter, coef, SE, R, D, to; P=1_000, n_elbo=10, max_iter=10, N=10_000, σ2=1.0) function clamp_ssr(ssr, max_value=709.7) # slightly below the threshold @@ -213,6 +209,9 @@ function train_cavi(p_causal, σ2_β, X_sd, i_iter, coef, SE, R, D, to; P=1_000, loss = -Inf prev_loss = -Inf + prev_loss_lower_limit = -Inf + prev_loss_upper_limit = Inf + current_elbo = -Inf prev_prev_loss = -Inf best_loss = -Inf cavi_loss = Float32[] @@ -230,56 +229,74 @@ function train_cavi(p_causal, σ2_β, X_sd, i_iter, coef, SE, R, D, to; P=1_000, # monitor loss convergence loss = 0.0 + elbo_loss = Float32[] @timeit to "elbo estimate" begin @inbounds for z in 1:n_elbo z = rand(Normal(0, 1), P) # loss_old = loss_old + elbo(z, q_μ, log.(q_var), coef, SE, R, σ2_β, p_causal, to) - loss = loss + elbo(z, q_μ, log.(q_var), coef, Σ_reg, SRSinv, σ2_β, p_causal, to) + # loss = loss + elbo(z, q_μ, log.(q_var), coef, Σ_reg, SRSinv, σ2_β, p_causal, to) + current_elbo = elbo(z, q_μ, q_var, coef, Σ_reg, SRSinv, σ2_β, p_causal, to) + loss = loss + current_elbo # @info "$(ltime()) loss_new = $loss, loss_old = $loss_old" + push!(elbo_loss, current_elbo) end end loss = loss / n_elbo - - if isnan(loss) == true - error("NaN loss detected.") + # 95% CI + current_loss_lower_limit = loss - 1.96 * std(elbo_loss) / sqrt(length(elbo_loss)) + current_loss_upper_limit = loss + 1.96 * std(elbo_loss) / sqrt(length(elbo_loss)) + + if (isnan(loss) == true || isinf(loss) == true) + #error("NaN or Inf loss detected.") + @info "$(ltime()) NaN or Inf loss detected.)" + @info "$(ltime()) $elbo_loss)" break end - @info "$(ltime()) iteration $i, loss = $(round(loss; digits = 2)) (bigger numbers are better)" - # ak: stopping criterion for oscillation - if (prev_loss > prev_prev_loss && loss < prev_loss) || - (prev_loss < prev_prev_loss && loss > prev_loss) - # ak: we want to keep q_μ, q_α, q_var, q_odds from i-1 iteration - @info "$(ltime()) Oscillation detected. Stopping at iteration $i." - cavi_iter = i - break - end + @info "$(ltime()) iteration $(i-1), loss = $(round(loss; digits = 2)) [$(round(current_loss_lower_limit; digits = 2)):$(round(current_loss_upper_limit; digits = 2))] (bigger numbers are better)" + # @info "$(ltime()) iteration $i, loss = $(round(loss; digits = 2)) (bigger numbers are better)" - # ak: stopping criterion for insufficient improvement (10%) - # ak: added a small constant to avoid division by zero - relative_improvement = abs(loss - prev_loss) / (abs(prev_loss) + 1e-8) - if relative_improvement < 0.01 - # ak: we want to keep q_μ, q_α, q_var, q_odds from i-1 iteration - @info "$(ltime()) Insufficient improvement. Stopping at iteration $i." - cavi_iter = i - break + if ( i > 1 ) + if ( current_loss_lower_limit > prev_loss_upper_limit ) + @info "$(ltime()) Sufficient improvement." + + elseif (prev_loss_lower_limit < current_loss_lower_limit && prev_loss_upper_limit > current_loss_lower_limit && current_loss_upper_limit > prev_loss_upper_limit) + @info "$(ltime()) Small improvement -- keep going" + else + @info "$(ltime()) No improvement." + cavi_iter = i-1 + break + end + + # ak: stopping criterion for insufficient improvement (10%) + # ak: added a small constant to avoid division by zero + relative_improvement = loss - prev_loss / (abs(prev_loss) + 1e-8) + # if relative_improvement < 0.01 + # # ak: we want to keep q_μ, q_α, q_var, q_odds from i-1 iteration + # @info "$(ltime()) Insufficient improvement. Stopping at iteration $i." + # cavi_iter = i + # break + # end end @timeit to "push cavi loss" push!(cavi_loss, Float32(loss)) # ak: update the previous losses for the next iteration - prev_prev_loss = prev_loss - prev_loss = loss + #prev_prev_loss = prev_loss + #prev_loss = loss q_μ_best = copy(q_μ) q_var_best = copy(q_var) q_α_best = copy(q_α) q_odds_best = copy(q_odds) best_loss = copy(loss) + prev_loss_lower_limit = copy(current_loss_lower_limit) + prev_loss_upper_limit = copy(current_loss_upper_limit) @info "$(ltime()) CAVI updates at iteration $i" @timeit to "update q_var" q_var .= σ2 ./ (diag(XtX) .+ 1 ./ σ2_β) ## ak: eq 8; \s^2_k; does not depend on alpha and mu from previous + q_sd .= sqrt.(q_var) @timeit to "update q_μ" begin @inbounds for k in 1:P J = setdiff(1:P, k) @@ -300,7 +317,7 @@ function train_cavi(p_causal, σ2_β, X_sd, i_iter, coef, SE, R, D, to; P=1_000, # with probability q_alpha, additive effect beta is normal with mean q_mu and variance q_var # return q_μ, q_α, q_var, q_odds, loss, cavi_loss - return q_μ_best, q_α_best, q_var_best, q_odds_best, best_loss, cavi_loss, cavi_iter + return q_μ_best, q_α_best, q_var_best, q_odds_best, best_loss, prev_loss_lower_limit, prev_loss_upper_limit, cavi_loss, cavi_iter end @@ -384,6 +401,8 @@ function train_until_convergence(coef::Vector, SE::Vector, R::AbstractArray, D:: end prev_loss = -Inf + prev_ci_lower = -Inf + prev_ci_upper = Inf model_init = deepcopy(model) prev_model = deepcopy(model) prev_prev_model = deepcopy(model) @@ -398,7 +417,7 @@ function train_until_convergence(coef::Vector, SE::Vector, R::AbstractArray, D:: # cavi_q_u is cavi trained estimated betas, and coef is from iteration before # q_μ, q_α, q_var, odds, new_loss, cavi_losses = train_cavi(cavi_q_μ, cavi_q_α, cavi_q_var, nn_p_causal, nn_σ2_β, X_sd, i, coef, SE, R, D) @timeit to "train_cavi" begin - q_μ, q_α, q_var, odds, new_loss, cavi_losses, cavi_iter = train_cavi( + q_μ, q_α, q_var, odds, new_loss, new_ci_lower, new_ci_upper, cavi_losses, cavi_iter = train_cavi( nn_p_causal, nn_σ2_β, X_sd, @@ -414,7 +433,7 @@ function train_until_convergence(coef::Vector, SE::Vector, R::AbstractArray, D:: @info "$(ltime()) Training CAVI finished" end - cavi_iter = cavi_iter - 1 + #cavi_iter = cavi_iter - 1 @info "$(ltime()) $cavi_iter updates for outer-loop iteration $i" @timeit to "GC" begin @@ -425,16 +444,33 @@ function train_until_convergence(coef::Vector, SE::Vector, R::AbstractArray, D:: @debug "$(ltime()) difference from n, n-1 (%) = $(abs(new_loss - prev_loss) / abs(prev_loss))" end + println("CHECKING FOR OUTER-LOOP CONVERGENCE") + println("previous range: [$prev_ci_lower:$prev_ci_upper]") + println("current range: [$new_ci_lower:$new_ci_upper]") + # check for convergence - if abs(new_loss - prev_loss) / abs(prev_loss) < threshold - @info "$(ltime()) converged!" - break + if i > 1 + # if abs(new_loss - prev_loss) / abs(prev_loss) < threshold + #if (( new_ci_upper < prev_ci_lower ) || ( prev_ci_lower < new_ci_lower && prev_ci_upper > new_ci_upper ) || + # ( new_ci_lower < prev_ci_lower && new_ci_upper > prev_ci_upper )) + # @info "$(ltime()) converged!" + # break + #end + + if (( new_ci_lower > prev_ci_upper ) || (prev_ci_lower < new_ci_lower && prev_ci_upper > new_ci_lower && new_ci_upper > prev_ci_upper)) + @info "$(ltime()) outer CAVI hasn't converged yet" + else + @info "$(ltime()) outer CAVI has converged!" + break + end end # plot_cavi_losses(cavi_losses, i) # savefig("cavi_loss_iter$i.png") prev_loss = copy(new_loss) + prev_ci_lower = copy(new_ci_lower) + prev_ci_upper = copy(new_ci_upper) ## ak: Set α(i) =α and μ(i) =μ cavi_q_μ = copy(q_μ) From 193ee202b8fdd1f361fc3b6ac32965420bbf180f Mon Sep 17 00:00:00 2001 From: akim126 Date: Fri, 26 Apr 2024 16:39:43 -0400 Subject: [PATCH 3/5] undo weigth change --- src/train.jl | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/src/train.jl b/src/train.jl index e66db05..135a397 100644 --- a/src/train.jl +++ b/src/train.jl @@ -49,7 +49,7 @@ end yhat[:, 2] .= 1.0 ./ (1.0 .+ exp.(-yhat[:, 2])) ``` """ -function fit_heritability_nn(model, q_var, q_α, G, i=1; max_epochs=50, patience=10, mse_improvement_threshold=0.1, test_ratio=0.2, num_splits=5, weight_slab=0.2, weight_causal=0.8) +function fit_heritability_nn(model, q_var, q_α, G, i=1; max_epochs=50, patience=10, mse_improvement_threshold=0.1, test_ratio=0.2, num_splits=5, weight_slab=1.0, weight_causal=1.0) # RMSE function loss(model, x, y_slab, y_causal) ## ak: need two losses for slab variance and percent causal From 0a47f4653205862c726d36fd94efa37c4782ae21 Mon Sep 17 00:00:00 2001 From: akim126 Date: Mon, 29 Apr 2024 11:21:20 -0400 Subject: [PATCH 4/5] apply LogSumExp trick to log_prior --- src/loss.jl | 20 ++++++++++++-------- 1 file changed, 12 insertions(+), 8 deletions(-) diff --git a/src/loss.jl b/src/loss.jl index d72715b..a8c80e2 100644 --- a/src/loss.jl +++ b/src/loss.jl @@ -29,16 +29,19 @@ end function log_prior_lse(β::Vector, σ2_β::Vector, p_causal::Vector) P = length(β) - # prob_slab = 0.10 - # L = prob_slab * 1_000 - # h2 = 0.10 spike_σ2 = 1e-6 slab_dist = Normal.(0, sqrt.(σ2_β .+ spike_σ2)) - #slab_dist = Normal.(0, sqrt.(σ2_β)) spike_dist = Normal.(0, sqrt(spike_σ2)) - densities = pdf.(slab_dist, β) .* p_causal .+ pdf.(spike_dist, β) .* (1 .- p_causal) - logprobs = log.(densities) - ##logprobs = log.(pdf.(slab_dist, β) .* p_causal .+ pdf.(spike_dist, β) .* (1 .- p_causal)) + + # compute log probabilities for slab and spike components using vectorized operations + log_prob_slab = logpdf.(slab_dist, β) .+ log.(p_causal) + log_prob_spike = logpdf.(spike_dist, β) .+ log.(1 .- p_causal) + + # applying the Log-Sum-Exp trick using vectorized operations + max_log_prob = max.(log_prob_slab, log_prob_spike) + logprobs = max_log_prob .+ log.(exp.(log_prob_slab .- max_log_prob) .+ exp.(log_prob_spike .- max_log_prob)) + + # sum of log probabilities return sum(logprobs) end @@ -130,7 +133,8 @@ joint_log_prob( """ joint_log_prob(β::Vector, coef::Vector, SE::Vector, R::Matrix, σ2_β::Vector, p_causal::Vector, to) = rss(β, coef, SE, R, to) + log_prior(β, σ2_β, p_causal) -joint_log_prob(β::Vector, coef::Vector, Σ::AbstractPDMat, SRSinv::Matrix, σ2_β::Vector, p_causal::Vector, to) = rss(β, coef, Σ, SRSinv, to) + log_prior(β, σ2_β, p_causal) +#joint_log_prob(β::Vector, coef::Vector, Σ::AbstractPDMat, SRSinv::Matrix, σ2_β::Vector, p_causal::Vector, to) = rss(β, coef, Σ, SRSinv, to) + log_prior(β, σ2_β, p_causal) +joint_log_prob(β::Vector, coef::Vector, Σ::AbstractPDMat, SRSinv::Matrix, σ2_β::Vector, p_causal::Vector, to) = rss(β, coef, Σ, SRSinv, to) + log_prior_lse(β, σ2_β, p_causal) """ elbo(z, q_μ, log_q_var, coef, SE, R, σ2_β, p_causal) From c0ae5a8f483134e1807ec47ae7d3ab47e679f41c Mon Sep 17 00:00:00 2001 From: akim126 Date: Mon, 29 Apr 2024 11:23:21 -0400 Subject: [PATCH 5/5] set appropriate threshold for NN relative improvement --- src/train.jl | 9 ++++++++- 1 file changed, 8 insertions(+), 1 deletion(-) diff --git a/src/train.jl b/src/train.jl index 135a397..0337a9a 100644 --- a/src/train.jl +++ b/src/train.jl @@ -144,6 +144,13 @@ function fit_heritability_nn(model, q_var, q_α, G, i=1; max_epochs=50, patience test_loss = loss(model, best_test_data[1], log.(best_test_data[2]), logit.(best_test_data[3])) push!(test_losses, test_loss) + if (epoch == 1) + mse_improvement_threshold = test_loss * scaling_factor + #if ( mse_improvement_threshold < 0.01) + # mse_improvement_threshold = 0.01 + #end + end + # check for improvement in loss mse_improvement = (test_loss - best_loss) / test_loss #println("MSE IMPROVEMENT: $mse_improvement") @@ -457,7 +464,7 @@ function train_until_convergence(coef::Vector, SE::Vector, R::AbstractArray, D:: # break #end - if (( new_ci_lower > prev_ci_upper ) || (prev_ci_lower < new_ci_lower && prev_ci_upper > new_ci_lower && new_ci_upper > prev_ci_upper)) + if (( new_ci_lower > prev_ci_upper )) #|| (prev_ci_lower < new_ci_lower && prev_ci_upper > new_ci_lower && new_ci_upper > prev_ci_upper)) @info "$(ltime()) outer CAVI hasn't converged yet" else @info "$(ltime()) outer CAVI has converged!"