import matplotlib.pyplot as plt import networkx as nx import matplotlib.colors as mcolors from matplotlib.colors import LinearSegmentedColormap def plot_graph(graph, length, beams, score): fig, ax = plt.subplots(figsize=(3+1.2*beams**length, max(5, 2+length)), dpi=300, facecolor='white') # Create positions for each node pos = nx.nx_agraph.graphviz_layout(graph, prog="dot") # Normalize the colors along the range of token scores if score == 'token': scores = [data['tokenscore'] for _, data in graph.nodes(data=True) if data['token'] is not None] elif score == 'sequence': scores = [data['sequencescore'] for _, data in graph.nodes(data=True) if data['token'] is not None] vmin = min(scores) vmax = max(scores) norm = mcolors.Normalize(vmin=vmin, vmax=vmax) cmap = LinearSegmentedColormap.from_list('rg', ["r", "y", "g"], N=256) # Draw the nodes nx.draw_networkx_nodes(graph, pos, node_size=2000, node_shape='o', alpha=1, linewidths=4, node_color=scores, cmap=cmap) # Draw the edges nx.draw_networkx_edges(graph, pos) # Draw the labels if score == 'token': labels = {node: data['token'].split('_')[0] + f"n{data['tokenscore']:.2f}%" for node, data in graph.nodes(data=True) if data['token'] is not None} elif score == 'sequence': labels = {node: data['token'].split('_')[0] + f"n{data['sequencescore']:.2f}" for node, data in graph.nodes(data=True) if data['token'] is not None} nx.draw_networkx_labels(graph, pos, labels=labels, font_size=10) plt.box(False) # Add a colorbar sm = plt.cm.ScalarMappable(cmap=cmap, norm=norm) sm.set_array([]) if score == 'token': fig.colorbar(sm, ax=ax, orientation='vertical', pad=0, label='Token probability (%)') elif score == 'sequence': fig.colorbar(sm, ax=ax, orientation='vertical', pad=0, label='Sequence score') plt.show() # Plot graph plot_graph(graph, length, 1.5, 'token')