class SimplePlot(TkFigure):
"""
Simple plot - single plot in frame with axes
"""
_y_axis_expansion_factor = 0.1
def __init__(self, root: tk.Misc, xdata: list[float], ydata: list[float],
xlabel: str = '', ylabel: str = '', title: str = '',
config: dict | None = None, fig_size: tuple[int, int] | None = None, fig_dpi: int = None):
super().__init__(
root=root,
config=config,
fig_size=fig_size,
fig_dpi=fig_dpi
)
self.plot_list: list[plt.Line2D] = []
# Add axes
self.ax1 = self.fig.add_subplot(111)
self.ax1.set_autoscaley_on(True)
self.ax1.set_autoscalex_on(True)
self.ax1.set_xlabel(xlabel)
self.ax1.set_ylabel(ylabel)
self.ax1.set_title(title)
if xdata:
self.plot(xdata, ydata)
def plot(self, *args, **kwargs) -> list[plt.Line2D]:
lines = self.ax1.plot(*args, **kwargs)
self.plot_list.extend(lines)
self.update_axes()
return lines
def update_labels(self, x_label: str | None = None, y_label: str | None = None,
title: str | None = None, legend: bool = False):
if x_label:
self.ax1.set_xlabel(x_label)
if y_label:
self.ax1.set_ylabel(y_label)
if title:
self.ax1.set_title(title)
if legend:
self.ax1.legend()
else:
self.ax1.legend([]).set_visible(False)
def plot_from_data(self, x_data: list[ndarray], y_data: list[ndarray], x_label: str = '', y_label: str = '',
title: str = '', labels: list[str] | None = None, **kwargs):
labels = [f"data #{n + 1}" for n in range(len(x_data))] if labels is None else labels
self.reset_plot()
for xdata, ydata, label in zip(x_data, y_data, labels):
lines = self.ax1.plot(np.ravel(xdata), np.ravel(ydata), label=label, **kwargs)
self.plot_list.extend(lines)
self.update_labels(x_label=x_label, y_label=y_label, title=title, legend=True if len(labels) > 1 else False)
self.update_axes()
def update_from_data(self, x_data: list[ndarray], y_data: list[ndarray], x_label: str | None = None,
y_label: str | None = None, title: str | None = None, legend: list[str] | None = None,
**kwargs):
if len(x_data) == len(self.plot_list):
# replace lines
legend = [None for _n in range(len(x_data))] if legend is None else legend
for xdata, ydata, label, line in zip(x_data, y_data, legend, self.plot_list):
line.set_data(np.ravel(xdata), np.ravel(ydata))
if label:
line.set_label(label)
self.update_labels(x_label=x_label, y_label=y_label, title=title, legend=True if len(legend) > 1 else False)
self.update_axes()
else:
self.plot_from_data(x_data, y_data, x_label, y_label, title, legend, **kwargs)
def remove_lines(self):
for obj in self.plot_list:
obj.remove()
self.plot_list.clear()
def reset_plot(self):
# self.ax1.set_xlabel(self.xaxis.get())
# self.ax1.set_ylabel(self.yaxis.get())
# self.ax1.set_title('')
self.ax1.set_prop_cycle(None) # reset colours
self.ax1.legend([]).set_visible(False)
self.remove_lines()
def _relim(self):
if not any(len(line.get_xdata()) for line in self.plot_list):
return
max_x_val = max(np.max(x) for line in self.plot_list if len(x := line.get_xdata()) > 0)
min_x_val = min(np.min(x) for line in self.plot_list if len(x := line.get_xdata()) > 0)
max_y_val = max(np.max(y) for line in self.plot_list if len(y := line.get_ydata()) > 0)
min_y_val = min(np.min(y) for line in self.plot_list if len(y := line.get_ydata()) > 0)
# expand y-axis slightly beyond data
y_diff = max_y_val - min_y_val
if y_diff == 0:
y_diff = max_y_val + 0.01
y_axis_max = max_y_val + self._y_axis_expansion_factor * y_diff
y_axis_min = min_y_val - self._y_axis_expansion_factor * y_diff
# max_y_val = 1.05 * max_y_val if max_y_val > 0 else max_y_val * 0.98
# min_y_val = 0.95 * min_y_val if min_y_val > 0 else min_y_val * 1.02
self.ax1.axis((min_x_val, max_x_val, y_axis_min, y_axis_max))
self.ax1.autoscale_view()
def update_axes(self):
self._relim()
self._update()