@@ -11,15 +11,15 @@ mutable struct LikelihoodInformed{FT <: Real} <: PairedDataContainerProcessor
1111 dim_criterion:: Tuple{Symbol, <:Number}
1212 α:: FT
1313 grad_type:: Symbol
14- use_prior_samples :: Bool
14+ use_data_as_samples :: Bool
1515end
1616
17- function likelihood_informed (retain_KL; alpha = 0.0 , grad_type = :localsl , use_prior_samples = true )
17+ function likelihood_informed (; retain_KL, alpha = 0.0 , grad_type = :localsl , use_data_as_samples = false )
1818 if grad_type ∉ [:linreg , :localsl ]
1919 @error " Unknown grad_type=$grad_type "
2020 end
2121
22- LikelihoodInformed (nothing , nothing , nothing , (:retain_KL , retain_KL), alpha, grad_type, use_prior_samples )
22+ LikelihoodInformed (nothing , nothing , nothing , (:retain_KL , retain_KL), alpha, grad_type, use_data_as_samples )
2323end
2424
2525get_encoder_mat (li:: LikelihoodInformed ) = li. encoder_mat
@@ -45,14 +45,13 @@ function initialize_processor!(
4545 else
4646 get_structure_vec (output_structure_vectors, :observation )
4747 end
48- samples_in, samples_out = if li. use_prior_samples
49- @assert α ≈ 0.0
48+ samples_in, samples_out = if li. use_data_as_samples
49+ (in_data, out_data)
50+ else
5051 (
51- get_structure_vec (input_structure_vectors, :prior_samples_in ),
52- get_structure_vec (output_structure_vectors, :prior_samples_out ),
52+ get_structure_vec (input_structure_vectors, :samples_in ),
53+ get_structure_vec (output_structure_vectors, :samples_out ),
5354 )
54- else
55- (in_data, out_data)
5655 end
5756 obs_noise_cov = get_structure_mat (output_structure_matrices, :obs_noise_cov )
5857 noise_cov_inv = inv (obs_noise_cov)
@@ -79,10 +78,10 @@ function initialize_processor!(
7978
8079 li. encoder_mat = if apply_to == " in" || α ≈ 0
8180 decomp = if apply_to == " in"
82- eigen (mean (grad' * noise_cov_inv * ((1 - α)obs_noise_cov + α^ 2 * (y - g) * (y - g)' ) * noise_cov_inv * grad for (g, grad) in zip (eachcol (samples_out), grads)), sortby = (- ))
81+ eigen (hermitianpart ( mean (grad' * noise_cov_inv * ((1 - α)obs_noise_cov + α^ 2 * (y - g) * (y - g)' ) * noise_cov_inv * grad for (g, grad) in zip (eachcol (samples_out), grads) )), sortby = (- ))
8382 else
8483 @assert apply_to == " out"
85- eigen (mean (grad * grad' for grad in grads), obs_noise_cov, sortby = (- ))
84+ eigen (hermitianpart ( mean (grad * grad' for grad in grads) ), obs_noise_cov, sortby = (- ))
8685 end
8786
8887 if li. dim_criterion[1 ] == :retain_KL
@@ -93,6 +92,7 @@ function initialize_processor!(
9392 @assert li. dim_criterion[1 ] == :dimension
9493 trunc_val = li. dim_criterion[2 ]
9594 end
95+ @info " truncating at $trunc_val /$(length (sv_cumsum)) retaining $(100.0 * sv_cumsum[trunc_val]) % of the KL divergence reduction"
9696 li. encoder_mat = decomp. vectors[:, 1 : trunc_val]'
9797 else
9898 @assert apply_to == " out"
@@ -135,13 +135,18 @@ function initialize_processor!(
135135 if li. dim_criterion[1 ] == :retain_KL
136136 retain_KL = li. dim_criterion[2 ]
137137 ref = f (M, zeros (output_dim, 0 ))
138- if f (M, Vs) / ref ≤ 1 - retain_KL
138+ val = f (M, Vs)
139+ if val / ref ≤ 1 - retain_KL
140+ @info " truncating at $k /$output_dim retaining $(100.0 * (1 - val/ ref)) % of the KL divergence reduction"
139141 break # TODO : Start bisecting?
140142 else
141143 k *= 2
142144 end
143145 else
144146 @assert li. dim_criterion[1 ] == :dimension
147+ ref = f (M, zeros (output_dim, 0 ))
148+ val = f (M, Vs)
149+ @info " truncating at $(li. dim_criterion[2 ]) /$output_dim retaining $(100.0 * val/ ref) % of the KL divergence reduction"
145150 break
146151 end
147152 end
0 commit comments