Visualization
epilearn.visualize provides two low-level plotting helpers, unchanged in 0.1.0:
plot_series() for time series and
plot_graph() for a node-coloured network. Both are
re-exported at package level, so from epilearn.visualize import plot_series,
plot_graph works.
Note
Neither helper calls plt.show(): the call inside plot_graph is commented
out so the module stays usable in headless environments. Nothing appears on
screen — save the figure yourself (or rely on the inline renderer in a
notebook).
import numpy as np
import matplotlib.pyplot as plt
from epilearn.data import Dataset
from epilearn.visualize import plot_series, plot_graph
dataset = Dataset()
dataset.load_toy_dataset()
# --- time series: one line per column ---------------------------------
series = dataset.y[:, :3].numpy() # first three regions
plot_series(series, columns=['region 0', 'region 1', 'region 2'],
fig_size=(12, 4))
# -> writes ./plot_series.png in the current working directory
# --- network: plot_graph wants an EDGE INDEX, not an adjacency matrix --
edge_index = np.array(np.nonzero(dataset.graph.numpy())) # (2, n_edges)
last_day = dataset.y[-1].numpy()
states = (last_day > np.median(last_day)).astype(int) # class per node
pos = plot_graph(states, edge_index,
classes=['below median', 'above median'])
plt.savefig('graph_last_day.png') # plot_graph does not save for you
print(len(pos)) # 47 node positions, reusable layout
On the bundled toy dataset this prints 47 and produces plot_series.png
(45 KB) and graph_last_day.png (880 KB).
Plot_Series
- epilearn.visualize.plot.plot_series(x: array, columns: list, fig_size=None)
Plots time series data from an array using Seaborn’s line plot functionality. Each series is represented as a separate line on the plot.
- Parameters:
x (np.array) – A NumPy array containing the data points of the time series.
columns (list) – List of strings representing the column names or labels for each time series in the array.
- Returns:
Displays the plot directly and returns None.
- Return type:
None
Warning
plot_series hard-codes its output path: the last line of the function is
plt.savefig("plot_series.png"). There is no save_path argument, so the
figure always lands in ./plot_series.png relative to the current working
directory, and each call overwrites the previous one. To control where it
goes, os.chdir() into the target directory first, or move the file
afterwards.
It also plots with seaborn.relplot, which creates its own figure — so a
preceding plt.figure(figsize=...) has no effect on the saved output and the
fig_size argument is effectively ignored.
Plot_Graph
- epilearn.visualize.plot.plot_graph(states: array, graph: array, classes=None, layout=None)
Plots a graph using NetworkX where nodes are colored based on their state and annotated with class labels. The graph’s layout can be specified; otherwise, a spring layout is used.
- Parameters:
states (np.array) – Array containing the state of each node used to determine the color of nodes.
graph (np.array) – A 2D array where each column represents an edge between two nodes (start and end).
classes (list, optional) – List of class labels corresponding to each node. Default is None.
layout (dict, optional) – A dictionary defining the position of each node for custom layouts. If None, a spring layout is used. Default is None.
- Returns:
The positions of the nodes in the plot, useful for customizing the layout or for further graphical analysis.
- Return type:
dict
Warning
Three sharp edges in plot_graph:
graphis an edge index of shape(2, n_edges)— columniis the edge(graph[0, i], graph[1, i]). Passing an(N, N)adjacency matrix silently builds the wrong graph. Convert withnp.array(np.nonzero(adjacency)).classesis declared optional but is required in practice: the body doeslabels[node] = classes[node], so leaving it asNoneraisesTypeError: 'NoneType' object is not subscriptable.The internal colour map has only five entries (
0–4), so every value instatesmust be an integer in0..4; anything else raisesKeyError.
The return value is the NetworkX position dict, which you can feed back in as
layout= to keep the same node placement across frames.
Task-level plotting
Note
Renamed in 0.1.0: Forecast.plot_forecasts is now
Forecast.plot_preds. The old name no longer exists.
Forecast.plot_preds draws predictions against targets together with the
conformal/adaptive uncertainty band:
task.plot_preds(eval_results, n_show=None, figsize=(15, 7), save_path=None,
backend='matplotlib', region_idx=0, horizon_idx=-1,
interactive=False)
eval_results is the dict returned by task.evaluate_model(...); it must
contain predictions and targets (shape (time, regions, horizon)), and
picks up adaptive_lower / adaptive_upper when present. Unlike the two
helpers above this one does take a save_path, and it honours a 'plotly'
backend for interactive HTML output. It returns (fig, ax) for matplotlib
and fig for plotly.