Skip to content

Commit 8152241

Browse files
committed
Add structure vectors
1 parent 8ef3330 commit 8152241

4 files changed

Lines changed: 61 additions & 18 deletions

File tree

src/Utilities.jl

Lines changed: 57 additions & 17 deletions
Original file line numberDiff line numberDiff line change
@@ -25,6 +25,7 @@ export create_encoder_schedule,
2525

2626

2727
const 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+
81100
function 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)
130150
end
131151

132152
function _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
258297
end
259298

260299
# Functions to encode/decode with initialized schedule
@@ -427,5 +466,6 @@ end
427466
include("Utilities/canonical_correlation.jl")
428467
include("Utilities/decorrelator.jl")
429468
include("Utilities/elementwise_scaler.jl")
469+
include("Utilities/likelihood_informed.jl")
430470

431471
end # module

src/Utilities/canonical_correlation.jl

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -172,6 +172,8 @@ initialize_processor!(
172172
out_data::MM,
173173
input_structure_matrices,
174174
output_structure_matrices,
175+
input_structure_vectors,
176+
output_structure_vectors,
175177
apply_to::AS,
176178
) where {MM <: AbstractMatrix, AS <: AbstractString} = initialize_processor!(cc, in_data, out_data, apply_to)
177179

src/Utilities/decorrelator.jl

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -124,6 +124,7 @@ function initialize_processor!(
124124
dd::Decorrelator,
125125
data::MM,
126126
structure_matrices::Dict{Symbol, <:StructureMatrix},
127+
::Dict{Symbol, <:StructureVector},
127128
) where {MM <: AbstractMatrix}
128129
if length(get_data_mean(dd)) == 0
129130
push!(get_data_mean(dd), vec(mean(data, dims = 2)))

src/Utilities/elementwise_scaler.jl

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -156,7 +156,7 @@ $(TYPEDSIGNATURES)
156156
157157
Computes and populates the `shift` and `scale` fields for the `ElementwiseScaler`
158158
"""
159-
initialize_processor!(es::ElementwiseScaler, data::MM, structure_matrices) where {MM <: AbstractMatrix} =
159+
initialize_processor!(es::ElementwiseScaler, data::MM, structure_matrices, structure_vectors) where {MM <: AbstractMatrix} =
160160
initialize_processor!(es, data)
161161

162162

0 commit comments

Comments
 (0)