@@ -25,6 +25,7 @@ export create_encoder_schedule,
2525
2626
2727const StructureMatrix = Union{UniformScaling, AbstractMatrix}
28+ const StructureVector = Union{AbstractVector, AbstractMatrix} # In case of a matrix, the columns should be seen as vectors
2829
2930"""
3031$(DocStringExtensions. TYPEDSIGNATURES)
@@ -78,6 +79,24 @@ Base.:(==)(a::PDCP, b::PDCP) where {PDCP <: PairedDataContainerProcessor} =
7879
7980# ###
8081
82+ function get_structure_vec (structure_vecs, name = nothing )
83+ if isnothing (name)
84+ if size (structure_vecs) == 1
85+ return only (values (structure_vecs))
86+ elseif isempty (structure_vecs)
87+ @error " Please provide a structure vector."
88+ else
89+ @error " Structure vectors $(collect (keys (structure_vecs))) are present. Please indicate which to use."
90+ end
91+ else
92+ if haskey (structure_vecs, name)
93+ return structure_vecs[name]
94+ else
95+ @error " Structure vector $name not found. Options: $(collect (keys (structure_vecs))) ."
96+ end
97+ end
98+ end
99+
81100function get_structure_mat (structure_mats, name = nothing )
82101 if isnothing (name)
83102 if size (structure_mats) == 1
@@ -123,25 +142,28 @@ function _initialize_and_encode_data!(
123142 proc:: PairedDataContainerProcessor ,
124143 data,
125144 structure_mats,
145+ structure_vecs,
126146 apply_to:: AS ,
127147) where {AS <: AbstractString }
128- initialize_processor! (proc, get_data (data)... , structure_mats... , apply_to)
148+ initialize_processor! (proc, get_data (data)... , structure_mats... , structure_vecs ... , apply_to)
129149 return _encode_data (proc, data, apply_to)
130150end
131151
132152function _initialize_and_encode_data! (
133153 proc:: DataContainerProcessor ,
134154 data,
135155 structure_mats,
156+ structure_vecs,
136157 apply_to:: AS ,
137158) where {AS <: AbstractString }
138159 input_data, output_data = get_data (data)
139160 input_structure_mats, output_structure_mats = structure_mats
161+ input_structure_vecs, output_structure_vecs = structure_vecs
140162
141163 if apply_to == " in"
142- initialize_processor! (proc, input_data, input_structure_mats)
164+ initialize_processor! (proc, input_data, input_structure_mats, input_structure_vecs )
143165 elseif apply_to == " out"
144- initialize_processor! (proc, output_data, output_structure_mats)
166+ initialize_processor! (proc, output_data, output_structure_mats, output_structure_vecs )
145167 else
146168 bad_apply_to (apply_to)
147169 end
@@ -210,23 +232,31 @@ function initialize_and_encode_with_schedule!(
210232 io_pairs:: PDC ;
211233 input_structure_mats = Dict {Symbol, StructureMatrix} (),
212234 output_structure_mats = Dict {Symbol, StructureMatrix} (),
235+ input_structure_vecs = Dict {Symbol, StructureVector} (),
236+ output_structure_vecs = Dict {Symbol, StructureVector} (),
213237 input_cov:: Union{Nothing, StructureMatrix} = nothing ,
214238 obs_noise_cov:: Union{Nothing, StructureMatrix} = nothing ,
239+ observation:: Union{Nothing, StructureVector} = nothing ,
240+ prior_samples_in:: Union{Nothing, StructureVector} = nothing ,
241+ prior_samples_out:: Union{Nothing, StructureVector} = nothing ,
215242) where {
216243 VV <: AbstractVector ,
217244 PDC <: PairedDataContainer ,
218245}
219246 processed_io_pairs = deepcopy (io_pairs)
220247
221- processed_input_structure_mats = deepcopy (input_structure_mats)
222- if ! isnothing (input_cov)
223- processed_input_structure_mats[:input_cov ] = input_cov
224- end
248+ input_structure_mats = deepcopy (input_structure_mats)
249+ ! isnothing (input_cov) && (input_structure_mats[:input_cov ] = input_cov)
225250
226- processed_output_structure_mats = deepcopy (output_structure_mats)
227- if ! isnothing (obs_noise_cov)
228- processed_output_structure_mats[:obs_noise_cov ] = obs_noise_cov
229- end
251+ output_structure_mats = deepcopy (output_structure_mats)
252+ ! isnothing (obs_noise_cov) && (output_structure_mats[:obs_noise_cov ] = obs_noise_cov)
253+
254+ input_structure_vecs = deepcopy (input_structure_vecs)
255+ ! isnothing (prior_samples_in) && (input_structure_vecs[:prior_samples_in ] = prior_samples_in)
256+
257+ output_structure_vecs = deepcopy (output_structure_vecs)
258+ ! isnothing (observation) && (output_structure_vecs[:observation ] = observation)
259+ ! isnothing (prior_samples_out) && (output_structure_vecs[:prior_samples_out ] = prior_samples_out)
230260
231261 # apply_to is the string "in", "out" etc.
232262 for (processor, apply_to) in encoder_schedule
@@ -235,26 +265,35 @@ function initialize_and_encode_with_schedule!(
235265 processed = _initialize_and_encode_data! (
236266 processor,
237267 processed_io_pairs,
238- (processed_input_structure_mats, processed_output_structure_mats),
268+ (input_structure_mats, output_structure_mats),
269+ (input_structure_vecs, output_structure_vecs),
239270 apply_to,
240271 )
241272
242273 if apply_to == " in"
243- processed_input_structure_mats = Dict (
274+ input_structure_mats = Dict (
244275 name => encode_structure_matrix (processor, mat)
245- for (name, mat) in processed_input_structure_mats
276+ for (name, mat) in input_structure_mats
277+ )
278+ input_structure_vecs = Dict (
279+ name => encode_data (processor, vec)
280+ for (name, vec) in input_structure_vecs
246281 )
247282 processed_io_pairs = PairedDataContainer (processed, get_outputs (processed_io_pairs))
248283 elseif apply_to == " out"
249- processed_output_structure_mats = Dict (
284+ output_structure_mats = Dict (
250285 name => encode_structure_matrix (processor, mat)
251- for (name, mat) in processed_output_structure_mats
286+ for (name, mat) in output_structure_mats
287+ )
288+ output_structure_vecs = Dict (
289+ name => encode_data (processor, vec)
290+ for (name, vec) in output_structure_vecs
252291 )
253292 processed_io_pairs = PairedDataContainer (get_inputs (processed_io_pairs), processed)
254293 end
255294 end
256295
257- return processed_io_pairs, processed_input_structure_mats, processed_output_structure_mats
296+ return processed_io_pairs, input_structure_mats, output_structure_mats, input_structure_vecs, output_structure_vecs
258297end
259298
260299# Functions to encode/decode with initialized schedule
427466include (" Utilities/canonical_correlation.jl" )
428467include (" Utilities/decorrelator.jl" )
429468include (" Utilities/elementwise_scaler.jl" )
469+ include (" Utilities/likelihood_informed.jl" )
430470
431471end # module
0 commit comments