Skip to content

simple_plot

a tkinter frame with a single plot

SimplePlot

Bases: TkFigure

Simple plot - single plot in frame with axes

Source code in mmg_toolbox/tkguis/widgets/simple_plot.py
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()