From 0302ba961ef668782dbc1c63a01bbcb9386564da Mon Sep 17 00:00:00 2001 From: Kilian Volmer <13285635+kilianvolmer@users.noreply.github.com> Date: Wed, 30 Sep 2026 10:47:25 +0200 Subject: [PATCH 1/6] NEW: Code and examples and docs for the timeSeries plotting function --- docs/source/python/m-plot.rst | 81 ++- .../python/m-simulation_model_usage.rst | 9 + pycode/examples/plot/plotSimulationResults.py | 128 +++++ .../examples/simulation/ode_secir_groups.py | 106 ++-- .../examples/simulation/ode_secir_mobility.py | 77 +-- .../examples/simulation/ode_secir_simple.py | 60 +-- pycode/memilio-plot/README.md | 9 +- pycode/memilio-plot/memilio/plot/__init__.py | 4 + .../memilio/plot/plotTimeSeries.py | 479 ++++++++++++++++++ pycode/memilio-plot/pyproject.toml | 3 +- 10 files changed, 780 insertions(+), 176 deletions(-) create mode 100644 pycode/examples/plot/plotSimulationResults.py create mode 100644 pycode/memilio-plot/memilio/plot/plotTimeSeries.py diff --git a/docs/source/python/m-plot.rst b/docs/source/python/m-plot.rst index e13a5572aa..3286e4d539 100644 --- a/docs/source/python/m-plot.rst +++ b/docs/source/python/m-plot.rst @@ -19,7 +19,7 @@ Dependencies Required python packages: - pandas>=1.2.2 -- matplotlib +- matplotlib>=3.6 - numpy>=1.22, !=1.25.* - openpyxl - xlrd @@ -32,3 +32,82 @@ Required python packages: - h5py - imageio - datetime + +Plotting simulation results +--------------------------- + +The module ``memilio.plot.plotTimeSeries`` provides a standard plot for ``TimeSeries`` objects as returned by the +simulations of the :doc:`MEmilio Python bindings `: + +.. code-block:: python + + from datetime import date + + import numpy as np + + from memilio.plot.plotTimeSeries import plot_time_series + from memilio.simulation import AgeGroup, Damping + from memilio.simulation.osecir import InfectionState, Model, simulate + + # ODE SECIR model with two age groups + groups = ['0-19', '20+'] + populations = [15000, 68000] + num_groups = len(groups) + + model = Model(num_groups) + for i, population in enumerate(populations): + group = AgeGroup(i) + # time spent in the compartments (days) + model.parameters.TimeExposed[group] = 3.2 + model.parameters.TimeInfectedNoSymptoms[group] = 2. + model.parameters.TimeInfectedSymptoms[group] = 6. + model.parameters.TimeInfectedSevere[group] = 12. + model.parameters.TimeInfectedCritical[group] = 8. + # transmission and transition probabilities + model.parameters.TransmissionProbabilityOnContact[group] = 1.0 + model.parameters.RelativeTransmissionNoSymptoms[group] = 0.67 + model.parameters.RiskOfInfectionFromSymptomatic[group] = 0.25 + model.parameters.MaxRiskOfInfectionFromSymptomatic[group] = 0.5 + model.parameters.RecoveredPerInfectedNoSymptoms[group] = 0.09 + model.parameters.SeverePerInfectedSymptoms[group] = 0.2 + model.parameters.CriticalPerSevere[group] = 0.25 + model.parameters.DeathsPerCritical[group] = 0.3 + # initial populations + model.populations[group, InfectionState.Exposed] = 100 + model.populations[group, InfectionState.InfectedNoSymptoms] = 50 + model.populations[group, InfectionState.InfectedSymptoms] = 50 + model.populations[group, InfectionState.InfectedSevere] = 20 + model.populations[group, InfectionState.InfectedCritical] = 10 + model.populations[group, InfectionState.Recovered] = 10 + model.populations.set_difference_from_group_total_AgeGroup( + (group, InfectionState.Susceptible), population) + + # contacts: one contact per person and day, reduced by 90% from day 30 on + model.parameters.ContactPatterns.cont_freq_mat[0].baseline = np.ones( + (num_groups, num_groups)) + model.parameters.ContactPatterns.cont_freq_mat.add_damping(Damping( + coeffs=np.ones((num_groups, num_groups)) * 0.9, t=30., level=0, type=0)) + model.check_constraints() + + result = simulate(0, 100, 0.1, model) + ax = plot_time_series( + result, labels=InfectionState.values(), groups=groups, + select=['Exposed', 'InfectedSymptoms', 'Dead'], + start_date=date(2020, 3, 1), title='ODE SECIR simulation') + ax.figure.savefig('secir.pdf') + +- ``labels`` names the compartments. It accepts strings or the ``InfectionState.values()`` of a model. Without labels, + the elements are named ``C1``, ``C2``, ... . +- ``groups`` (the number or the names of, e.g., age groups) tells the function how the elements of the ``TimeSeries`` + are arranged. By default, the compartments are summed over all groups; with ``sum_groups=False`` one line per group + and compartment is drawn. +- ``select`` restricts the plot to some compartments, ``log_scale`` switches to a logarithmic axis and ``start_date`` + turns the time axis into calendar dates. Colors are fixed per compartment, so a selection does not change them. +- The function returns a matplotlib ``Axes``. Pass ``ax=`` to draw into an existing figure (e.g., a subplot grid) and + use ``ax.figure.savefig`` to save the figure. The legend is placed right of the axes, so create your own figures with + ``layout='constrained'`` or save them with ``bbox_inches='tight'``. + +``time_series_to_dataframe`` converts a ``TimeSeries`` into a tidy pandas ``DataFrame`` with the columns ``Time``, +``Group``, ``Compartment`` and ``Value`` (and ``Date`` if a start date is given) for use with other libraries such as +seaborn, plotnine or altair. A complete example is given in +`examples/plot/plotSimulationResults.py `_. diff --git a/docs/source/python/m-simulation_model_usage.rst b/docs/source/python/m-simulation_model_usage.rst index ef74891927..bb3f5725f4 100755 --- a/docs/source/python/m-simulation_model_usage.rst +++ b/docs/source/python/m-simulation_model_usage.rst @@ -259,6 +259,15 @@ pythonic interface. Now you can use the usual data handling options and make use of the easy visualization tools that are part of Python. Some plotting functions specific to MEmilio and created as part of the project are combined in the :doc:`MEmilio Plot Package `. +For a quick look at the result, it provides a standard plot for ``TimeSeries`` objects: + +.. code-block:: python + + from memilio.plot.plotTimeSeries import plot_time_series + + ax = plot_time_series( + result, labels=oseir.InfectionState.values(), groups=num_groups) + ax.figure.savefig('result.pdf') Additional resources --------------------- diff --git a/pycode/examples/plot/plotSimulationResults.py b/pycode/examples/plot/plotSimulationResults.py new file mode 100644 index 0000000000..e557275b02 --- /dev/null +++ b/pycode/examples/plot/plotSimulationResults.py @@ -0,0 +1,128 @@ +############################################################################# +# Copyright (C) 2020-2026 MEmilio +# +# Authors: Kilian Volmer +# +# Contact: Martin J. Kuehn +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +############################################################################# +""" +Example demonstrating the standard TimeSeries plot of the MEmilio plot +package on the result of an ODE SEIR simulation with two age groups. +""" +import argparse +import os +from datetime import date + +import matplotlib.pyplot as plt +import numpy as np + +from memilio.plot.plotTimeSeries import (plot_time_series, + time_series_to_dataframe) +from memilio.simulation import AgeGroup, Damping +from memilio.simulation.oseir import InfectionState, Model, simulate + +AGE_GROUPS = ['0-19', '20+'] +START_DATE = date(2020, 3, 1) + + +def run_ode_seir_simulation(days=100, dt=0.1): + """ Runs the ODE SEIR model with two age groups. + + :param days: Number of days to simulate. (Default value = 100) + :param dt: Initial time step. (Default value = 0.1) + :returns: Simulation result as TimeSeries. + """ + group_populations = [15000, 68000] + num_groups = len(group_populations) + model = Model(num_groups) + + for i, population in enumerate(group_populations): + group = AgeGroup(i) + model.parameters.TimeExposed[group] = 5.2 + model.parameters.TimeInfected[group] = 6. + model.parameters.TransmissionProbabilityOnContact[group] = 1. * i + model.populations[group, InfectionState.Exposed] = 100 + model.populations[group, InfectionState.Infected] = 50 + model.populations[group, InfectionState.Recovered] = 10 + model.populations.set_difference_from_group_total_AgeGroup( + (group, InfectionState.Susceptible), population) + + model.parameters.ContactPatterns.cont_freq_mat[0].baseline = np.ones( + (num_groups, num_groups)) + model.parameters.ContactPatterns.cont_freq_mat[0].minimum = np.zeros( + (num_groups, num_groups)) + model.parameters.ContactPatterns.cont_freq_mat.add_damping(Damping( + coeffs=np.ones((num_groups, num_groups)) * 0.9, t=30.0, level=0, + type=0)) + + model.check_constraints() + return simulate(0, days, dt, model) + + +def plot_results(result, output_path='.', show_plot=False): + """ Creates the standard plots of the simulation result. + + :param result: TimeSeries returned by the simulation. + :param output_path: Directory the figures are written to. + (Default value = '.') + :param show_plot: Whether to show the figures interactively. + (Default value = False) + """ + # All compartments, summed over the age groups. The compartment names + # are taken from the InfectionState enum of the model. + ax = plot_time_series( + result, labels=InfectionState.values(), groups=AGE_GROUPS, + title='ODE SEIR simulation') + ax.figure.savefig( + os.path.join(output_path, 'seir_compartments.png'), dpi=150) + + # Two panels in one figure: a selection of compartments per age group + # on a date axis, and a selection on a logarithmic axis. + fig, axes = plt.subplots(1, 2, figsize=(12, 4.5), layout='constrained') + plot_time_series( + result, labels=InfectionState.values(), groups=AGE_GROUPS, + sum_groups=False, select=['Exposed', 'Infected'], + start_date=START_DATE, ax=axes[0], + title='Exposed and infected per age group') + plot_time_series( + result, labels=InfectionState.values(), groups=AGE_GROUPS, + select=['Exposed', 'Infected', 'Recovered'], log_scale=True, + start_date=START_DATE, ax=axes[1], title='Logarithmic scale') + fig.savefig(os.path.join(output_path, 'seir_selection.pdf')) + + # The same data as a tidy data frame, e.g. for other plotting libraries. + df = time_series_to_dataframe( + result, labels=InfectionState.values(), groups=AGE_GROUPS, + start_date=START_DATE) + print(df.head()) + print(df.pivot_table(index='Date', columns='Compartments', values='Values', + aggfunc='sum', observed=False).tail()) + + if show_plot: + plt.show() + plt.close('all') + + +if __name__ == '__main__': + arg_parser = argparse.ArgumentParser( + 'plotSimulationResults', + description='Plots the result of an ODE SEIR simulation with the ' + 'standard TimeSeries plot of the MEmilio plot package.') + arg_parser.add_argument('-p', '--show_plot', action='store_true', + help='Show the figures interactively.') + arg_parser.add_argument('-o', '--output_path', default='.', + help='Directory the figures are written to.') + args = arg_parser.parse_args() + plot_results(run_ode_seir_simulation(), args.output_path, args.show_plot) diff --git a/pycode/examples/simulation/ode_secir_groups.py b/pycode/examples/simulation/ode_secir_groups.py index a20f1562f4..e85718b1a7 100644 --- a/pycode/examples/simulation/ode_secir_groups.py +++ b/pycode/examples/simulation/ode_secir_groups.py @@ -19,31 +19,27 @@ ############################################################################# import argparse import os -from datetime import date, datetime +from datetime import date import matplotlib.pyplot as plt import numpy as np -import pandas as pd -from memilio.simulation import AgeGroup, ContactMatrix, Damping, UncertainContactMatrix -from memilio.simulation.osecir import Index_InfectionState +from memilio.plot.plotTimeSeries import plot_time_series +from memilio.simulation import AgeGroup, Damping from memilio.simulation.osecir import InfectionState as State -from memilio.simulation.osecir import (Model, Simulation, - interpolate_simulation_result, simulate) +from memilio.simulation.osecir import (Model, interpolate_simulation_result, + simulate) def run_ode_secir_groups_simulation(show_plot=True): """Runs the c++ ODE SECIHURD model using mulitple age groups and plots the results - :param show_plot: (Default value = True) + :param show_plot: Whether to show the figures interactively. + (Default value = True) """ - # Define Comartment names - compartments = [ - 'Susceptible', 'Exposed', 'InfectedNoSymptoms', 'InfectedSymptoms', - 'InfectedSevere', 'InfectedCritical', 'Recovered', 'Dead'] # Define age Groups groups = ['0-4', '5-14', '15-34', '35-59', '60-79', '80+'] # Define population of age groups @@ -55,11 +51,10 @@ def run_ode_secir_groups_simulation(show_plot=True): start_year = 2019 dt = 0.1 num_groups = len(groups) - num_compartments = len(compartments) # set contact frequency matrix data_dir = os.path.join(os.path.dirname( - __file__), "..", "..", "..", "data") + __file__), "..", "..", "..", "data", "Germany") baseline_contact_matrix0 = os.path.join( data_dir, "contacts/baseline_home.txt") baseline_contact_matrix1 = os.path.join( @@ -133,79 +128,38 @@ def run_ode_secir_groups_simulation(show_plot=True): # interpolate results result = interpolate_simulation_result(result) - # print(result.get_last_value()) - num_time_points = result.get_num_time_points() - result_array = result.as_ndarray() - t = result_array[0, :] - group_data = np.transpose(result_array[1:, :]) - - # sum over all groups - data = np.zeros((num_time_points, num_compartments)) - for i in range(num_groups): - data += group_data[:, i * num_compartments: (i + 1) * num_compartments] - - # Plot Results - datelist = np.array( - pd.date_range( - datetime(start_year, start_month, start_day), - periods=days, freq='D').strftime('%m-%d').tolist()) - - tick_range = (np.arange(int(days / 10) + 1) * 10) - tick_range[-1] -= 1 - fig, ax = plt.subplots() - ax.plot(t, data[:, 0], label='#Susceptible') - ax.plot(t, data[:, 1], label='#Exposed') - ax.plot(t, data[:, 2], label='#InfectedNoSymptoms') - ax.plot(t, data[:, 3], label='#InfectedSymptoms') - ax.plot(t, data[:, 4], label='#Hospitalzed') - ax.plot(t, data[:, 5], label='#InfectedCritical') - ax.plot(t, data[:, 6], label='#Recovered') - ax.plot(t, data[:, 7], label='#Dead') - ax.set_title("ODE SECIR simulation results (entire population)") - ax.set_xticks(tick_range) - ax.set_xticklabels(datelist[tick_range], rotation=45) - ax.legend() - fig.tight_layout - fig.savefig('osecir_by_compartments.pdf') - - # plot dynamics in each comparment by age group - fig, ax = plt.subplots(4, 2, figsize=(12, 15)) - - for i, title in zip(range(num_compartments), compartments): - - for j, group in enumerate(groups): - ax[int(np.floor(i / 2)), int(i % 2)].plot(t, - group_data[:, j*num_compartments+i], label=group) - - ax[int(np.floor(i / 2)), int(i % 2)].set_title(title, fontsize=10) - ax[int(np.floor(i / 2)), int(i % 2)].legend() - - ax[int(np.floor(i / 2)), int(i % 2)].set_xticks(tick_range) - ax[int(np.floor(i / 2)), int(i % 2) - ].set_xticklabels(datelist[tick_range], rotation=45) - plt.subplots_adjust(hspace=0.5, bottom=0.1, top=0.9) + start_date = date(start_year, start_month, start_day) + + # Plot results summed over all age groups, one line per compartment. The + # compartment names are taken from the InfectionState enum of the model. + ax = plot_time_series( + result, labels=State.values(), groups=groups, start_date=start_date, + title='ODE SECIR simulation results (entire population)') + ax.figure.savefig('osecir_by_compartments.pdf') + + # One panel per compartment with one line per age group. + fig, axes = plt.subplots(5, 2, figsize=(16, 18), layout='constrained') + for state, ax in zip(State.values(), axes.flat): + plot_time_series( + result, labels=State.values(), groups=groups, sum_groups=False, + select=state, start_date=start_date, ax=ax, title=state.name) fig.suptitle( 'ODE SECIR simulation results by age group in each compartment') fig.savefig('osecir_age_groups_in_compartments.pdf') - fig, ax = plt.subplots(4, 2, figsize=(12, 15)) - for i, title in zip(range(num_compartments), compartments): - ax[int(np.floor(i / 2)), int(i % 2)].plot(t, data[:, i]) - ax[int(np.floor(i / 2)), int(i % 2)].set_title(title, fontsize=10) - - ax[int(np.floor(i / 2)), int(i % 2)].set_xticks(tick_range) - ax[int(np.floor(i / 2)), int(i % 2) - ].set_xticklabels(datelist[tick_range], rotation=45) - plt.subplots_adjust(hspace=0.5, bottom=0.1, top=0.9) + # One panel per compartment, summed over all age groups. + fig, axes = plt.subplots(5, 2, figsize=(16, 18), layout='constrained') + for state, ax in zip(State.values(), axes.flat): + plot_time_series( + result, labels=State.values(), groups=groups, select=state, + start_date=start_date, ax=ax, title=state.name) fig.suptitle( 'ODE SECIR simulation results by compartment (entire population)') fig.savefig('osecir_all_parts.pdf') if show_plot: plt.show() - plt.close() - - # return data + plt.close('all') if __name__ == "__main__": diff --git a/pycode/examples/simulation/ode_secir_mobility.py b/pycode/examples/simulation/ode_secir_mobility.py index c07ce277fe..8c14049a73 100644 --- a/pycode/examples/simulation/ode_secir_mobility.py +++ b/pycode/examples/simulation/ode_secir_mobility.py @@ -19,11 +19,12 @@ ############################################################################# import argparse -import numpy as np import matplotlib.pyplot as plt +import numpy as np import memilio.simulation as mio import memilio.simulation.osecir as osecir +from memilio.plot.plotTimeSeries import plot_time_series def run_ode_secir_mobility_simulation(plot_results=True): @@ -115,58 +116,36 @@ def run_ode_secir_mobility_simulation(plot_results=True): region1_result = osecir.interpolate_simulation_result( sim.graph.get_node(1).property.result) - if (plot_results): - results = [region0_result.as_ndarray(), region1_result.as_ndarray()] - t = results[0][0, :] - tick_range = (np.arange(int((len(t) - 1) / 10) + 1) * 10) - tick_range[-1] -= 1 - - fig, ax = plt.subplots(figsize=(10, 6)) - for idx, result_region in enumerate(results): - region_label = f'Region {idx}' - ax.plot(t, result_region[1, :], - label=f'{region_label} - #Susceptible') - ax.plot(t, result_region[2, :], label=f'{region_label} - #Exposed') - ax.plot(t, result_region[3, :] + result_region[4, :], - label=f'{region_label} - #InfectedNoSymptoms') - ax.plot(t, result_region[5, :] + result_region[6, :], - label=f'{region_label} - #InfectedSymptoms') - ax.plot(t, result_region[7, :], - label=f'{region_label} - #Hospitalzed') - ax.plot(t, result_region[8, :], - label=f'{region_label} - #InfectedCritical') - ax.plot(t, result_region[9, :], - label=f'{region_label} - #Recovered') - ax.plot(t, result_region[10, :], label=f'{region_label} - #Dead') - - ax.set_title( - "ODE SECIR simulation results for both regions (entire population)") - ax.set_xticks(tick_range) - ax.legend(loc='upper right', bbox_to_anchor=(1, 0.6)) - plt.yscale('log') - fig.tight_layout + if plot_results: + region_results = [region0_result, region1_result] + region_labels = ['Region 0', 'Region 1'] + + # All compartments of each region on a logarithmic axis. + fig, axes = plt.subplots(1, 2, figsize=(16, 5), layout='constrained') + for region_result, region_label, ax in zip( + region_results, region_labels, axes): + plot_time_series( + region_result, labels=osecir.InfectionState.values(), + ax=ax, title=region_label) + fig.suptitle('ODE SECIR simulation results for both regions') fig.savefig('osecir_mobility_by_compartments.pdf') - fig, ax = plt.subplots(5, 2, figsize=(12, 15)) - compartments = [ - 'Susceptible', 'Exposed', 'InfectedNoSymptoms', - 'InfectedNoSymptomsConfirmed', 'InfectedSymptoms', - 'InfectedSymptomsConfirmed', 'InfectedSevere', 'InfectedCritical', - 'Recovered', 'Dead'] - num_compartments = len(compartments) - - for i, title in zip(range(num_compartments), compartments): - ax[int(np.floor(i / 2)), int(i % 2)].plot(t, - results[0][i+1, :], label="Region 0") - ax[int(np.floor(i / 2)), int(i % 2)].plot(t, - results[1][i+1, :], label="Region 1") - ax[int(np.floor(i / 2)), int(i % 2)].set_title(title, fontsize=10) - ax[int(np.floor(i / 2)), int(i % 2)].legend() - ax[int(np.floor(i / 2)), int(i % 2)].set_xticks(tick_range) - - plt.subplots_adjust(hspace=0.5, bottom=0.1, top=0.9) + # Stack the results of the regions into one array with the elements + # ordered region by region (the interpolated results share the same + # time points), so that the regions can be passed as groups. + results = np.vstack([region0_result.as_ndarray(), + region1_result.as_ndarray()[1:]]) + + # One panel per compartment with one line per region. + fig, axes = plt.subplots(5, 2, figsize=(14, 16), layout='constrained') + for state, ax in zip(osecir.InfectionState.values(), axes.flat): + plot_time_series( + results, labels=osecir.InfectionState.values(), + groups=region_labels, sum_groups=False, select=state, ax=ax, + title=state.name) fig.suptitle('Simulation results for each region in each compartment') fig.savefig('osecir_region_results_compartments.pdf') + plt.close('all') if __name__ == "__main__": diff --git a/pycode/examples/simulation/ode_secir_simple.py b/pycode/examples/simulation/ode_secir_simple.py index 8ba5de6c32..e389edd5fe 100644 --- a/pycode/examples/simulation/ode_secir_simple.py +++ b/pycode/examples/simulation/ode_secir_simple.py @@ -18,31 +18,27 @@ # limitations under the License. ############################################################################# import argparse -from datetime import date, datetime +from datetime import date import matplotlib.pyplot as plt import numpy as np -import pandas as pd -from memilio.simulation import AgeGroup, ContactMatrix, Damping, UncertainContactMatrix -from memilio.simulation.osecir import Index_InfectionState +from memilio.plot.plotTimeSeries import plot_time_series +from memilio.simulation import AgeGroup, Damping from memilio.simulation.osecir import InfectionState as State -from memilio.simulation.osecir import (Model, Simulation, - interpolate_simulation_result, simulate) +from memilio.simulation.osecir import (Model, interpolate_simulation_result, + simulate) def run_ode_secir_simulation(show_plot=True): """Runs the c++ ODE SECIHURD model using one age group and plots the results - :param show_plot: (Default value = True) + :param show_plot: Whether to show the figure interactively. + (Default value = True) """ - # Define Comartment names - compartments = [ - 'Susceptible', 'Exposed', 'InfectedNoSymptoms', 'InfectedSymptoms', - 'InfectedSevere', 'InfectedCritical', 'Recovered', 'Dead'] # Define population of age groups populations = [83000] @@ -52,7 +48,6 @@ def run_ode_secir_simulation(show_plot=True): start_year = 2019 dt = 0.1 num_groups = 1 - num_compartments = len(compartments) # Initialize Parameters model = Model(1) @@ -112,39 +107,14 @@ def run_ode_secir_simulation(show_plot=True): result = interpolate_simulation_result(result) print(result.get_last_value()) - num_time_points = result.get_num_time_points() - result_array = result.as_ndarray() - t = result_array[0, :] - group_data = np.transpose(result_array[1:, :]) - - # sum over all groups - data = np.zeros((num_time_points, num_compartments)) - for i in range(num_groups): - data += group_data[:, i * num_compartments: (i + 1) * num_compartments] - - # Plot Results - datelist = np.array( - pd.date_range( - datetime(start_year, start_month, start_day), - periods=days, freq='D').strftime('%m-%d').tolist()) - - tick_range = (np.arange(int(days / 10) + 1) * 10) - tick_range[-1] -= 1 - fig, ax = plt.subplots() - ax.plot(t, data[:, 0], label='#Susceptible') - ax.plot(t, data[:, 1], label='#Exposed') - ax.plot(t, data[:, 2], label='#Carrying') - ax.plot(t, data[:, 3], label='#InfectedSymptoms') - ax.plot(t, data[:, 4], label='#Hospitalzed') - ax.plot(t, data[:, 5], label='#InfectedCritical') - ax.plot(t, data[:, 6], label='#Recovered') - ax.plot(t, data[:, 7], label='#Died') - ax.set_title("ODE SECIR model simulation") - ax.set_xticks(tick_range) - ax.set_xticklabels(datelist[tick_range], rotation=45) - ax.legend() - fig.tight_layout - fig.savefig('osecir_simple.pdf') + + # Plot results: one line per compartment (infection state) on a date + # axis. The names are taken from the InfectionState enum of the model. + ax = plot_time_series( + result, labels=State.values(), + start_date=date(start_year, start_month, start_day), + title="ODE SECIR model simulation") + ax.figure.savefig('osecir_simple.pdf') if show_plot: plt.show() diff --git a/pycode/memilio-plot/README.md b/pycode/memilio-plot/README.md index 0c27bfcc16..b71a999891 100644 --- a/pycode/memilio-plot/README.md +++ b/pycode/memilio-plot/README.md @@ -18,7 +18,9 @@ Introduction ------------ This package provides modules and scripts to plot epidemiological or simulation data as returned -by other packages of the MEmilio software. +by other packages of the MEmilio software. The module ``plotTimeSeries`` provides a standard plot +for ``TimeSeries`` objects (e.g., simulation results of the ``memilio-simulation`` package), ``plotMap`` +visualizes regional data on maps and ``createGIF`` animates map plots over time. Installation ------------ @@ -47,15 +49,14 @@ Dependencies Required python packages: - pandas>=1.2.2 -- matplotlib -- numpy>=1.22,<1.25 +- matplotlib>=3.6 +- numpy>=1.22,!=1.25.* - openpyxl - xlrd - requests - pyxlsb - wget - folium -- matplotlib - mapclassify - geopandas - h5py diff --git a/pycode/memilio-plot/memilio/plot/__init__.py b/pycode/memilio-plot/memilio/plot/__init__.py index 68d983e289..419e854091 100644 --- a/pycode/memilio-plot/memilio/plot/__init__.py +++ b/pycode/memilio-plot/memilio/plot/__init__.py @@ -20,4 +20,8 @@ """ Functions to plot and visualize map data and simulation trajectories. + +The module ``plotTimeSeries`` provides a standard plot for ``TimeSeries`` +objects, ``plotMap`` plots regional data on maps and ``createGIF`` animates +map plots over time. """ diff --git a/pycode/memilio-plot/memilio/plot/plotTimeSeries.py b/pycode/memilio-plot/memilio/plot/plotTimeSeries.py new file mode 100644 index 0000000000..eac7ef3339 --- /dev/null +++ b/pycode/memilio-plot/memilio/plot/plotTimeSeries.py @@ -0,0 +1,479 @@ +############################################################################# +# Copyright (C) 2020-2026 MEmilio +# +# Authors: Kilian Volmer +# +# Contact: Martin J. Kuehn +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +############################################################################# +""" +:strong:`plotTimeSeries.py` +Standard plot of a MEmilio ``TimeSeries`` such as the result of a simulation. + +The functions in this module accept a ``memilio.simulation.TimeSeries`` (or +any object with an ``as_ndarray()`` method returning an array of shape +``(1 + num_elements, num_time_points)`` whose first row holds the time +points) as well as a plain 2-D numpy array of that shape. The module does +not import ``memilio.simulation`` itself, so it can be used with the plot +package alone. +""" +from __future__ import annotations + +import math +import warnings +from collections.abc import Sequence + +import matplotlib.dates as mdates +import matplotlib.pyplot as plt +import numpy as np +import pandas as pd +from matplotlib.ticker import FuncFormatter, NullFormatter + +# Categorical palette in a fixed, colorblind-safe order. Compartment i is +# always drawn with _COLORS[i % 8]; from the ninth compartment on, the line +# style changes instead of introducing new hues. +_COLORS = ['#2a78d6', '#eb6834', '#1baf7a', '#eda100', + '#e87ba4', '#008300', '#4a3aa7', '#e34948'] +_LINESTYLES = ['-', '--', '-.', ':'] + +_GRID_COLOR = '#e1e0d9' +_AXIS_COLOR = '#c3c2b7' +_TICK_COLOR = '#898781' +_TEXT_COLOR = '#52514e' +_TITLE_COLOR = '#0b0b0b' + +# Legends with more rows than this are split into several columns. +_MAX_LEGEND_ROWS = 20 + + +def _to_array(time_series) -> np.ndarray: + """ Converts the input into a copied 2-D array (time points in row 0). + + :param time_series: TimeSeries-like object (with ``as_ndarray()``) or + a 2-D array-like of shape (1 + num_elements, num_time_points). + :returns: Copied numpy array of the same shape. + """ + if hasattr(time_series, 'as_ndarray'): + # as_ndarray() of the bindings is a view into C++ memory; copy it so + # the plot data stays valid independent of the TimeSeries object. + data = np.array(time_series.as_ndarray(), dtype=float, copy=True) + else: + data = np.array(time_series, dtype=float, copy=True) + if data.ndim != 2 or data.shape[0] < 2: + raise ValueError( + 'Expected a TimeSeries or a 2-D array of shape ' + '(1 + num_elements, num_time_points) with the time points in ' + f'the first row, got shape {data.shape}.') + return data + + +def _names(items) -> list[str]: + """ Converts a sequence of strings or objects with a ``name`` attribute + (e.g. members of an InfectionState enum) into a list of strings. + + :param items: Iterable of strings or named objects, or a non-iterable + object with a ``values()`` method returning such an iterable (for + example the ``InfectionState`` enums of the bindings). + :returns: List of names. + """ + if isinstance(items, (str, bytes)): + raise TypeError('Expected a sequence of names, not a single string.') + if not hasattr(items, '__iter__'): + if not callable(getattr(items, 'values', None)): + raise TypeError( + f'Expected a sequence of names, got {type(items).__name__}.') + items = items.values() + return [str(getattr(item, 'name', item)) for item in items] + + +def _prepare(time_series, labels, groups) -> tuple[ + np.ndarray, np.ndarray, list[str], list[str] | None]: + """ Validates the input and reshapes the values by group and compartment. + + :param time_series: See :func:`plot_time_series`. + :param labels: See :func:`plot_time_series`. + :param groups: See :func:`plot_time_series`. + :returns: Tuple of the time points (shape (num_time_points,)), the values + (shape (num_groups, num_compartments, num_time_points)), the + compartment names and the group names (None if no groups are given). + """ + data = _to_array(time_series) + times = data[0] + values = data[1:] + num_elements = values.shape[0] + + if groups is None: + group_names = None + elif isinstance(groups, (int, np.integer)) and not isinstance( + groups, bool): + if groups < 1: + raise ValueError('groups must be a positive number of groups.') + group_names = [f'Group {i}' for i in range(int(groups))] + else: + group_names = _names(groups) + if not group_names: + raise ValueError('groups must not be empty.') + num_groups = 1 if group_names is None else len(group_names) + + if labels is None: + if num_elements % num_groups != 0: + raise ValueError( + f'The number of elements ({num_elements}) is not divisible ' + f'by the number of groups ({num_groups}).') + # Same default names as TimeSeries.print_table. + labels = [f'C{i + 1}' for i in range(num_elements // num_groups)] + else: + labels = _names(labels) + if len(labels) * num_groups != num_elements: + if groups is None: + raise ValueError( + f'Got {len(labels)} labels for {num_elements} elements. ' + 'Provide one label per element, or use the groups ' + 'argument if the elements are resolved by (age) groups.') + raise ValueError( + f'{len(labels)} labels times {num_groups} groups does not ' + f'match the number of elements ({num_elements}).') + if len(set(labels)) != len(labels): + raise ValueError('labels must be unique.') + if group_names is not None and len(set(group_names)) != num_groups: + raise ValueError('groups must be unique.') + + # Elements are ordered group-major: g0c0, g0c1, ..., g1c0, g1c1, ... + values = values.reshape(num_groups, len(labels), values.shape[1]) + return times, values, labels, group_names + + +def _dates(times: np.ndarray, start_date) -> pd.DatetimeIndex: + """ Converts simulation times (in days) into dates, rounded to seconds. + + :param times: Time points in days. + :param start_date: Date (``datetime.date``, ``datetime.datetime``, + ``pandas.Timestamp`` or ISO string) that corresponds to time 0. + :returns: DatetimeIndex with one entry per time point. + """ + dates = pd.Timestamp(start_date) + pd.to_timedelta(times, unit='D') + return dates.round('s') + + +def time_series_to_dataframe( + time_series, labels=None, groups=None, start_date=None +) -> pd.DataFrame: + """ Converts a TimeSeries into a tidy (long-form) pandas DataFrame. + + The frame has one row per time point and element with the columns + ``Time``, ``Date`` (only if ``start_date`` is given), ``Group`` (only if + ``groups`` is given), ``Compartment`` and ``Value``. ``Compartment`` and + ``Group`` are ordered categoricals in the order of the TimeSeries, so the + frame can be used directly with grammar-of-graphics libraries such as + seaborn, plotnine or altair. A wide table (one column per compartment, + summed over groups) is obtained by + ``df.pivot_table(index='Time', columns='Compartment', values='Value', + aggfunc='sum')``. + + :param time_series: ``memilio.simulation.TimeSeries`` (or any object with + ``as_ndarray()``) or a 2-D array of shape + (1 + num_elements, num_time_points) with the time points in row 0. + :param labels: Names of the compartments (elements). Either strings or + objects with a ``name`` attribute, e.g. ``InfectionState.values()`` + of a model. If None, the elements are named ``C1``, ``C2``, ... + (Default value = None) + :param groups: Number of groups or their names (e.g. age groups) if the + elements are resolved by groups. The elements are expected in the + order of the bindings, i.e. all compartments of the first group, + then all compartments of the second group, and so on. + (Default value = None) + :param start_date: Date corresponding to time 0. If given, a ``Date`` + column (rounded to full seconds) is added. (Default value = None) + :returns: Long-form DataFrame. + """ + times, values, labels, group_names = _prepare( + time_series, labels, groups) + num_groups, num_compartments, num_time_points = values.shape + + frame = {'Time': np.tile(times, num_groups * num_compartments)} + if start_date is not None: + frame['Date'] = _dates(frame['Time'], start_date) + if group_names is not None: + frame['Groups'] = pd.Categorical.from_codes( + np.repeat(np.arange(num_groups), + num_compartments * num_time_points), + categories=group_names, ordered=True) + frame['Compartments'] = pd.Categorical.from_codes( + np.tile(np.repeat(np.arange(num_compartments), num_time_points), + num_groups), + categories=labels, ordered=True) + frame['Values'] = values.reshape(-1) + return pd.DataFrame(frame) + + +def _resolve_select(select, labels: list[str]) -> list[int]: + """ Resolves the compartments to plot into sorted, unique indices. + + :param select: None (all), a single name/index or a sequence of names + or indices. + :param labels: All compartment names. + :returns: Sorted list of compartment indices. + """ + if select is None: + return list(range(len(labels))) + if isinstance(select, (str, int, np.integer)) or hasattr(select, 'name'): + select = [select] + indices = set() + for item in select: + if isinstance(item, (int, np.integer)) and not isinstance(item, bool): + if not 0 <= item < len(labels): + raise ValueError(f'Compartment index {item} out of range.') + indices.add(int(item)) + else: + name = str(getattr(item, 'name', item)) + if name not in labels: + raise ValueError( + f'Unknown compartment {name!r}. ' + f'Available: {", ".join(labels)}.') + indices.add(labels.index(name)) + return sorted(indices) + + +def _color_for(index: int, name: str, colors) -> str: + """ Color of a line, keyed by the position of its compartment (or group). + + :param index: Index of the compartment (or group). + :param name: Name of the compartment (or group). + :param colors: None, a sequence of colors (by index) or a dictionary + mapping names to colors. + :returns: Matplotlib color. + """ + if isinstance(colors, dict): + if name in colors: + return colors[name] + elif colors is not None and index < len(colors): + return colors[index] + return _COLORS[index % len(_COLORS)] + + +def _style_for(index: int) -> str: + """ Line style of a line, changing once the palette has been used up. + + :param index: Index of the compartment (or group). + :returns: Matplotlib line style. + """ + return _LINESTYLES[(index // len(_COLORS)) % len(_LINESTYLES)] + + +def _format_count(value, _position=None) -> str: + """ Tick formatter with thousands separators for large numbers. + + :param value: Tick value. + :param _position: Tick position (unused). (Default value = None) + :returns: Formatted tick label. + """ + if abs(value) >= 1000: + return f'{value:,.0f}' + if 0 < abs(value) < 1e-3: + return f'{value:.0e}'.replace('e-0', 'e-') + return f'{value:g}' + + +def _apply_style(ax, log_scale: bool, use_dates: bool): + """ Applies the style to an axes: recessive grid and spines, + readable tick labels. + + :param ax: Matplotlib axes. + :param log_scale: Whether the y axis is logarithmic. + :param use_dates: Whether the x axis shows dates. + """ + for side in ('top', 'right'): + ax.spines[side].set_visible(False) + for side in ('left', 'bottom'): + ax.spines[side].set_color(_AXIS_COLOR) + ax.tick_params(colors=_TICK_COLOR, labelcolor=_TEXT_COLOR) + ax.xaxis.label.set_color(_TEXT_COLOR) + ax.yaxis.label.set_color(_TEXT_COLOR) + ax.grid(True, axis='y', color=_GRID_COLOR, linewidth=0.8, linestyle='-') + ax.grid(False, axis='x') + ax.set_axisbelow(True) + ax.margins(x=0) + + if log_scale: + ax.set_yscale('log', nonpositive='clip') + ax.yaxis.set_minor_formatter(NullFormatter()) + ax.yaxis.set_major_formatter(FuncFormatter(_format_count)) + + if use_dates: + locator = mdates.AutoDateLocator() + ax.xaxis.set_major_locator(locator) + ax.xaxis.set_major_formatter(mdates.ConciseDateFormatter(locator)) + + +def _draw_legend(ax, num_lines: int): + """ Draws the legend right of the axes, split into columns if it has + many entries. + + :param ax: Matplotlib axes. + :param num_lines: Number of legend entries. + """ + ncol = max(1, math.ceil(num_lines / _MAX_LEGEND_ROWS)) + ax.legend(frameon=False, ncol=ncol, loc='upper left', + bbox_to_anchor=(1.01, 1.0), borderaxespad=0.0) + + +def _collect_lines(values: np.ndarray, labels: list[str], + group_names: list[str] | None, selected: list[int], + sum_groups: bool, colors) -> list[tuple]: + """ Determines the lines to draw. + + :param values: Values of shape (num_groups, num_compartments, + num_time_points). + :param labels: Compartment names. + :param group_names: Group names or None. + :param selected: Indices of the compartments to draw. + :param sum_groups: Whether to sum the compartments over the groups. + :param colors: See :func:`plot_time_series`. + :returns: List of (color, line style, label, y values) per line. + """ + lines = [] + for index in selected: + name = labels[index] + if group_names is None or sum_groups: + lines.append((_color_for(index, name, colors), _style_for(index), + name, values[:, index, :].sum(axis=0))) + elif len(selected) == 1: + # A single compartment: the groups take the colors. + for group_index, group_name in enumerate(group_names): + lines.append((_color_for(group_index, group_name, colors), + _style_for(group_index), + f'{name} ({group_name})', + values[group_index, index, :])) + else: + for group_index, group_name in enumerate(group_names): + lines.append((_color_for(index, name, colors), + _LINESTYLES[group_index % len(_LINESTYLES)], + f'{name} ({group_name})', + values[group_index, index, :])) + styles = [(color, style) for color, style, _, _ in lines] + if len(set(styles)) < len(styles): + warnings.warn( + 'Some lines share color and line style and cannot be told ' + 'apart. Use select to plot fewer compartments or groups.', + stacklevel=3) + return lines + + +def plot_time_series( + time_series, labels=None, *, groups=None, sum_groups: bool = True, + select=None, ax=None, title: str | None = None, + xlabel: str | None = None, ylabel: str = 'Number of individuals', + log_scale: bool = False, start_date=None, + legend: bool = True, + colors: Sequence[str] | dict[str, str] | None = None, + figsize: tuple[float, float] = (8, 4.5), **plot_kwargs): + """ Plots the elements of a TimeSeries over time, one line per compartment. + + Example:: + + from memilio.simulation.oseir import InfectionState, simulate + from memilio.plot.plotTimeSeries import plot_time_series + + # model: oseir.Model with two age groups + result = simulate(0, 100, 0.1, model) + ax = plot_time_series(result, labels=InfectionState.values(), + groups=['0-19', '20+'], title='ODE SEIR') + ax.figure.savefig('seir.pdf') + + Every compartment has a fixed color. For more than eight compartments, + the colors are reused with a different line style; consider ``select`` + to plot only the compartments of interest. + + :param time_series: ``memilio.simulation.TimeSeries`` (or any object with + ``as_ndarray()``) or a 2-D array of shape + (1 + num_elements, num_time_points) with the time points in row 0. + :param labels: Names of the compartments (elements). Either strings or + objects with a ``name`` attribute, e.g. ``InfectionState.values()`` + of a model. If None, the elements are named ``C1``, ``C2``, ... + (Default value = None) + :param groups: Number of groups or their names (e.g. age groups) if the + elements are resolved by groups, i.e. if the TimeSeries has + ``len(labels) * num_groups`` elements ordered group by group. + (Default value = None) + :param sum_groups: If True, the compartments are summed over all groups + and one line per compartment is drawn. If False, one line per group + and compartment is drawn, labeled ``' ()'``. The + groups are distinguished by line style, or by color if a single + compartment is selected. A warning is issued if the lines cannot be + told apart. (Default value = True) + :param select: Compartment(s) to plot, given by name or index. None plots + all compartments. Colors are assigned by the position of a compartment + in the TimeSeries, so a selection does not change the colors. + (Default value = None) + :param ax: Matplotlib axes to draw into. If None, a new figure with + constrained layout is created. When drawing into your own figure, + create it with ``layout='constrained'`` or save it with + ``bbox_inches='tight'`` so that the legend right of the axes is not + cut off. (Default value = None) + :param title: Title of the plot. (Default value = None) + :param xlabel: Label of the x axis. If None, 'Time (days)' or 'Date' is + used. (Default value = None) + :param ylabel: Label of the y axis. + (Default value = 'Number of individuals') + :param log_scale: Use a logarithmic y axis. (Default value = False) + :param start_date: Date corresponding to time 0 (``datetime.date`` or + similar). If given, the x axis shows dates. (Default value = None) + :param legend: Whether to draw a legend. It is placed right of the axes. + (Default value = True) + :param colors: Sequence of colors indexed by compartment, or a dictionary + mapping compartment names to colors. Compartments not covered use the + default palette. If a single compartment is plotted per group, the + colors refer to the groups instead. (Default value = None) + :param figsize: Size of the figure in inches if a new figure is created. + (Default value = (8, 4.5)) + :param plot_kwargs: Additional keyword arguments passed to + ``matplotlib.axes.Axes.plot`` for every line, e.g. ``linewidth``. + :returns: The matplotlib axes containing the plot. Use ``ax.figure`` to + access and save the figure. + """ + times, values, labels, group_names = _prepare( + time_series, labels, groups) + selected = _resolve_select(select, labels) + x = _dates(times, start_date) if start_date is not None else times + + lines = _collect_lines(values, labels, group_names, selected, + sum_groups, colors) + + if ax is None: + _, ax = plt.subplots(figsize=figsize, layout='constrained') + + line_kwargs = {'linewidth': 2.0, 'solid_capstyle': 'round', + 'solid_joinstyle': 'round'} + for alias, full_name in (('lw', 'linewidth'), ('ls', 'linestyle'), + ('c', 'color')): + if alias in plot_kwargs: + plot_kwargs[full_name] = plot_kwargs.pop(alias) + plot_kwargs.pop('label', None) + line_kwargs.update(plot_kwargs) + for color, style, name, y in lines: + ax.plot(x, y, label=name, + **{'color': color, 'linestyle': style, **line_kwargs}) + + _apply_style(ax, log_scale, start_date is not None) + + if xlabel is None: + xlabel = 'Date' if start_date is not None else 'Time (days)' + ax.set_xlabel(xlabel) + ax.set_ylabel(ylabel) + + if title is not None: + ax.set_title(title, fontweight='bold', color=_TITLE_COLOR) + if legend and lines: + _draw_legend(ax, len(lines)) + return ax diff --git a/pycode/memilio-plot/pyproject.toml b/pycode/memilio-plot/pyproject.toml index 9838362755..a6acc3f05d 100644 --- a/pycode/memilio-plot/pyproject.toml +++ b/pycode/memilio-plot/pyproject.toml @@ -16,7 +16,8 @@ maintainers = [ dependencies = [ "setuptools>=68", "pandas>=1.2.2", - "matplotlib", + # layout engines (constrained layout of figures) were added in 3.6 + "matplotlib>=3.6", # smaller numpy versions cause a security issue, 1.25 does not work together with pyfakefs "numpy>=1.22,!=1.25.*", "openpyxl", From 3dd7301c9ceceeb9b5041a9c325e2ed7d5b15b35 Mon Sep 17 00:00:00 2001 From: Kilian Volmer <13285635+kilianvolmer@users.noreply.github.com> Date: Wed, 30 Sep 2026 13:39:13 +0200 Subject: [PATCH 2/6] FIX docstrings --- .../memilio-plot/memilio/plot/plotTimeSeries.py | 16 ++++++++-------- 1 file changed, 8 insertions(+), 8 deletions(-) diff --git a/pycode/memilio-plot/memilio/plot/plotTimeSeries.py b/pycode/memilio-plot/memilio/plot/plotTimeSeries.py index eac7ef3339..c991347d60 100644 --- a/pycode/memilio-plot/memilio/plot/plotTimeSeries.py +++ b/pycode/memilio-plot/memilio/plot/plotTimeSeries.py @@ -83,7 +83,7 @@ def _names(items) -> list[str]: (e.g. members of an InfectionState enum) into a list of strings. :param items: Iterable of strings or named objects, or a non-iterable - object with a ``values()`` method returning such an iterable (for + object with a ``values()`` method returning such an iterable (for example the ``InfectionState`` enums of the bindings). :returns: List of names. """ @@ -172,13 +172,13 @@ def time_series_to_dataframe( """ Converts a TimeSeries into a tidy (long-form) pandas DataFrame. The frame has one row per time point and element with the columns - ``Time``, ``Date`` (only if ``start_date`` is given), ``Group`` (only if - ``groups`` is given), ``Compartment`` and ``Value``. ``Compartment`` and - ``Group`` are ordered categoricals in the order of the TimeSeries, so the - frame can be used directly with grammar-of-graphics libraries such as - seaborn, plotnine or altair. A wide table (one column per compartment, - summed over groups) is obtained by - ``df.pivot_table(index='Time', columns='Compartment', values='Value', + ``Time``, ``Date`` (only if ``start_date`` is given), ``Groups`` (only + if ``groups`` is given), ``Compartments`` and ``Values``. + ``Compartments`` and ``Groups`` are ordered categoricals in the order of + the TimeSeries, so the frame can be used directly with + grammar-of-graphics libraries such as seaborn, plotnine or altair. A wide + table (one column per compartment, summed over groups) is obtained by + ``df.pivot_table(index='Time', columns='Compartments', values='Values', aggfunc='sum')``. :param time_series: ``memilio.simulation.TimeSeries`` (or any object with From 5ddb32a0dbb64caefe0da3e4bf474b3c7a11c04c Mon Sep 17 00:00:00 2001 From: Kilian Volmer <13285635+kilianvolmer@users.noreply.github.com> Date: Wed, 30 Sep 2026 14:23:51 +0200 Subject: [PATCH 3/6] Add tests for plot TimeSeries --- .../tests/test_plot_plotTimeSeries.py | 451 ++++++++++++++++++ 1 file changed, 451 insertions(+) create mode 100644 pycode/memilio-plot/tests/test_plot_plotTimeSeries.py diff --git a/pycode/memilio-plot/tests/test_plot_plotTimeSeries.py b/pycode/memilio-plot/tests/test_plot_plotTimeSeries.py new file mode 100644 index 0000000000..6b856c246d --- /dev/null +++ b/pycode/memilio-plot/tests/test_plot_plotTimeSeries.py @@ -0,0 +1,451 @@ +############################################################################# +# Copyright (C) 2020-2026 MEmilio +# +# Authors: Kilian Volmer +# +# Contact: Martin J. Kuehn +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +############################################################################# +import datetime as dt +import enum +import unittest +import warnings + +import matplotlib +import matplotlib.pyplot as plt +import numpy as np +import pandas as pd +from matplotlib.layout_engine import ConstrainedLayoutEngine +from matplotlib.ticker import NullFormatter + +import memilio.plot.plotTimeSeries as pts + +try: + import memilio.simulation as mio +except ImportError: + # The plot package is tested without the compiled bindings. + mio = None + +matplotlib.use('Agg') + + +class _StubTimeSeries: + """Minimal stand-in for memilio.simulation.TimeSeries.""" + + def __init__(self, data): + self._data = np.asarray(data, dtype=float) + + def as_ndarray(self): + """Returns the data as array (1 + num_elements, num_time_points).""" + return self._data + + +class _State(enum.Enum): + """Iterable enum with named members, like Python's own enums.""" + Susceptible = 0 + Infected = 1 + Recovered = 2 + + +class _NamedMember: + """Object with a name attribute, like a member of a bindings enum.""" + + def __init__(self, name): + self.name = name + + +class _BindingsLikeEnum: + """Mimics the InfectionState enums of the bindings: not iterable, but + with a static values() method returning members with a name.""" + + @staticmethod + def values(): + """Returns the members of the enum.""" + return [_NamedMember('Susceptible'), _NamedMember('Infected'), + _NamedMember('Recovered')] + + +class TestPlotTimeSeries(unittest.TestCase): + """Tests for plot_time_series and time_series_to_dataframe.""" + + def setUp(self): + """Creates a TimeSeries with 2 groups x 3 compartments = 6 elements + in group-major order, where element i has the values (i+1)*10 + t.""" + self.times = np.array([0.0, 0.5, 1.0, 2.0, 3.0]) + self.num_compartments = 3 + self.values = np.array( + [(i + 1) * 10.0 + self.times for i in range(6)]) + self.data = np.vstack([self.times, self.values]) + self.ts = _StubTimeSeries(self.data) + self.labels = ['S', 'I', 'R'] + self.groups = ['0-4', '5+'] + + def tearDown(self): + """Closes all figures created by the test.""" + plt.close('all') + + def _labels(self, ax): + """Returns the labels of the lines in drawing order.""" + return [line.get_label() for line in ax.get_lines()] + + def _legend_texts(self, ax): + """Returns the entries of the legend.""" + return [t.get_text() for t in ax.get_legend().get_texts()] + + def _legend_is_outside(self, ax): + """Returns whether the legend is anchored right of the axes.""" + ax.figure.canvas.draw() + anchor = ax.get_legend().get_bbox_to_anchor() + return anchor.x0 >= ax.get_window_extent().x1 + + # ---- input handling ------------------------------------------------- + + def test_default_labels(self): + """Test that elements without labels are named C1, C2, ... as in + TimeSeries.print_table, and that the legend lists them.""" + ax = pts.plot_time_series(self.ts) + expected = ['C1', 'C2', 'C3', 'C4', 'C5', 'C6'] + self.assertEqual(self._labels(ax), expected) + self.assertEqual(self._legend_texts(ax), expected) + + def test_given_labels(self): + """Test that one line per element is drawn with the given labels and + the unchanged time points and values, for a TimeSeries and for a + plain array.""" + labels = ['a', 'b', 'c', 'd', 'e', 'f'] + for time_series in (self.ts, self.data): + ax = pts.plot_time_series(time_series, labels) + self.assertEqual(self._labels(ax), labels) + for line, row in zip(ax.get_lines(), self.values): + np.testing.assert_array_equal(line.get_xdata(), self.times) + np.testing.assert_array_equal(line.get_ydata(), row) + + def test_labels_from_enum(self): + """Test that labels can be taken from the members of an enum, from + a bindings-like enum class or from its values().""" + data = self.data[:4] # 3 elements + expected = ['Susceptible', 'Infected', 'Recovered'] + for labels in (_State, _BindingsLikeEnum, _BindingsLikeEnum.values()): + ax = pts.plot_time_series(_StubTimeSeries(data), labels) + self.assertEqual(self._labels(ax), expected) + + def test_input_is_copied(self): + """Test that the plot keeps its data if the input array is modified + afterwards, since as_ndarray() of the bindings is a view.""" + data = self.data.copy() + ax = pts.plot_time_series(_StubTimeSeries(data), groups=2) + data[1:] = 0.0 + self.assertGreater(ax.get_lines()[0].get_ydata().sum(), 0.0) + + def test_invalid_input(self): + """Test that wrong shapes, mismatching numbers of labels and groups, + duplicate names and wrong argument types raise errors.""" + cases = [ + ('1-D input', {'time_series': np.array([0.0, 1.0, 2.0])}, + ValueError), + ('no elements', {'time_series': np.zeros((1, 3))}, ValueError), + ('too few labels', {'labels': ['only', 'two']}, ValueError), + ('labels times groups mismatch', + {'labels': self.labels, 'groups': 3}, ValueError), + ('group names mismatch', + {'labels': self.labels, 'groups': ['x']}, ValueError), + ('elements not divisible by groups', {'groups': 4}, ValueError), + ('zero groups', {'groups': 0}, ValueError), + ('empty groups', {'groups': []}, ValueError), + ('duplicate labels', + {'labels': ['a', 'a', 'b', 'c', 'd', 'e']}, ValueError), + ('labels as single string', {'labels': 'abcdef'}, TypeError), + ('groups as bool', {'groups': True}, TypeError), + ('groups as float', {'groups': 2.0}, TypeError), + ] + for description, kwargs, error in cases: + with self.subTest(description): + kwargs = {'time_series': self.ts, **kwargs} + with self.assertRaises(error): + pts.plot_time_series(**kwargs) + + # ---- groups --------------------------------------------------------- + + def test_sum_groups(self): + """Test that the compartments are summed over the groups by + default.""" + ax = pts.plot_time_series(self.ts, self.labels, groups=self.groups) + self.assertEqual(self._labels(ax), self.labels) + for c, line in enumerate(ax.get_lines()): + expected = self.values[c] + self.values[c + self.num_compartments] + np.testing.assert_allclose(line.get_ydata(), expected) + + def test_groups_by_count(self): + """Test that groups given as a number are named Group 0, Group 1, + ... and that default labels are per compartment.""" + ax = pts.plot_time_series(self.ts, groups=2) + self.assertEqual(self._labels(ax), ['C1', 'C2', 'C3']) + ax = pts.plot_time_series(self.ts, groups=2, sum_groups=False) + self.assertEqual(self._labels(ax)[:2], + ['C1 (Group 0)', 'C1 (Group 1)']) + + def test_per_group_lines(self): + """Test one line per group and compartment: the color follows the + compartment and the line style the group; without groups, + sum_groups has no effect.""" + ax = pts.plot_time_series( + self.ts, self.labels, groups=self.groups, sum_groups=False) + lines = ax.get_lines() + self.assertEqual( + self._labels(ax), + ['S (0-4)', 'S (5+)', 'I (0-4)', 'I (5+)', 'R (0-4)', 'R (5+)']) + self.assertEqual(lines[0].get_color(), lines[1].get_color()) + self.assertNotEqual(lines[0].get_color(), lines[2].get_color()) + self.assertNotEqual(lines[0].get_linestyle(), lines[1].get_linestyle()) + self.assertEqual(lines[0].get_linestyle(), lines[2].get_linestyle()) + np.testing.assert_allclose(lines[1].get_ydata(), self.values[3]) + + ax = pts.plot_time_series(self.ts, sum_groups=False) + self.assertEqual(len(ax.get_lines()), 6) + + def test_single_compartment_per_group(self): + """Test that the groups take the colors (default or given) if a + single compartment is plotted per group.""" + ax = pts.plot_time_series( + self.ts, self.labels, groups=self.groups, sum_groups=False, + select='I') + lines = ax.get_lines() + self.assertEqual(self._labels(ax), ['I (0-4)', 'I (5+)']) + self.assertEqual(lines[0].get_color(), '#2a78d6') + self.assertEqual(lines[1].get_color(), '#eb6834') + self.assertEqual(lines[0].get_linestyle(), lines[1].get_linestyle()) + np.testing.assert_allclose(lines[0].get_ydata(), self.values[1]) + np.testing.assert_allclose(lines[1].get_ydata(), self.values[4]) + + ax = pts.plot_time_series( + self.ts, self.labels, groups=self.groups, sum_groups=False, + select='I', colors=['black', 'gray']) + self.assertEqual([line.get_color() for line in ax.get_lines()], + ['black', 'gray']) + + def test_ambiguous_lines_warn(self): + """Test that a warning is issued if lines share color and line + style (more groups than line styles), and not otherwise.""" + data = np.vstack([self.times, np.ones((10, len(self.times)))]) + with self.assertWarns(UserWarning): + pts.plot_time_series(data, ['a', 'b'], groups=5, sum_groups=False) + with warnings.catch_warnings(): + warnings.simplefilter('error') + pts.plot_time_series(data, ['a', 'b'], groups=5) + pts.plot_time_series( + data, ['a', 'b'], groups=5, sum_groups=False, select='a') + + # ---- selection and colors ------------------------------------------ + + def test_select(self): + """Test selecting compartments by name, index or enum member, in the + order of the TimeSeries and without changing their colors.""" + ax_all = pts.plot_time_series(self.ts, self.labels, groups=2) + colors = {line.get_label(): line.get_color() + for line in ax_all.get_lines()} + ax = pts.plot_time_series( + self.ts, self.labels, groups=2, select=['R', 'I']) + self.assertEqual(self._labels(ax), ['I', 'R']) + for line in ax.get_lines(): + self.assertEqual(line.get_color(), colors[line.get_label()]) + ax = pts.plot_time_series(self.ts, self.labels, groups=2, select=2) + self.assertEqual(self._labels(ax), ['R']) + ax = pts.plot_time_series( + self.ts, self.labels, groups=2, select=_NamedMember('I')) + self.assertEqual(self._labels(ax), ['I']) + ax = pts.plot_time_series( + self.ts, self.labels, groups=2, select=[0, 'S', np.int64(1)]) + self.assertEqual(self._labels(ax), ['S', 'I']) + with self.assertRaises(ValueError): + pts.plot_time_series(self.ts, self.labels, groups=2, select='X') + with self.assertRaises(ValueError): + pts.plot_time_series(self.ts, self.labels, groups=2, select=[7]) + + def test_custom_colors(self): + """Test overriding the colors by a dictionary of compartment names + or by a sequence indexed by compartment.""" + ax = pts.plot_time_series( + self.ts, self.labels, groups=2, colors={'I': 'black'}) + self.assertEqual([line.get_color() for line in ax.get_lines()], + ['#2a78d6', 'black', '#1baf7a']) + ax = pts.plot_time_series( + self.ts, self.labels, groups=2, colors=['red', 'green']) + self.assertEqual([line.get_color() for line in ax.get_lines()], + ['red', 'green', '#1baf7a']) + + def test_more_than_eight_compartments(self): + """Test that from the ninth compartment on, the colors are reused + with a different line style.""" + data = np.vstack([self.times, np.ones((10, len(self.times)))]) + ax = pts.plot_time_series(data) + lines = ax.get_lines() + self.assertEqual(len(lines), 10) + self.assertEqual(lines[8].get_color(), lines[0].get_color()) + self.assertEqual(lines[0].get_linestyle(), '-') + self.assertEqual(lines[8].get_linestyle(), '--') + self.assertEqual(len(self._legend_texts(ax)), 10) + + # ---- legend --------------------------------------------------------- + + def test_legend_right_of_axes(self): + """Test that the legend is placed right of the axes, also for a + single line and for a user-provided axes.""" + ax = pts.plot_time_series(self.ts, groups=2) + self.assertTrue(self._legend_is_outside(ax)) + ax = pts.plot_time_series(self.ts, self.labels, groups=2, select='I') + self.assertEqual(self._legend_texts(ax), ['I']) + self.assertTrue(self._legend_is_outside(ax)) + _, axes = plt.subplots(2, 2) + ax = pts.plot_time_series(self.ts, ax=axes[0, 0]) + self.assertEqual(len(self._legend_texts(ax)), 6) + self.assertTrue(self._legend_is_outside(ax)) + + def test_no_legend(self): + """Test that no legend is drawn if it is disabled or if there are no + lines.""" + ax = pts.plot_time_series(self.ts, groups=2, legend=False) + self.assertIsNone(ax.get_legend()) + ax = pts.plot_time_series(self.ts, groups=2, select=[]) + self.assertEqual(len(ax.get_lines()), 0) + self.assertIsNone(ax.get_legend()) + + # ---- axes, labels and options -------------------------------------- + + def test_new_figure(self): + """Test that a new figure of the given size with constrained layout + and default axis labels is created if no axes is given.""" + ax = pts.plot_time_series(self.ts, groups=2, figsize=(10, 3)) + np.testing.assert_array_equal( + ax.figure.get_size_inches(), [10.0, 3.0]) + self.assertIsInstance( + ax.figure.get_layout_engine(), ConstrainedLayoutEngine) + self.assertEqual(ax.get_xlabel(), 'Time (days)') + self.assertEqual(ax.get_ylabel(), 'Number of individuals') + self.assertEqual(ax.get_title(), '') + + def test_given_axes_and_labels(self): + """Test drawing into a given axes with a title, axis labels and line + keyword arguments.""" + _, ax_in = plt.subplots() + ax = pts.plot_time_series( + self.ts, self.labels, groups=2, ax=ax_in, title='T', + xlabel='x', ylabel='y', linewidth=4.0) + self.assertIs(ax, ax_in) + self.assertEqual(ax.get_title(), 'T') + self.assertEqual(ax.get_xlabel(), 'x') + self.assertEqual(ax.get_ylabel(), 'y') + self.assertEqual(ax.get_yscale(), 'linear') + self.assertFalse(ax.spines['top'].get_visible()) + for line in ax.get_lines(): + self.assertEqual(line.get_linewidth(), 4.0) + + def test_plot_kwargs(self): + """Test that keyword arguments and their matplotlib aliases are + passed to all lines and override the defaults, except the label.""" + ax = pts.plot_time_series( + self.ts, groups=2, lw=3.0, color='black', ls=':', label='x') + for line in ax.get_lines(): + self.assertEqual(line.get_linewidth(), 3.0) + self.assertEqual(line.get_color(), 'black') + self.assertEqual(line.get_linestyle(), ':') + self.assertEqual(self._labels(ax), ['C1', 'C2', 'C3']) + + def test_log_scale(self): + """Test the logarithmic y axis with the count formatter on the major + ticks and no labels on the minor ticks.""" + ax = pts.plot_time_series(self.ts, groups=2, log_scale=True) + self.assertEqual(ax.get_yscale(), 'log') + self.assertEqual(ax.yaxis.get_major_formatter()(1000.0), '1,000') + self.assertIsInstance(ax.yaxis.get_minor_formatter(), NullFormatter) + + def test_start_date(self): + """Test that the time points are converted to dates from a date + object or an ISO string, with 'Date' as axis label.""" + for start_date in (dt.date(2020, 3, 1), '2020-03-01'): + ax = pts.plot_time_series( + self.ts, groups=2, start_date=start_date) + self.assertEqual(ax.get_xlabel(), 'Date') + xdata = pd.DatetimeIndex(ax.get_lines()[0].get_xdata()) + self.assertEqual(xdata[0], pd.Timestamp('2020-03-01')) + self.assertEqual(xdata[1], pd.Timestamp('2020-03-01 12:00')) + self.assertEqual(xdata[-1], pd.Timestamp('2020-03-04')) + + def test_tick_formatter(self): + """Test the y tick labels: thousands separators for large numbers, + plain fractions and a consistent notation for very small values.""" + ax = pts.plot_time_series(self.ts, groups=2) + formatter = ax.yaxis.get_major_formatter() + self.assertEqual(formatter(1234567.0), '1,234,567') + self.assertEqual(formatter(12.5), '12.5') + self.assertEqual(formatter(0.0), '0') + self.assertEqual(formatter(0.001), '0.001') + self.assertEqual(formatter(1e-4), '1e-4') + self.assertEqual(formatter(1e-5), '1e-5') + + # ---- dataframe ------------------------------------------------------ + + def test_to_dataframe(self): + """Test the long-form data frame with groups: columns, categorical + order, values per group and compartment, and summing by pivoting.""" + df = pts.time_series_to_dataframe(self.ts, self.labels, self.groups) + self.assertEqual(list(df.columns), + ['Time', 'Groups', 'Compartments', 'Values']) + self.assertEqual(len(df), 6 * len(self.times)) + self.assertEqual(list(df['Compartments'].cat.categories), self.labels) + self.assertEqual(list(df['Groups'].cat.categories), self.groups) + self.assertTrue(df['Compartments'].cat.ordered) + # Element 4 = second group, second compartment. + subset = df[(df['Groups'] == '5+') & (df['Compartments'] == 'I')] + np.testing.assert_array_equal(subset['Time'].to_numpy(), self.times) + np.testing.assert_allclose( + subset['Values'].to_numpy(), self.values[4]) + wide = df.pivot_table(index='Time', columns='Compartments', + values='Values', aggfunc='sum', observed=False) + np.testing.assert_allclose( + wide['S'].to_numpy(), self.values[0] + self.values[3]) + + def test_to_dataframe_without_groups_with_dates(self): + """Test the data frame without groups and with a date column that is + rounded to full seconds.""" + data = np.vstack([[0.0, 1.253706], np.ones((2, 2))]) + df = pts.time_series_to_dataframe( + data, start_date=dt.date(2021, 1, 1)) + self.assertEqual(list(df.columns), + ['Time', 'Date', 'Compartments', 'Values']) + self.assertEqual(list(df['Compartments'].cat.categories), + ['C1', 'C2']) + self.assertEqual(df['Date'].iloc[0], pd.Timestamp('2021-01-01')) + self.assertEqual(df['Date'].iloc[1], + pd.Timestamp('2021-01-02 06:05:20')) + + # ---- real bindings -------------------------------------------------- + + @unittest.skipIf(mio is None, 'memilio.simulation not installed') + def test_with_bindings_time_series(self): + """Test plotting and converting a TimeSeries of the bindings; skipped + if memilio.simulation is not installed.""" + ts = mio.TimeSeries(4) + ts.add_time_point(0.0, np.r_[100.0, 10.0, 200.0, 20.0]) + ts.add_time_point(1.0, np.r_[90.0, 20.0, 180.0, 40.0]) + ax = pts.plot_time_series(ts, ['S', 'I'], groups=['a', 'b']) + lines = ax.get_lines() + self.assertEqual(self._labels(ax), ['S', 'I']) + np.testing.assert_allclose(lines[0].get_ydata(), [300.0, 270.0]) + np.testing.assert_allclose(lines[1].get_ydata(), [30.0, 60.0]) + df = pts.time_series_to_dataframe(ts, ['S', 'I'], groups=2) + self.assertEqual(len(df), 8) + + +if __name__ == '__main__': + unittest.main() From 6f72f72b6511da339a42af31fa809ddb27927ede Mon Sep 17 00:00:00 2001 From: Kilian Volmer <13285635+kilianvolmer@users.noreply.github.com> Date: Wed, 30 Sep 2026 14:26:05 +0200 Subject: [PATCH 4/6] FIX: Use colors suggested by style guide --- .../memilio/plot/plotTimeSeries.py | 12 +- .../tests/test_plot_plotTimeSeries.py | 115 +++++++++++++++--- 2 files changed, 107 insertions(+), 20 deletions(-) diff --git a/pycode/memilio-plot/memilio/plot/plotTimeSeries.py b/pycode/memilio-plot/memilio/plot/plotTimeSeries.py index c991347d60..a247247b6a 100644 --- a/pycode/memilio-plot/memilio/plot/plotTimeSeries.py +++ b/pycode/memilio-plot/memilio/plot/plotTimeSeries.py @@ -40,11 +40,11 @@ import pandas as pd from matplotlib.ticker import FuncFormatter, NullFormatter -# Categorical palette in a fixed, colorblind-safe order. Compartment i is -# always drawn with _COLORS[i % 8]; from the ninth compartment on, the line -# style changes instead of introducing new hues. -_COLORS = ['#2a78d6', '#eb6834', '#1baf7a', '#eda100', - '#e87ba4', '#008300', '#4a3aa7', '#e34948'] +# Colors of matplotlib's 'tab10' colormap, the colorblind-friendly palette. +# Compartment i is always drawn with _COLORS[i % 10]; from the eleventh +# compartment on, the line style changes instead of introducing new hues. +_COLORS = ['#1f77b4', '#ff7f0e', '#2ca02c', '#d62728', '#9467bd', + '#8c564b', '#e377c2', '#7f7f7f', '#bcbd22', '#17becf'] _LINESTYLES = ['-', '--', '-.', ':'] _GRID_COLOR = '#e1e0d9' @@ -391,7 +391,7 @@ def plot_time_series( groups=['0-19', '20+'], title='ODE SEIR') ax.figure.savefig('seir.pdf') - Every compartment has a fixed color. For more than eight compartments, + Every compartment has a fixed color. For more than ten compartments, the colors are reused with a different line style; consider ``select`` to plot only the compartments of interest. diff --git a/pycode/memilio-plot/tests/test_plot_plotTimeSeries.py b/pycode/memilio-plot/tests/test_plot_plotTimeSeries.py index 6b856c246d..e1cd57148c 100644 --- a/pycode/memilio-plot/tests/test_plot_plotTimeSeries.py +++ b/pycode/memilio-plot/tests/test_plot_plotTimeSeries.py @@ -115,8 +115,10 @@ def test_default_labels(self): """Test that elements without labels are named C1, C2, ... as in TimeSeries.print_table, and that the legend lists them.""" ax = pts.plot_time_series(self.ts) + # Check that one line per element is drawn and named by its number. expected = ['C1', 'C2', 'C3', 'C4', 'C5', 'C6'] self.assertEqual(self._labels(ax), expected) + # Check that the legend contains the same names in the same order. self.assertEqual(self._legend_texts(ax), expected) def test_given_labels(self): @@ -126,7 +128,10 @@ def test_given_labels(self): labels = ['a', 'b', 'c', 'd', 'e', 'f'] for time_series in (self.ts, self.data): ax = pts.plot_time_series(time_series, labels) + # Check that the lines are labeled in the order of the elements. self.assertEqual(self._labels(ax), labels) + # Check that every line shows the time points and the values of + # its element without modification. for line, row in zip(ax.get_lines(), self.values): np.testing.assert_array_equal(line.get_xdata(), self.times) np.testing.assert_array_equal(line.get_ydata(), row) @@ -138,6 +143,7 @@ def test_labels_from_enum(self): expected = ['Susceptible', 'Infected', 'Recovered'] for labels in (_State, _BindingsLikeEnum, _BindingsLikeEnum.values()): ax = pts.plot_time_series(_StubTimeSeries(data), labels) + # Check that the names of the members are used as labels. self.assertEqual(self._labels(ax), expected) def test_input_is_copied(self): @@ -146,6 +152,7 @@ def test_input_is_copied(self): data = self.data.copy() ax = pts.plot_time_series(_StubTimeSeries(data), groups=2) data[1:] = 0.0 + # Check that the line still holds the original (non-zero) values. self.assertGreater(ax.get_lines()[0].get_ydata().sum(), 0.0) def test_invalid_input(self): @@ -172,6 +179,8 @@ def test_invalid_input(self): for description, kwargs, error in cases: with self.subTest(description): kwargs = {'time_series': self.ts, **kwargs} + # Check that the invalid input is rejected with the expected + # error type instead of producing a wrong plot. with self.assertRaises(error): pts.plot_time_series(**kwargs) @@ -181,7 +190,10 @@ def test_sum_groups(self): """Test that the compartments are summed over the groups by default.""" ax = pts.plot_time_series(self.ts, self.labels, groups=self.groups) + # Check that there is one line per compartment, not per element. self.assertEqual(self._labels(ax), self.labels) + # Check that each line is the sum of the compartment over both + # groups, i.e. of the elements c and c + num_compartments. for c, line in enumerate(ax.get_lines()): expected = self.values[c] + self.values[c + self.num_compartments] np.testing.assert_allclose(line.get_ydata(), expected) @@ -190,8 +202,10 @@ def test_groups_by_count(self): """Test that groups given as a number are named Group 0, Group 1, ... and that default labels are per compartment.""" ax = pts.plot_time_series(self.ts, groups=2) + # Check that the default labels count the compartments of one group. self.assertEqual(self._labels(ax), ['C1', 'C2', 'C3']) ax = pts.plot_time_series(self.ts, groups=2, sum_groups=False) + # Check that the groups are named by their index. self.assertEqual(self._labels(ax)[:2], ['C1 (Group 0)', 'C1 (Group 1)']) @@ -202,16 +216,23 @@ def test_per_group_lines(self): ax = pts.plot_time_series( self.ts, self.labels, groups=self.groups, sum_groups=False) lines = ax.get_lines() + # Check that the lines are ordered by compartment, then by group, + # and labeled with both names. self.assertEqual( self._labels(ax), ['S (0-4)', 'S (5+)', 'I (0-4)', 'I (5+)', 'R (0-4)', 'R (5+)']) + # Check that lines of the same compartment share the color and + # lines of different compartments do not. self.assertEqual(lines[0].get_color(), lines[1].get_color()) self.assertNotEqual(lines[0].get_color(), lines[2].get_color()) + # Check that the line style distinguishes the groups. self.assertNotEqual(lines[0].get_linestyle(), lines[1].get_linestyle()) self.assertEqual(lines[0].get_linestyle(), lines[2].get_linestyle()) + # Check that a line shows the values of its element (group 1, S). np.testing.assert_allclose(lines[1].get_ydata(), self.values[3]) ax = pts.plot_time_series(self.ts, sum_groups=False) + # Check that without groups still one line per element is drawn. self.assertEqual(len(ax.get_lines()), 6) def test_single_compartment_per_group(self): @@ -221,16 +242,22 @@ def test_single_compartment_per_group(self): self.ts, self.labels, groups=self.groups, sum_groups=False, select='I') lines = ax.get_lines() + # Check that only the selected compartment is drawn, once per group. self.assertEqual(self._labels(ax), ['I (0-4)', 'I (5+)']) - self.assertEqual(lines[0].get_color(), '#2a78d6') - self.assertEqual(lines[1].get_color(), '#eb6834') + # Check that the groups get the first two palette colors and the + # same (solid) line style. + self.assertEqual(lines[0].get_color(), '#1f77b4') + self.assertEqual(lines[1].get_color(), '#ff7f0e') self.assertEqual(lines[0].get_linestyle(), lines[1].get_linestyle()) + # Check that the lines show the values of compartment I in group 0 + # and group 1, respectively. np.testing.assert_allclose(lines[0].get_ydata(), self.values[1]) np.testing.assert_allclose(lines[1].get_ydata(), self.values[4]) ax = pts.plot_time_series( self.ts, self.labels, groups=self.groups, sum_groups=False, select='I', colors=['black', 'gray']) + # Check that given colors are assigned to the groups in this case. self.assertEqual([line.get_color() for line in ax.get_lines()], ['black', 'gray']) @@ -238,8 +265,12 @@ def test_ambiguous_lines_warn(self): """Test that a warning is issued if lines share color and line style (more groups than line styles), and not otherwise.""" data = np.vstack([self.times, np.ones((10, len(self.times)))]) + # Check that 5 groups x 2 compartments with only 4 line styles + # triggers the warning. with self.assertWarns(UserWarning): pts.plot_time_series(data, ['a', 'b'], groups=5, sum_groups=False) + # Check that summing the groups or selecting a single compartment + # (groups colored) does not warn. with warnings.catch_warnings(): warnings.simplefilter('error') pts.plot_time_series(data, ['a', 'b'], groups=5) @@ -256,9 +287,14 @@ def test_select(self): for line in ax_all.get_lines()} ax = pts.plot_time_series( self.ts, self.labels, groups=2, select=['R', 'I']) + # Check that the selection is drawn in the order of the TimeSeries, + # not in the order given. self.assertEqual(self._labels(ax), ['I', 'R']) + # Check that the compartments keep the colors of the full plot. for line in ax.get_lines(): self.assertEqual(line.get_color(), colors[line.get_label()]) + # Check that a single index, a single named member and a mixture of + # indices and names (with duplicates) are accepted. ax = pts.plot_time_series(self.ts, self.labels, groups=2, select=2) self.assertEqual(self._labels(ax), ['R']) ax = pts.plot_time_series( @@ -267,6 +303,7 @@ def test_select(self): ax = pts.plot_time_series( self.ts, self.labels, groups=2, select=[0, 'S', np.int64(1)]) self.assertEqual(self._labels(ax), ['S', 'I']) + # Check that an unknown name and an index out of range are rejected. with self.assertRaises(ValueError): pts.plot_time_series(self.ts, self.labels, groups=2, select='X') with self.assertRaises(ValueError): @@ -277,24 +314,34 @@ def test_custom_colors(self): or by a sequence indexed by compartment.""" ax = pts.plot_time_series( self.ts, self.labels, groups=2, colors={'I': 'black'}) + # Check that only the named compartment changes its color. self.assertEqual([line.get_color() for line in ax.get_lines()], - ['#2a78d6', 'black', '#1baf7a']) + ['#1f77b4', 'black', '#2ca02c']) ax = pts.plot_time_series( self.ts, self.labels, groups=2, colors=['red', 'green']) + # Check that the sequence is used by index and the remaining + # compartment falls back to the palette. self.assertEqual([line.get_color() for line in ax.get_lines()], - ['red', 'green', '#1baf7a']) + ['red', 'green', '#2ca02c']) - def test_more_than_eight_compartments(self): - """Test that from the ninth compartment on, the colors are reused - with a different line style.""" - data = np.vstack([self.times, np.ones((10, len(self.times)))]) + def test_more_than_ten_compartments(self): + """Test that the ten colors of tab10 are used in order and that from + the eleventh compartment on, they are reused with a different line + style.""" + data = np.vstack([self.times, np.ones((11, len(self.times)))]) ax = pts.plot_time_series(data) lines = ax.get_lines() - self.assertEqual(len(lines), 10) - self.assertEqual(lines[8].get_color(), lines[0].get_color()) - self.assertEqual(lines[0].get_linestyle(), '-') - self.assertEqual(lines[8].get_linestyle(), '--') - self.assertEqual(len(self._legend_texts(ax)), 10) + # Check that all eleven compartments are drawn and listed. + self.assertEqual(len(lines), 11) + self.assertEqual(len(self._legend_texts(ax)), 11) + # Check that the first ten lines have distinct, solid colors. + colors = [line.get_color() for line in lines[:10]] + self.assertEqual(len(set(colors)), 10) + self.assertTrue(all(line.get_linestyle() == '-' + for line in lines[:10])) + # Check that the eleventh line reuses the first color, but dashed. + self.assertEqual(lines[10].get_color(), lines[0].get_color()) + self.assertEqual(lines[10].get_linestyle(), '--') # ---- legend --------------------------------------------------------- @@ -302,12 +349,16 @@ def test_legend_right_of_axes(self): """Test that the legend is placed right of the axes, also for a single line and for a user-provided axes.""" ax = pts.plot_time_series(self.ts, groups=2) + # Check that the legend of a new figure is anchored right of the + # axes. self.assertTrue(self._legend_is_outside(ax)) ax = pts.plot_time_series(self.ts, self.labels, groups=2, select='I') + # Check that a single line still gets a legend, right of the axes. self.assertEqual(self._legend_texts(ax), ['I']) self.assertTrue(self._legend_is_outside(ax)) _, axes = plt.subplots(2, 2) ax = pts.plot_time_series(self.ts, ax=axes[0, 0]) + # Check that the same holds for an axes of a user-created grid. self.assertEqual(len(self._legend_texts(ax)), 6) self.assertTrue(self._legend_is_outside(ax)) @@ -315,8 +366,10 @@ def test_no_legend(self): """Test that no legend is drawn if it is disabled or if there are no lines.""" ax = pts.plot_time_series(self.ts, groups=2, legend=False) + # Check that legend=False suppresses the legend. self.assertIsNone(ax.get_legend()) ax = pts.plot_time_series(self.ts, groups=2, select=[]) + # Check that an empty selection draws neither lines nor a legend. self.assertEqual(len(ax.get_lines()), 0) self.assertIsNone(ax.get_legend()) @@ -326,10 +379,14 @@ def test_new_figure(self): """Test that a new figure of the given size with constrained layout and default axis labels is created if no axes is given.""" ax = pts.plot_time_series(self.ts, groups=2, figsize=(10, 3)) + # Check that the figure has the requested size. np.testing.assert_array_equal( ax.figure.get_size_inches(), [10.0, 3.0]) + # Check that constrained layout is used, so that the legend right of + # the axes is not cut off. self.assertIsInstance( ax.figure.get_layout_engine(), ConstrainedLayoutEngine) + # Check the default axis labels and that there is no title. self.assertEqual(ax.get_xlabel(), 'Time (days)') self.assertEqual(ax.get_ylabel(), 'Number of individuals') self.assertEqual(ax.get_title(), '') @@ -341,12 +398,16 @@ def test_given_axes_and_labels(self): ax = pts.plot_time_series( self.ts, self.labels, groups=2, ax=ax_in, title='T', xlabel='x', ylabel='y', linewidth=4.0) + # Check that the given axes is used and returned. self.assertIs(ax, ax_in) + # Check that title and axis labels are set as given. self.assertEqual(ax.get_title(), 'T') self.assertEqual(ax.get_xlabel(), 'x') self.assertEqual(ax.get_ylabel(), 'y') + # Check that the axes is styled (linear scale, top spine hidden). self.assertEqual(ax.get_yscale(), 'linear') self.assertFalse(ax.spines['top'].get_visible()) + # Check that the line width is passed to all lines. for line in ax.get_lines(): self.assertEqual(line.get_linewidth(), 4.0) @@ -355,17 +416,23 @@ def test_plot_kwargs(self): passed to all lines and override the defaults, except the label.""" ax = pts.plot_time_series( self.ts, groups=2, lw=3.0, color='black', ls=':', label='x') + # Check that the aliases lw and ls as well as color are applied to + # every line instead of the defaults. for line in ax.get_lines(): self.assertEqual(line.get_linewidth(), 3.0) self.assertEqual(line.get_color(), 'black') self.assertEqual(line.get_linestyle(), ':') + # Check that a given label is ignored in favor of the compartments. self.assertEqual(self._labels(ax), ['C1', 'C2', 'C3']) def test_log_scale(self): """Test the logarithmic y axis with the count formatter on the major ticks and no labels on the minor ticks.""" ax = pts.plot_time_series(self.ts, groups=2, log_scale=True) + # Check that the y axis is logarithmic. self.assertEqual(ax.get_yscale(), 'log') + # Check that the major ticks keep the thousands separators and the + # minor ticks are not labeled. self.assertEqual(ax.yaxis.get_major_formatter()(1000.0), '1,000') self.assertIsInstance(ax.yaxis.get_minor_formatter(), NullFormatter) @@ -375,7 +442,10 @@ def test_start_date(self): for start_date in (dt.date(2020, 3, 1), '2020-03-01'): ax = pts.plot_time_series( self.ts, groups=2, start_date=start_date) + # Check that the axis label switches to 'Date'. self.assertEqual(ax.get_xlabel(), 'Date') + # Check that time 0 is the start date and fractional days are + # converted to times of day. xdata = pd.DatetimeIndex(ax.get_lines()[0].get_xdata()) self.assertEqual(xdata[0], pd.Timestamp('2020-03-01')) self.assertEqual(xdata[1], pd.Timestamp('2020-03-01 12:00')) @@ -386,10 +456,14 @@ def test_tick_formatter(self): plain fractions and a consistent notation for very small values.""" ax = pts.plot_time_series(self.ts, groups=2) formatter = ax.yaxis.get_major_formatter() + # Check that large numbers get thousands separators without + # decimals. self.assertEqual(formatter(1234567.0), '1,234,567') + # Check that small numbers are printed plainly. self.assertEqual(formatter(12.5), '12.5') self.assertEqual(formatter(0.0), '0') self.assertEqual(formatter(0.001), '0.001') + # Check that values below 1e-3 use the same short exponent notation. self.assertEqual(formatter(1e-4), '1e-4') self.assertEqual(formatter(1e-5), '1e-5') @@ -399,17 +473,23 @@ def test_to_dataframe(self): """Test the long-form data frame with groups: columns, categorical order, values per group and compartment, and summing by pivoting.""" df = pts.time_series_to_dataframe(self.ts, self.labels, self.groups) + # Check the columns and that there is one row per element and time + # point. self.assertEqual(list(df.columns), ['Time', 'Groups', 'Compartments', 'Values']) self.assertEqual(len(df), 6 * len(self.times)) + # Check that compartments and groups are ordered categoricals in + # the order of the TimeSeries. self.assertEqual(list(df['Compartments'].cat.categories), self.labels) self.assertEqual(list(df['Groups'].cat.categories), self.groups) self.assertTrue(df['Compartments'].cat.ordered) - # Element 4 = second group, second compartment. + # Check that the rows of group 1, compartment I hold the time points + # and the values of element 4. subset = df[(df['Groups'] == '5+') & (df['Compartments'] == 'I')] np.testing.assert_array_equal(subset['Time'].to_numpy(), self.times) np.testing.assert_allclose( subset['Values'].to_numpy(), self.values[4]) + # Check that pivoting with sum reproduces the group totals. wide = df.pivot_table(index='Time', columns='Compartments', values='Values', aggfunc='sum', observed=False) np.testing.assert_allclose( @@ -421,10 +501,14 @@ def test_to_dataframe_without_groups_with_dates(self): data = np.vstack([[0.0, 1.253706], np.ones((2, 2))]) df = pts.time_series_to_dataframe( data, start_date=dt.date(2021, 1, 1)) + # Check that there is a Date column but no Groups column and that + # the default compartment names are used. self.assertEqual(list(df.columns), ['Time', 'Date', 'Compartments', 'Values']) self.assertEqual(list(df['Compartments'].cat.categories), ['C1', 'C2']) + # Check that time 0 maps to the start date and 1.253706 days are + # rounded to full seconds. self.assertEqual(df['Date'].iloc[0], pd.Timestamp('2021-01-01')) self.assertEqual(df['Date'].iloc[1], pd.Timestamp('2021-01-02 06:05:20')) @@ -440,9 +524,12 @@ def test_with_bindings_time_series(self): ts.add_time_point(1.0, np.r_[90.0, 20.0, 180.0, 40.0]) ax = pts.plot_time_series(ts, ['S', 'I'], groups=['a', 'b']) lines = ax.get_lines() + # Check that the two compartments are summed over the groups a and + # b of the bindings' TimeSeries. self.assertEqual(self._labels(ax), ['S', 'I']) np.testing.assert_allclose(lines[0].get_ydata(), [300.0, 270.0]) np.testing.assert_allclose(lines[1].get_ydata(), [30.0, 60.0]) + # Check that the data frame has one row per element and time point. df = pts.time_series_to_dataframe(ts, ['S', 'I'], groups=2) self.assertEqual(len(df), 8) From dfd0ff4e63d2903e64ea86d61e811f799c7aa6dd Mon Sep 17 00:00:00 2001 From: Kilian Volmer <13285635+kilianvolmer@users.noreply.github.com> Date: Fri, 2 Oct 2026 13:31:07 +0200 Subject: [PATCH 5/6] Apply suggestions from code review --- docs/source/python/m-plot.rst | 2 +- pycode/examples/plot/plotSimulationResults.py | 2 +- pycode/examples/simulation/ode_secir_groups.py | 3 +++ pycode/examples/simulation/ode_secir_mobility.py | 3 +++ pycode/examples/simulation/ode_secir_simple.py | 3 +++ pycode/memilio-plot/memilio/plot/plotTimeSeries.py | 3 ++- pycode/memilio-simulation/pyproject.toml | 3 ++- 7 files changed, 15 insertions(+), 4 deletions(-) diff --git a/docs/source/python/m-plot.rst b/docs/source/python/m-plot.rst index 3286e4d539..38a93e8dbb 100644 --- a/docs/source/python/m-plot.rst +++ b/docs/source/python/m-plot.rst @@ -108,6 +108,6 @@ simulations of the :doc:`MEmilio Python bindings `: ``layout='constrained'`` or save them with ``bbox_inches='tight'``. ``time_series_to_dataframe`` converts a ``TimeSeries`` into a tidy pandas ``DataFrame`` with the columns ``Time``, -``Group``, ``Compartment`` and ``Value`` (and ``Date`` if a start date is given) for use with other libraries such as +``Groups``, ``Compartments`` and ``Values`` (and ``Date`` if a start date is given) for use with other libraries such as seaborn, plotnine or altair. A complete example is given in `examples/plot/plotSimulationResults.py `_. diff --git a/pycode/examples/plot/plotSimulationResults.py b/pycode/examples/plot/plotSimulationResults.py index e557275b02..5c7790d248 100644 --- a/pycode/examples/plot/plotSimulationResults.py +++ b/pycode/examples/plot/plotSimulationResults.py @@ -52,7 +52,7 @@ def run_ode_seir_simulation(days=100, dt=0.1): group = AgeGroup(i) model.parameters.TimeExposed[group] = 5.2 model.parameters.TimeInfected[group] = 6. - model.parameters.TransmissionProbabilityOnContact[group] = 1. * i + model.parameters.TransmissionProbabilityOnContact[group] = 1. * (i+1) model.populations[group, InfectionState.Exposed] = 100 model.populations[group, InfectionState.Infected] = 50 model.populations[group, InfectionState.Recovered] = 10 diff --git a/pycode/examples/simulation/ode_secir_groups.py b/pycode/examples/simulation/ode_secir_groups.py index e85718b1a7..2ca79bb0c0 100644 --- a/pycode/examples/simulation/ode_secir_groups.py +++ b/pycode/examples/simulation/ode_secir_groups.py @@ -17,6 +17,9 @@ # See the License for the specific language governing permissions and # limitations under the License. ############################################################################# + +# This example requires memilio.simulation and memilio.plot to be installed! + import argparse import os from datetime import date diff --git a/pycode/examples/simulation/ode_secir_mobility.py b/pycode/examples/simulation/ode_secir_mobility.py index 8c14049a73..9ebd1cdd06 100644 --- a/pycode/examples/simulation/ode_secir_mobility.py +++ b/pycode/examples/simulation/ode_secir_mobility.py @@ -17,6 +17,9 @@ # See the License for the specific language governing permissions and # limitations under the License. ############################################################################# + +# This example requires memilio.simulation and memilio.plot to be installed! + import argparse import matplotlib.pyplot as plt diff --git a/pycode/examples/simulation/ode_secir_simple.py b/pycode/examples/simulation/ode_secir_simple.py index e389edd5fe..9ea4b478fa 100644 --- a/pycode/examples/simulation/ode_secir_simple.py +++ b/pycode/examples/simulation/ode_secir_simple.py @@ -17,6 +17,9 @@ # See the License for the specific language governing permissions and # limitations under the License. ############################################################################# + +# This example requires memilio.simulation and memilio.plot to be installed! + import argparse from datetime import date diff --git a/pycode/memilio-plot/memilio/plot/plotTimeSeries.py b/pycode/memilio-plot/memilio/plot/plotTimeSeries.py index a247247b6a..63b4892c4c 100644 --- a/pycode/memilio-plot/memilio/plot/plotTimeSeries.py +++ b/pycode/memilio-plot/memilio/plot/plotTimeSeries.py @@ -438,7 +438,8 @@ def plot_time_series( :param figsize: Size of the figure in inches if a new figure is created. (Default value = (8, 4.5)) :param plot_kwargs: Additional keyword arguments passed to - ``matplotlib.axes.Axes.plot`` for every line, e.g. ``linewidth``. + ``matplotlib.axes.Axes.plot`` for every line, e.g. ``linewidth``. This + overwrites the default settings. :returns: The matplotlib axes containing the plot. Use ``ax.figure`` to access and save the figure. """ diff --git a/pycode/memilio-simulation/pyproject.toml b/pycode/memilio-simulation/pyproject.toml index c120d78d74..90d7cc4eb5 100644 --- a/pycode/memilio-simulation/pyproject.toml +++ b/pycode/memilio-simulation/pyproject.toml @@ -3,7 +3,7 @@ name = "memilio-simulation" dynamic = ["version"] description = "Part of MEmilio project, Python bindings to the C++ libraries that contain the models and simulations." readme = "README.md" -requires-python = ">=3.8" +requires-python = ">=3.9" license = "Apache-2.0" authors = [{ name = "MEmilio Team" }] maintainers = [ @@ -31,6 +31,7 @@ classifiers = [ [project.optional-dependencies] dev = [] +plot = ["memilio.plot"] [project.urls] Homepage = "https://github.com/SciCompMod/memilio" From 1b9c5f7901b4e9ada7f7ec304a7ac1fb9f822a5f Mon Sep 17 00:00:00 2001 From: Kilian Volmer <13285635+kilianvolmer@users.noreply.github.com> Date: Mon, 5 Oct 2026 13:08:08 +0200 Subject: [PATCH 6/6] Remove misleading comment --- pycode/examples/simulation/ode_secir_mobility.py | 1 - 1 file changed, 1 deletion(-) diff --git a/pycode/examples/simulation/ode_secir_mobility.py b/pycode/examples/simulation/ode_secir_mobility.py index 9ebd1cdd06..4e9109033b 100644 --- a/pycode/examples/simulation/ode_secir_mobility.py +++ b/pycode/examples/simulation/ode_secir_mobility.py @@ -123,7 +123,6 @@ def run_ode_secir_mobility_simulation(plot_results=True): region_results = [region0_result, region1_result] region_labels = ['Region 0', 'Region 1'] - # All compartments of each region on a logarithmic axis. fig, axes = plt.subplots(1, 2, figsize=(16, 5), layout='constrained') for region_result, region_label, ax in zip( region_results, region_labels, axes):