Source code for devctrl.data.plot

from matplotlib.axes import Axes
import numpy as np
import matplotlib.pyplot as plt
from .data import NDArray, Data, Columns
from typing import Optional, Callable

import logging
log = logging.getLogger(__name__)


DEFAULT_COLOR = "#000000"
DEFAULT_SCATTER_IN_COLORS_DOT_KW = dict(color="black", marker=".", s=8)
DEFAULT_SCATTER_IN_COLORS_KW = dict(marker="o", s=12)
[docs] def plot_data_on_ax(data: NDArray, columns: Columns, ax: Axes, key: str, xkey: str, ylabel: str|bool=True, scale_factor: None|float|int=None, error_shade=False, data_getter_obj=None, is_normalized=False, mode=None, max_ylim: None|tuple[float,float]=None, zero_hline=None, zero_hline_keys: Optional[list[str]]=None, scatter_in_colors: Optional[list[str]]=None, scatter_in_colors_kw: dict=DEFAULT_SCATTER_IN_COLORS_KW, scatter_in_colors_dot_kw: Optional[dict]=DEFAULT_SCATTER_IN_COLORS_DOT_KW, **plot_kw ): """ Plot column data on a single axis :param data: Numpy array with <columns> :param ax: Matplotlib axes :param key: Which key of <columns> to plot :param columns: List of column keys :param xkey: Which key of <columns> to plot on the x-axis (wl or E) :param ylabel: Either a label (str) or True to use the PrsData.get_column_tex_label_with_unit_and_scale() or False to not use a label :param scale_factor: Which scaling to apply to the ylabel. :param error_shade: Whether to include the errors by drawing a semi-transparent shade around the plot line :param plot_kw: Keyword arguments to pass to the axes.plot() function :param mode: Whether to use reflection or transmission labels :param zero_hline: Whether to plot a horizontal line at zero. If None, automatically plots one for certain keys :param zero_hline_keys: List of column keys for which to automatically plot a zero hline :param scatter_in_colors: If given, scatter the data in the given colors instead of plotting the line :param scatter_in_colors_kw: Keywords of the scatter plots. Colors are taken from scatter_in_colors list :param scatter_in_colors_dot_kw: If not None, scatter single-colored dots in addition to the colored dots """ if data_getter_obj is None: data_getter_obj = Data if zero_hline is None: zero_hline = key in zero_hline_keys if zero_hline == True: ax.axhline(0, color="black", linewidth=plt.rcParams["axes.linewidth"]) x_idx = columns.index(xkey) data_idx = columns.index(key) ydata = data[:, data_idx] # # apply scaling # if scale_factor is not None: # log.debug(f"Applying scale: {scale_factor} to key '{key}'") # ydata = ydata.copy()/scale_factor # ydata = remove_outliers(ydata, threshold=0.5) # ydata = data[:, data_idx] ylim_before = ax.get_ylim() if error_shade: if f"s{key}" in columns: sdata_idx = columns.index(f"s{key}") sydata = data[:, sdata_idx] # if scale_factor is not None: # sydata = sydata.copy()/scale_factor # determine the color for the shade, use semi-transparent color of line linecolor = plot_kw["color"] if "color" in plot_kw else data_getter_obj.get_column_color(key, default=DEFAULT_COLOR, mode=mode) color = "#00000030" if not linecolor.startswith("#"): # if a named color was given, try getting the #rrggbb try: import matplotlib.colors linecolor = matplotlib.colors.get_named_colors_mapping()[linecolor] if not type(linecolor) == str: log.debug(f"Non-string color: {linecolor}") linecolor = color raise ValueError() except: pass if linecolor.startswith("#"): if len(linecolor) in [7,9]: # #rrggbb(aa) color = linecolor[:7] + "33" # make semi-transparent elif len(linecolor) in [4,5]: # #rgb(a): color = linecolor[:4] + "3" # make semi-transparent ax.fill_between(data[:, x_idx], ydata-sydata, ydata+sydata, color=color) else: log.warning(f"Can not draw error shade, since 's{key}' is not in columns") if scatter_in_colors is not None: line = [] for i in range(ydata.size): collection_color = ax.scatter(data[i, x_idx], ydata[i], **(dict(color=scatter_in_colors[i]) | scatter_in_colors_kw)) if scatter_in_colors_dot_kw is not None: collection_dot = ax.scatter(data[i, x_idx], ydata[i], **(scatter_in_colors_dot_kw)) line.append(collection_dot) line.append(collection_color) else: line, = ax.plot(data[:, x_idx], ydata, **(dict(color=data_getter_obj.get_column_color(key, default=DEFAULT_COLOR, mode=mode)) | plot_kw)) if error_shade: # set y limits # if the error_shade is at some point much larger than the line, we dont want to show it (likely an outlier) # therefore we ignore the error shade ylimits ymin = np.nanmin(ydata) ymax = np.nanmax(ydata) # extra margins: 5% each side margin = 0.05 * (ymax-ymin) # if the limits were previously larger, dont re-set them (might be form another plot line) # defaults are (0.0, 1.0), ignore those default_lim = ylim_before[0] == 0.0 and ylim_before[1] == 1.0 ymin = min(ymin-margin, ylim_before[0]) if not default_lim else ymin-margin ymax = max(ymax+margin, ylim_before[1]) if not default_lim else ymax+margin if not (np.isnan(ymin) or np.isinf(ymin) or np.isnan(ymax) or np.isinf(ymax)): ax.set_ylim(ymin, ymax) else: log.debug(f"Not setting limits since at least one is invalid: ymin={ymin}, ymax={ymax}") if max_ylim: # hard limit to given values ylims = ax.get_ylim() ymin = max(ylims[0], max_ylim[0]) ymax = min(ylims[1], max_ylim[1]) ax.set_ylim(ymin, ymax) # label if ylabel == True: # if True, use auto ylabel if is_normalized: ylabel = data_getter_obj.get_column_tex_label(key, mode=mode) elif scale_factor == False: ylabel = data_getter_obj.get_column_tex_label_with_unit(key, mode=mode) else: ylabel = data_getter_obj.get_column_tex_label_with_unit_and_scale(key, scale=scale_factor, mode=mode) if ylabel: # if string ax.set_ylabel(ylabel) return line