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)