Skip to content

Commit 8d33f8f

Browse files
committed
Adapt tests to new functions
1 parent 90e26e3 commit 8d33f8f

4 files changed

Lines changed: 31 additions & 16 deletions

File tree

src/Emulator.jl

Lines changed: 6 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -174,7 +174,12 @@ function Emulator(
174174

175175
enc_schedule = create_encoder_schedule(encoder_schedule)
176176
(encoded_io_pairs, encoded_input_structure_mat, encoded_output_structure_mat) =
177-
encode_with_schedule!(enc_schedule, input_output_pairs, input_structure_mat, output_structure_mat)
177+
initialize_and_encode_with_schedule!(
178+
enc_schedule,
179+
input_output_pairs,
180+
input_structure_mat,
181+
output_structure_mat,
182+
)
178183

179184
# build the machine learning tool in the encoded space
180185
build_models!(

src/Utilities.jl

Lines changed: 10 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -14,7 +14,14 @@ using ..DataContainers
1414
export get_training_points
1515

1616
export PairedDataContainerProcessor, DataContainerProcessor
17-
export create_encoder_schedule, initialize_and_encode_with_schedule!, encode_with_schedule, decode_with_schedule
17+
export create_encoder_schedule,
18+
initialize_and_encode_with_schedule!,
19+
encode_with_schedule,
20+
decode_with_schedule,
21+
encode_data,
22+
encode_structure_matrix,
23+
decode_data,
24+
decode_structure_matrix
1825

1926

2027

@@ -101,7 +108,7 @@ function _initialize_and_encode_data!(
101108
apply_to::AS,
102109
) where {P <: DataProcessor, AS <: AbstractString}
103110
if proc isa PairedDataContainerProcessor
104-
initialize_processor!(proc, data..., structure_mats..., apply_to)
111+
initialize_processor!(proc, get_data(data)..., structure_mats..., apply_to)
105112
else
106113
input_data, output_data = get_data(data)
107114
input_structure_mat, output_structure_mat = structure_mats
@@ -250,7 +257,7 @@ function encode_with_schedule(
250257
structure_matrix::USorM,
251258
in_or_out::AS,
252259
) where {VV <: AbstractVector, USorM <: Union{UniformScaling, AbstractMatrix}, AS <: AbstractString}
253-
if !(in_or_out ["in", "out"])
260+
if in_or_out ["in", "out"]
254261
bad_in_or_out(in_or_out)
255262
end
256263
processed_structure_matrix = deepcopy(structure_matrix)

test/Emulator/runtests.jl

Lines changed: 6 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -44,11 +44,11 @@ struct MLTester <: Emulators.MachineLearningTool end
4444
@test get_io_pairs(em) == io_pairs
4545
default_encoder = (decorrelate_sample_cov(), "in_and_out") # for these inputs this is the default
4646
enc_sch = create_encoder_schedule(default_encoder)
47-
enc_io_pairs, enc_I_in, enc_I_out = encode_with_schedule!(enc_sch, io_pairs, 1.0 * I(p), 1.0 * I(d))
47+
enc_io_pairs, enc_I_in, enc_I_out = initialize_and_encode_with_schedule!(enc_sch, io_pairs, 1.0 * I(p), 1.0 * I(d))
4848
@test get_encoder_schedule(em)[1][1] == enc_sch[1][1] # inputs: proc
49-
@test get_encoder_schedule(em)[1][3] == enc_sch[1][3] # inputs: apply_to
49+
@test get_encoder_schedule(em)[1][2] == enc_sch[1][2] # inputs: apply_to
5050
@test get_encoder_schedule(em)[2][1] == enc_sch[2][1] # outputs...
51-
@test get_encoder_schedule(em)[2][3] == enc_sch[2][3]
51+
@test get_encoder_schedule(em)[2][2] == enc_sch[2][2]
5252
@test get_data(get_encoded_io_pairs(em)) == get_data(enc_io_pairs)
5353
@test get_data(get_encoded_io_pairs(em)) == get_data(enc_io_pairs)
5454

@@ -91,23 +91,23 @@ struct MLTester <: Emulators.MachineLearningTool end
9191
enc_sch1 = create_encoder_schedule([(decorrelate_sample_cov(), "in"), (decorrelate_structure_mat(), "out")])
9292
enc_sch2 = create_encoder_schedule([(decorrelate_structure_mat(), "in"), (decorrelate_structure_mat(), "out")])
9393
enc_sch3 = create_encoder_schedule([(decorrelate_structure_mat(), "in"), (decorrelate_sample_cov(), "out")])
94-
(_, _, _) = encode_with_schedule!(
94+
(_, _, _) = initialize_and_encode_with_schedule!(
9595
enc_sch1,
9696
io_pairs,
9797
1.0 * I(p),
9898
Σ, # obs noise cov becomes the output structure matrix
9999
)
100100
@test get_encoder_schedule(em1) == enc_sch1
101101

102-
(_, _, _) = encode_with_schedule!(
102+
(_, _, _) = initialize_and_encode_with_schedule!(
103103
enc_sch2,
104104
io_pairs,
105105
4.0 * I(p),
106106
2.0 * I(d), # out_struct_mat overrides obs noise cov
107107
)
108108
@test get_encoder_schedule(em2) == enc_sch2
109109

110-
(_, _, _) = encode_with_schedule!(enc_sch3, io_pairs, Γ, I(d))
110+
(_, _, _) = initialize_and_encode_with_schedule!(enc_sch3, io_pairs, Γ, I(d))
111111
@test get_encoder_schedule(em3) == enc_sch3
112112

113113
end

test/Utilities/runtests.jl

Lines changed: 9 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -143,7 +143,7 @@ end
143143
for (name, sch, ll_flag) in zip(test_names, schedules, lossless)
144144
encoder_schedule = create_encoder_schedule(sch)
145145
(encoded_io_pairs, encoded_prior_cov, encoded_obs_noise_cov) =
146-
encode_with_schedule!(encoder_schedule, io_pairs, prior_cov, obs_noise_cov)
146+
initialize_and_encode_with_schedule!(encoder_schedule, io_pairs, prior_cov, obs_noise_cov)
147147

148148
(decoded_io_pairs, decoded_prior_cov, decoded_obs_noise_cov) =
149149
decode_with_schedule(encoder_schedule, encoded_io_pairs, encoded_prior_cov, encoded_obs_noise_cov)
@@ -293,18 +293,21 @@ end
293293

294294
@test_logs (:warn,) create_encoder_schedule((canonical_correlation(), "bad"))
295295
@test_logs (:warn,) create_encoder_schedule((zscore_scale(), "bad"))
296-
func = x -> (get_inputs(x), get_outputs(x))
297-
bad_encoder_schedule = [(canonical_correlation(), func, "bad")]
298-
@test_throws ArgumentError encode_with_schedule!(bad_encoder_schedule, io_pairs, prior_cov, obs_noise_cov)
296+
bad_encoder_schedule = [(canonical_correlation(), "bad")]
297+
@test_throws ArgumentError initialize_and_encode_with_schedule!(
298+
bad_encoder_schedule,
299+
io_pairs,
300+
prior_cov,
301+
obs_noise_cov,
302+
)
299303
@test_throws ArgumentError decode_with_schedule(bad_encoder_schedule, io_pairs, prior_cov, obs_noise_cov)
300-
@test_throws ArgumentError decode_data(canonical_correlation(), func(io_pairs), "bad")
301304

302305

303306
encoder_schedule = create_encoder_schedule(schedule_builder)
304307

305308
# encode the data using the schedule
306309
(encoded_io_pairs, encoded_prior_cov, encoded_obs_noise_cov) =
307-
encode_with_schedule!(encoder_schedule, io_pairs, prior_cov, obs_noise_cov)
310+
initialize_and_encode_with_schedule!(encoder_schedule, io_pairs, prior_cov, obs_noise_cov)
308311

309312
# decode the data using the schedule
310313
(decoded_io_pairs, decoded_prior_cov, decoded_obs_noise_cov) =

0 commit comments

Comments
 (0)