Skip to content

Commit 39f7731

Browse files
authored
Do not convert last states to dask DataFrame. (#134)
1 parent 035cb0e commit 39f7731

7 files changed

Lines changed: 13 additions & 39 deletions

File tree

docs/source/changes.rst

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -13,6 +13,8 @@ all releases are available on `Anaconda.org
1313
- :gh:`131` moves the parsing of the virus strain infectiousness factor to the
1414
simulation.
1515
- :gh:`132` sets initialized countdowns to -9,999.
16+
- :gh:`134` changes that the last states are returned as a ``pandas.DataFrame`` and not
17+
as a ``dask.dataframe``.
1618

1719

1820
0.0.9 - 2021-05-28

src/sid/simulate.py

Lines changed: 1 addition & 29 deletions
Original file line numberDiff line numberDiff line change
@@ -644,8 +644,7 @@ def _simulate(
644644
time_series = _prepare_time_series(path, columns_to_keep, states)
645645
results["time_series"] = time_series
646646
if return_last_states:
647-
last_states = _prepare_last_states(path, states)
648-
results["last_states"] = last_states
647+
results["last_states"] = states
649648
if period_outputs:
650649
results["period_outputs"] = evaluated_period_outputs
651650

@@ -1067,33 +1066,6 @@ def _prepare_time_series(output_directory, columns_to_keep, last_states):
10671066
return time_series
10681067

10691068

1070-
def _prepare_last_states(output_directory, last_states):
1071-
"""Prepare the last_states for the simulation results.
1072-
1073-
Args:
1074-
output_directory (pathlib.Path): Path to output directory.
1075-
columns_to_keep (list): List of variables which should be kept.
1076-
last_states (pandas.DataFrame): The states from the last period.
1077-
1078-
Returns:
1079-
dask.dataframe: The DataFrame with the last states
1080-
1081-
1082-
"""
1083-
categoricals = {
1084-
column: last_states[column].cat.categories.shape[0]
1085-
for column in last_states.select_dtypes("category").columns
1086-
}
1087-
1088-
last_states.to_parquet(output_directory / "last_states" / "last_states.parquet")
1089-
last_states = dd.read_parquet(
1090-
output_directory / "last_states" / "last_states.parquet",
1091-
categories=categoricals,
1092-
engine="fastparquet",
1093-
)
1094-
return last_states
1095-
1096-
10971069
def _process_saved_columns(
10981070
saved_columns: Union[None, Dict[str, Union[bool, str, List[str]]]],
10991071
initial_state_columns: List[str],

tests/test_plotting.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -65,7 +65,7 @@ def test_plot_infection_rates_by_contact_models(params, initial_states, tmp_path
6565
result = simulate(params)
6666

6767
time_series = result["time_series"].compute()
68-
last_states = result["last_states"].compute()
68+
last_states = result["last_states"]
6969

7070
for df in [time_series, last_states]:
7171
assert isinstance(df, pd.DataFrame)

tests/test_rapid_tests.py

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -35,7 +35,7 @@ def test_simulate_rapid_tests(params, initial_states, tmp_path):
3535
result = simulate(params)
3636

3737
time_series = result["time_series"].compute()
38-
last_states = result["last_states"].compute()
38+
last_states = result["last_states"]
3939

4040
for df in [time_series, last_states]:
4141
assert isinstance(df, pd.DataFrame)
@@ -77,7 +77,7 @@ def test_simulate_rapid_tests_with_reaction_models(params, initial_states, tmp_p
7777
result = simulate(params)
7878

7979
time_series = result["time_series"].compute()
80-
last_states = result["last_states"].compute()
80+
last_states = result["last_states"]
8181

8282
for df in [time_series, last_states]:
8383
assert isinstance(df, pd.DataFrame)

tests/test_seasonality.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -26,7 +26,7 @@ def test_simulate_a_simple_model(params, initial_states, tmp_path):
2626
result = simulate(params)
2727

2828
time_series = result["time_series"].compute()
29-
last_states = result["last_states"].compute()
29+
last_states = result["last_states"]
3030

3131
for df in [time_series, last_states]:
3232
assert isinstance(df, pd.DataFrame)

tests/test_simulate.py

Lines changed: 5 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -31,7 +31,7 @@ def test_simulate_a_simple_model(params, initial_states, tmp_path):
3131
result = simulate(params)
3232

3333
time_series = result["time_series"].compute()
34-
last_states = result["last_states"].compute()
34+
last_states = result["last_states"]
3535

3636
for df in [time_series, last_states]:
3737
assert isinstance(df, pd.DataFrame)
@@ -55,7 +55,7 @@ def test_resume_a_simulation(params, initial_states, tmp_path):
5555
result = simulate(params)
5656

5757
time_series = result["time_series"].compute()
58-
last_states = result["last_states"].compute()
58+
last_states = result["last_states"]
5959

6060
for df in [time_series, last_states]:
6161
assert isinstance(df, pd.DataFrame)
@@ -77,7 +77,7 @@ def test_resume_a_simulation(params, initial_states, tmp_path):
7777
resumed_result = resumed_simulate(params)
7878

7979
resumed_time_series = resumed_result["time_series"].compute()
80-
resumed_last_states = resumed_result["last_states"].compute()
80+
resumed_last_states = resumed_result["last_states"]
8181

8282
for df in [resumed_time_series, resumed_last_states]:
8383
assert isinstance(df, pd.DataFrame)
@@ -110,7 +110,7 @@ def test_simulate_a_simple_model_without_assort_by(params, initial_states, tmp_p
110110
result = simulate(params)
111111

112112
time_series = result["time_series"].compute()
113-
last_states = result["last_states"].compute()
113+
last_states = result["last_states"]
114114

115115
for df in [time_series, last_states]:
116116
assert isinstance(df, pd.DataFrame)
@@ -352,7 +352,7 @@ def test_skipping_factorization_of_assort_by_variable_works(
352352
result = simulate(params)
353353

354354
time_series = result["time_series"].compute()
355-
last_states = result["last_states"].compute()
355+
last_states = result["last_states"]
356356

357357
assert "group_codes_households" not in time_series
358358
assert "group_codes_households" not in last_states

tests/test_time.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -100,7 +100,7 @@ def test_replace_date_with_period_in_simulation(params, initial_states, tmp_path
100100
result = simulate(params)
101101

102102
time_series = result["time_series"].compute()
103-
last_states = result["last_states"].compute()
103+
last_states = result["last_states"]
104104

105105
for df in [time_series, last_states]:
106106
assert isinstance(df, pd.DataFrame)

0 commit comments

Comments
 (0)