Skip to content

Commit 59fb512

Browse files
committed
Remove LocalEnsemble.load_all_gen_kw_data()
This commit removes the function, and replaces it with `LocalEnsemble.load_scalars()` as it gradually moves from pandas towards polars.
1 parent a048535 commit 59fb512

8 files changed

Lines changed: 71 additions & 77 deletions

File tree

src/ert/plugins/hook_implementations/workflows/csv_export.py

Lines changed: 4 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -90,7 +90,10 @@ def run(
9090
f"The ensemble '{ensemble.name}' does not have any data!"
9191
)
9292

93-
ensemble_data = ensemble.load_all_gen_kw_data()
93+
ensemble_data = ensemble.load_scalars().to_pandas().set_index("realization")
94+
ensemble_data.columns.name = None
95+
ensemble_data.index.name = "Realization"
96+
ensemble_data = ensemble_data.sort_index(axis=1)
9497

9598
if design_matrix_path is not None:
9699
design_matrix_data = loadDesignMatrix(design_matrix_path)

src/ert/storage/local_ensemble.py

Lines changed: 0 additions & 44 deletions
Original file line numberDiff line numberDiff line change
@@ -783,50 +783,6 @@ def _load_responses_lazy(
783783

784784
return pl.concat(loaded) if loaded else pl.DataFrame().lazy()
785785

786-
def load_all_gen_kw_data(
787-
self,
788-
group: str | None = None,
789-
realization_index: int | None = None,
790-
) -> pd.DataFrame:
791-
"""Loads scalar parameters (GEN_KWs) into a pandas DataFrame
792-
with columns <PARAMETER_GROUP>:<PARAMETER_NAME> and
793-
"Realization" as index.
794-
795-
Parameters
796-
----------
797-
group : str, optional
798-
Name of parameter group to load.
799-
relization_index : int, optional
800-
The realization to load.
801-
802-
Returns
803-
-------
804-
data : DataFrame
805-
A pandas DataFrame containing the GEN_KW data.
806-
807-
Notes
808-
-----
809-
Any provided keys that are not gen_kw will be ignored.
810-
"""
811-
if realization_index is not None:
812-
realizations = np.array([realization_index])
813-
else:
814-
ens_mask = (
815-
self.get_realization_mask_with_responses()
816-
+ self.get_realization_mask_with_parameters()
817-
)
818-
realizations = np.flatnonzero(ens_mask)
819-
820-
df = self.load_scalars(group, realizations)
821-
822-
if df.is_empty():
823-
return pd.DataFrame()
824-
825-
dataframe = df.to_pandas().set_index("realization")
826-
dataframe.columns.name = None
827-
dataframe.index.name = "Realization"
828-
return dataframe.sort_index(axis=1)
829-
830786
@require_write
831787
def save_parameters(
832788
self,

tests/ert/ui_tests/cli/analysis/test_es_update.py

Lines changed: 4 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -53,9 +53,9 @@ def test_that_posterior_has_lower_variance_than_prior():
5353
with open_storage("storage") as storage:
5454
experiment = storage.get_experiment_by_name("es-test")
5555
prior_ensemble = experiment.get_ensemble_by_name("iter-0")
56-
df_default = prior_ensemble.load_all_gen_kw_data()
56+
df_default = prior_ensemble.load_scalars()
5757
posterior_ensemble = experiment.get_ensemble_by_name("iter-1")
58-
df_target = posterior_ensemble.load_all_gen_kw_data()
58+
df_target = posterior_ensemble.load_scalars()
5959

6060
# The std for the ensemble should decrease
6161
assert float(
@@ -68,8 +68,8 @@ def test_that_posterior_has_lower_variance_than_prior():
6868
# generalized variance for the parameters.
6969
assert (
7070
0
71-
< np.linalg.det(df_target.cov().to_numpy())
72-
< np.linalg.det(df_default.cov().to_numpy())
71+
< np.linalg.det(df_target.to_pandas().cov().to_numpy())
72+
< np.linalg.det(df_default.to_pandas().cov().to_numpy())
7373
)
7474

7575

tests/ert/ui_tests/cli/test_cli.py

Lines changed: 10 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -540,7 +540,11 @@ def test_that_es_mda_on_poly_case_matches_snapshot(snapshot):
540540
experiment = storage.get_experiment_by_name("es-mda")
541541
for iter_nr in range(4):
542542
ensemble = experiment.get_ensemble_by_name(f"iter-{iter_nr}")
543-
data.append(ensemble.load_all_gen_kw_data())
543+
ensemble_data = ensemble.load_scalars().to_pandas().set_index("realization")
544+
ensemble_data.columns.name = None
545+
ensemble_data.index.name = "Realization"
546+
ensemble_data = ensemble_data.sort_index(axis=1)
547+
data.append(ensemble_data)
544548
result = pd.concat(
545549
data,
546550
keys=[f"iter-{iter_}" for iter_ in range(len(data))],
@@ -576,7 +580,11 @@ def test_that_enif_on_poly_case_matches_snapshot(snapshot):
576580
experiment = storage.get_experiment_by_name("enif")
577581
for iter_nr in range(2):
578582
ensemble = experiment.get_ensemble_by_name(f"iter-{iter_nr}")
579-
data.append(ensemble.load_all_gen_kw_data())
583+
ensemble_data = ensemble.load_scalars().to_pandas().set_index("realization")
584+
ensemble_data.columns.name = None
585+
ensemble_data.index.name = "Realization"
586+
ensemble_data = ensemble_data.sort_index(axis=1)
587+
data.append(ensemble_data)
580588
result = pd.concat(
581589
data,
582590
keys=[f"iter-{i}" for i in range(len(data))],

tests/ert/ui_tests/cli/test_update.py

Lines changed: 4 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -232,12 +232,10 @@ def test_update_lowers_generalized_variance_or_deactivates_observations(
232232
if success:
233233
with open_storage("storage") as storage:
234234
experiment = storage.get_experiment_by_name("experiment")
235-
prior = experiment.get_ensemble_by_name("iter-0").load_all_gen_kw_data()
236-
posterior = experiment.get_ensemble_by_name(
237-
"iter-1"
238-
).load_all_gen_kw_data()
235+
prior = experiment.get_ensemble_by_name("iter-0").load_scalars()
236+
posterior = experiment.get_ensemble_by_name("iter-1").load_scalars()
239237

240238
assert (
241-
np.linalg.det(posterior.cov().to_numpy())
242-
<= np.linalg.det(prior.cov().to_numpy()) + 0.001
239+
np.linalg.det(posterior.to_pandas().cov().to_numpy())
240+
<= np.linalg.det(prior.to_pandas().cov().to_numpy()) + 0.001
243241
)

tests/ert/ui_tests/gui/test_csv_export.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -68,7 +68,7 @@ def verify_exported_content(file_name, gui, ensemble_select):
6868
for name in ensemble_names:
6969
experiment = gui.notifier.storage.get_experiment_by_name("es_mda")
7070
ensemble = experiment.get_ensemble_by_name(name)
71-
gen_kw_data = ensemble.load_all_gen_kw_data()
71+
gen_kw_data = ensemble.load_scalars().to_pandas()
7272

7373
facade = LibresFacade.from_config_file("poly.ert")
7474
misfit_data = facade.load_all_misfit_data(ensemble)

tests/ert/ui_tests/gui/test_full_manual_update_workflow.py

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -88,9 +88,9 @@ def test_manual_analysis_workflow(ensemble_experiment_has_run, qtbot):
8888
10,
8989
)
9090

91-
df_prior = ensemble_prior.load_all_gen_kw_data()
91+
df_prior = ensemble_prior.load_scalars().to_pandas()
9292
ensemble_posterior = experiment.get_ensemble_by_name("iter-0_1")
93-
df_posterior = ensemble_posterior.load_all_gen_kw_data()
93+
df_posterior = ensemble_posterior.load_scalars().to_pandas()
9494

9595
# Making sure measured data works with failed realizations
9696
MeasuredData(experiment.get_ensemble_by_name("iter-0"), ["POLY_OBS"])

tests/ert/unit_tests/storage/test_local_storage.py

Lines changed: 46 additions & 17 deletions
Original file line numberDiff line numberDiff line change
@@ -1021,40 +1021,69 @@ def test_load_gen_kw_not_sorted(storage, tmpdir, snapshot):
10211021
)
10221022

10231023
sample_prior(ensemble, range(ensemble_size), random_seed=1234)
1024-
1025-
data = ensemble.load_all_gen_kw_data()
1024+
data = ensemble.load_scalars().to_pandas().set_index("realization")
1025+
data.columns.name = None
1026+
data.index.name = "Realization"
1027+
data = data.sort_index(axis=1)
10261028
snapshot.assert_match(data.round(12).to_csv(), "gen_kw_unsorted")
10271029

10281030

10291031
def test_gen_kw_collector(snake_oil_default_storage, snapshot):
1030-
data = snake_oil_default_storage.load_all_gen_kw_data()
1032+
data = snake_oil_default_storage.load_scalars().to_pandas().set_index("realization")
1033+
data.columns.name = None
1034+
data.index.name = "Realization"
1035+
data = data.sort_index(axis=1)
10311036
snapshot.assert_match(data.round(6).to_csv(), "gen_kw_collector.csv")
10321037

10331038
with pytest.raises(KeyError):
10341039
# realization 60:
10351040
_ = data.loc[60]
10361041

1037-
data = snake_oil_default_storage.load_all_gen_kw_data(
1038-
"SNAKE_OIL_PARAM",
1039-
)[["SNAKE_OIL_PARAM:OP1_PERSISTENCE", "SNAKE_OIL_PARAM:OP1_OFFSET"]]
1042+
data = (
1043+
snake_oil_default_storage.load_scalars(
1044+
"SNAKE_OIL_PARAM",
1045+
)
1046+
.to_pandas()
1047+
.set_index("realization")
1048+
)
1049+
data.columns.name = None
1050+
data.index.name = "Realization"
1051+
data = data.sort_index(axis=1)
1052+
data = data[["SNAKE_OIL_PARAM:OP1_PERSISTENCE", "SNAKE_OIL_PARAM:OP1_OFFSET"]]
10401053
snapshot.assert_match(data.round(6).to_csv(), "gen_kw_collector_2.csv")
10411054

10421055
with pytest.raises(KeyError):
10431056
_ = data["SNAKE_OIL_PARAM:OP1_DIVERGENCE_SCALE"]
10441057

10451058
realization_index = 3
1046-
data = snake_oil_default_storage.load_all_gen_kw_data(
1047-
"SNAKE_OIL_PARAM",
1048-
realization_index=realization_index,
1049-
)["SNAKE_OIL_PARAM:OP1_PERSISTENCE"]
1059+
data = (
1060+
snake_oil_default_storage.load_scalars(
1061+
"SNAKE_OIL_PARAM",
1062+
realization_index=realization_index,
1063+
)
1064+
.to_pandas()
1065+
.set_index("realization")
1066+
)
1067+
data.columns.name = None
1068+
data.index.name = "Realization"
1069+
data = data.sort_index(axis=1)
1070+
data = data["SNAKE_OIL_PARAM:OP1_PERSISTENCE"]
10501071
snapshot.assert_match(data.round(6).to_csv(), "gen_kw_collector_3.csv")
10511072

10521073
non_existing_realization_index = 150
10531074
with pytest.raises((IndexError, KeyError)):
1054-
_ = snake_oil_default_storage.load_all_gen_kw_data(
1055-
"SNAKE_OIL_PARAM",
1056-
realization_index=non_existing_realization_index,
1057-
)["SNAKE_OIL_PARAM:OP1_PERSISTENCE"]
1075+
data = (
1076+
snake_oil_default_storage.load_scalars(
1077+
"SNAKE_OIL_PARAM",
1078+
realization_index=non_existing_realization_index,
1079+
)
1080+
.to_pandas()
1081+
.set_index("realization")
1082+
)
1083+
data.columns.name = None
1084+
data.index.name = "Realization"
1085+
data = data.sort_index(axis=1)
1086+
data = data["SNAKE_OIL_PARAM:OP1_PERSISTENCE"]
10581087

10591088

10601089
def test_keyword_type_checks(snake_oil_default_storage):
@@ -1083,12 +1112,12 @@ def test_data_fetching_missing_key(snake_oil_case):
10831112
empty_case = experiment.create_ensemble(name="new_case", ensemble_size=25)
10841113

10851114
data = [
1086-
empty_case.load_all_gen_kw_data("nokey", None),
1115+
empty_case.load_scalars("nokey", None),
10871116
]
10881117

10891118
for dataframe in data:
1090-
assert isinstance(dataframe, DataFrame)
1091-
assert dataframe.empty
1119+
assert isinstance(dataframe, pl.DataFrame)
1120+
assert dataframe.is_empty()
10921121

10931122

10941123
@dataclass

0 commit comments

Comments
 (0)