Source code for pyorbb.plotting.orbital_diagram

import matplotlib.pyplot as plt
import matplotlib.patches as mpatches
import matplotlib.patheffects as pe
import numpy as np
import pyorbb


[docs] def anchored_text(ax, x, y, text, offset_axes=(0.0, 0.1), **kwargs): """ Draw text that is anchored to a point in data space with specified axes offsets. Returns the Artist instance. """ fig = ax.figure # initial dummy head (will be updated immediately) T = ax.text(x, y, text, path_effects=[pe.withStroke(linewidth=4, foreground='white')], **kwargs) def update_artist(event=None): # 1) get tail position in display coords anchor_disp = ax.transData.transform((x, y)) # display (pixel) coords # 2) convert tail display to axes coords # anchor_axes = ax.transAxes.inverted().transform(anchor_disp) # (0..1, 0..1) # update the positions of the arrow xyB = anchor_disp[0] + offset_axes[0]/2 * fig.dpi, anchor_disp[1] + offset_axes[1]/2 * fig.dpi # we have to set them like this, otherwise it does not work # new_disp = ax.transAxes.transform(xyB) new_data = ax.transData.inverted().transform(xyB) T.set_position(new_data) # Connect updates when y-limits change or figure is resized ax.callbacks.connect('ylim_changed', update_artist) fig.canvas.mpl_connect('resize_event', update_artist) # Do an initial update to set correct head position now update_artist() return T
[docs] def arrow_tail_with_axes_offset(ax, anchor, displacement_axes=(0.0, 0.1), anchor_axes=(0.0, 0.0), arrowstyle="<|-", alpha=1, **kwargs): """ Draw an arrow whose tail is fixed to `anchor` in data coords, but whose head is offset from the tail by `displacement_axes` measured in axes fraction (dx, dy in [0..1] of the axes width/height). Returns the ConnectionPatch instance. """ fig = ax.figure # initial dummy head (will be updated immediately) patch = mpatches.ConnectionPatch( anchor, (0, 0), coordsA=ax.transData, coordsB=ax.transData, arrowstyle=arrowstyle, shrinkA=0, shrinkB=0, **kwargs ) ax.add_patch(patch) def update_patch(event=None): # 1) get tail position in display coords anchor_disp = ax.transData.transform(anchor) # display (pixel) coords # 2) convert tail display to axes coords # anchor_axes = ax.transAxes.inverted().transform(anchor_disp) # (0..1, 0..1) # update the positions of the arrow xyA = anchor_disp[0] - displacement_axes[0]/2 * fig.dpi + anchor_axes[0]/2 * fig.dpi, anchor_disp[1] - displacement_axes[1]/2 * fig.dpi + anchor_axes[1]/2 * fig.dpi xyB = anchor_disp[0] + displacement_axes[0]/2 * fig.dpi + anchor_axes[0]/2 * fig.dpi, anchor_disp[1] + displacement_axes[1]/2 * fig.dpi + anchor_axes[1]/2 * fig.dpi # we have to set them like this, otherwise it does not work patch.xy1 = ax.transData.inverted().transform(xyA) patch.xy2 = ax.transData.inverted().transform(xyB) # Connect updates when y-limits change or figure is resized ax.callbacks.connect('ylim_changed', update_patch) fig.canvas.mpl_connect('resize_event', update_patch) # Do an initial update to set correct head position now update_patch() patch.orig_alpha = alpha patch.set_alpha(alpha) return patch
[docs] def draw_interaction(fmos, mos, connections, title=None, energy_type='energy', connection_types={}, ax=None, ylim=None, highlighted_orbitals=None, mo_column_name='Complex', fragment_order=None, warning_orbs=None, **kwargs): # load some parameters that we use for plotting # we take by default the parameters from rcParams unless they were provided with the function call arrow_length = kwargs.get('arrow_length', pyorbb.rcParams['plotting']['arrows']['arrow_length']) arrow_head_width = kwargs.get('arrow_head_width', pyorbb.rcParams['plotting']['arrows']['arrow_head_width']) arrow_head_length = kwargs.get('arrow_head_length', pyorbb.rcParams['plotting']['arrows']['arrow_head_length']) arrow_spacing = kwargs.get('arrow_spacing', pyorbb.rcParams['plotting']['arrows']['arrow_spacing']) arrow_color = kwargs.get('arrow_color', pyorbb.rcParams['plotting']['arrows']['arrow_color']) level_width = kwargs.get('level_width', pyorbb.rcParams['plotting']['levels']['level_width']) level_thickness = kwargs.get('level_thickness', pyorbb.rcParams['plotting']['levels']['level_thickness']) level_color = kwargs.get('level_color', pyorbb.rcParams['plotting']['levels']['level_color']) highlight_thickness = kwargs.get('highlight_thickness', pyorbb.rcParams['plotting']['levels']['highlight_thickness']) highlight_color = kwargs.get('highlight_color', pyorbb.rcParams['plotting']['levels']['highlight_color']) degeneracy_threshold = kwargs.get('degeneracy_threshold', pyorbb.rcParams['plotting']['levels']['degeneracy_threshold']) font_name = kwargs.get('font_name', pyorbb.rcParams['plotting']['font']['font_name']) font_size = kwargs.get('font_size', pyorbb.rcParams['plotting']['font']['font_size']) draw_mo_labels = kwargs.get('draw_mo_labels', pyorbb.rcParams['plotting']['labels']['draw_mo_labels']) draw_sfo_labels = kwargs.get('draw_sfo_labels', pyorbb.rcParams['plotting']['labels']['draw_sfo_labels']) orb_label_offset = kwargs.get('orb_label_offset', pyorbb.rcParams['plotting']['labels']['orb_label_offset']) axis_label_color = kwargs.get('axis_label_color', pyorbb.rcParams['plotting']['labels']['axis_label_color']) axis_label_font_size = kwargs.get('axis_label_font_size', pyorbb.rcParams['plotting']['labels']['axis_label_font_size']) label_color = kwargs.get('label_color', pyorbb.rcParams['plotting']['labels']['label_color']) alpha_range = kwargs.get('alpha_range', pyorbb.rcParams['plotting']['connections']['alpha_range']) OI_color = kwargs.get('OI_color', pyorbb.rcParams['plotting']['connections']['OI_color']) PR_color = kwargs.get('PR_color', pyorbb.rcParams['plotting']['connections']['PR_color']) sanitization_color = kwargs.get('sanitization_color', pyorbb.rcParams['plotting']['connections']['sanitization_color']) multiple_color = kwargs.get('multiple_color', pyorbb.rcParams['plotting']['connections']['multiple_color']) connection_width = kwargs.get('connection_width', pyorbb.rcParams['plotting']['connections']['line_width']) spine_color = kwargs.get('spine_color', pyorbb.rcParams['plotting']['spine']['spine_color']) if highlighted_orbitals is None: highlighted_orbitals = [] if ax is None: ax = plt.gca() if ylim is None: try: energies = [getattr(orb, energy_type) for orb in list(fmos)] + [orb.energy for orb in list(mos)] energy_span = max(energies) - min(energies) ax.set_ylim(min(energies) - .1 * energy_span, max(energies) + .1 * energy_span, auto=False) except ValueError: energy_span = 1 ax.set_ylim(0, 1, auto=False) energy_span *= 1.2 else: energy_span = ylim[1] - ylim[0] ax.set_ylim(*ylim, auto=False) frags = [] for fmo in fmos: if fmo.fragment not in frags: frags.append(fmo.fragment) if fragment_order is None: xtick_order = {mo_column_name: .5} for i, frag in enumerate(frags): if i == 0: xtick_order[frag] = -.5 else: xtick_order[frag] = i + .5 else: xtick_order = {mo_column_name: .5} for i, frag in enumerate(fragment_order): if i == 0: xtick_order[frag] = -.5 else: xtick_order[frag] = i + .5 ax.set_xlim(-1, len(frags), auto=False) sep_orbs = {frag: [fmo for fmo in fmos if fmo.fragment == frag] for frag in frags} sep_orbs[mo_column_name] = mos poss = {} for typ, sep_orbs_ in sep_orbs.items(): base_pos = xtick_order[typ] - .5 degenerates = [] for orb in sep_orbs_: if any(orb in degenerates_ for degenerates_ in degenerates): continue degenerates.append([orb]) for other_orb in sep_orbs_: if orb == other_orb: continue E1, E2 = orb.energy, other_orb.energy if orb in fmos: E1, E2 = getattr(orb, energy_type), getattr(other_orb, energy_type) if abs(E1 - E2) < (degeneracy_threshold * energy_span): degenerates[-1].append(other_orb) degenerates = [list(sorted(deg, key=lambda orb: orb.energy)) for deg in degenerates] for orb in sep_orbs_: orb_degenerate = [deg for deg in degenerates if orb in deg][0] deg_idx = orb_degenerate.index(orb) + 1 deg_degree = len(orb_degenerate) + 1 poss[orb] = base_pos + 1 / deg_degree * deg_idx ax.set_title(title) ax.set_ylabel('Orbital Energy / eV', size=axis_label_font_size, color=axis_label_color) ax.set_xticks(list(xtick_order.values()), list(xtick_order.keys()), color=spine_color) for i, artist in enumerate(ax.get_xticklabels()): if list(xtick_order.keys())[i] == mo_column_name: artist.is_MO = True else: artist.is_MO = False ax.spines[['top', 'bottom', 'right']].set_visible(False) ax.spines['left'].set_color(spine_color) ax.yaxis.label.set_color(spine_color) ax.tick_params(axis='y', colors=spine_color) ax.tick_params('x', labelsize=axis_label_font_size, labelcolor=label_color) ax.tick_params(bottom = False) for orb in poss: E = orb.energy if orb in fmos: E = getattr(orb, energy_type) is_MO = orb in mos orb_name = pyorbb.generate_label(orb) orb_index = orb.parent.orbitals.index(orb) # if our orbital is highlighted we draw an extra plot around it with a different color if orb in highlighted_orbitals: ax.plot([poss[orb]-level_width/2, poss[orb]+level_width/2], [E, E], c=highlight_color, linewidth=level_thickness + highlight_thickness, gid=f'{"MO" if is_MO else "FMO"}_{orb_index}') ax.plot([poss[orb]-level_width/2, poss[orb]+level_width/2], [E, E], c=level_color, linewidth=level_thickness, gid=f'{"MO" if is_MO else "FMO"}_{orb_index}') if (is_MO and draw_mo_labels) or (not is_MO and draw_sfo_labels): anchored_text(ax, poss[orb], E, orb_name, # [0, -arrow_length / 1.8 * energy_span], [0, orb_label_offset], ha='center', va='top', size=font_size, clip_on=True, gid=f'{"TEXTMO" if is_MO else "TEXTFMO"}_{orb_index}', fontname=font_name, color=level_color) if warning_orbs is not None: for w in warning_orbs: if w[0] is not orb: continue anchored_text(ax, poss[orb], E, '⚠', # [0, -arrow_length / 1.8 * energy_span], [orb_label_offset * 1.8, 0], ha='center', va='center', size=font_size*2.5, clip_on=True, gid=w[1], fontname=font_name, color='r') if not orb.occupied: continue for spin_part in orb.spin: break_on_one = False if spin_part == 'A': offset_x = -arrow_spacing displacement = arrow_length elif spin_part == 'B': offset_x = arrow_spacing displacement = -arrow_length if orb.spin == 'AB' and orb.occupation == 1: offset_x = 0 if orb.spin_pol in (0, 1): displacement = arrow_length elif orb.spin_pol == -1: displacement = -arrow_length break_on_one = True if orb.spin != 'AB': offset_x = 0 style = mpatches.ArrowStyle.CurveB( head_length=arrow_head_length, head_width=arrow_head_width ) ax.add_patch( arrow_tail_with_axes_offset( ax, (poss[orb], E), # anchor, [0, displacement], [offset_x, 0], arrowstyle=style, clip_on=True, gid=f'{"ARROWMO" if is_MO else "ARROWFMO"}_{orb_index}', color=arrow_color ) ) if break_on_one: break for fmo, mo in connections: psfo, pmo = poss[fmo], poss[mo] if psfo < pmo: psfo += level_width/2 pmo -= level_width/2 else: psfo -= level_width/2 pmo += level_width/2 sfo_index = fmo.parent.orbitals.index(fmo) mo_index = mo.parent.orbitals.index(mo) typ = connection_types.get((fmo, mo), 'Multiple') if typ == 'Multiple': c = multiple_color elif typ == 'PR': c = PR_color elif typ == 'OI': c = OI_color elif typ == 'Sanitization': c = sanitization_color ax.plot([psfo, pmo], [getattr(fmo, energy_type), mo.energy], c=c, linewidth=connection_width, alpha=np.clip(fmo.mulliken_contribution(mo), *alpha_range), gid=f'MIX_{sfo_index} -> {mo_index}', zorder=-10)