Plot histogram of top-k value from the given array. Args: tag (str): histogram title. array (tensor): a tensor to draw top k value from. k (int): number of top values to draw from array. Defaut to 10. class_names (list of strings, optional):
(tag, array, k=10, class_names=None, figsize=None)
| 90 | |
| 91 | |
| 92 | def plot_topk_histogram(tag, array, k=10, class_names=None, figsize=None): |
| 93 | """ |
| 94 | Plot histogram of top-k value from the given array. |
| 95 | Args: |
| 96 | tag (str): histogram title. |
| 97 | array (tensor): a tensor to draw top k value from. |
| 98 | k (int): number of top values to draw from array. |
| 99 | Defaut to 10. |
| 100 | class_names (list of strings, optional): |
| 101 | a list of names for values in array. |
| 102 | figsize (Optional[float, float]): the figure size of the confusion matrix. |
| 103 | If None, default to [6.4, 4.8]. |
| 104 | Returns: |
| 105 | fig (matplotlib figure): a matplotlib figure of the histogram. |
| 106 | """ |
| 107 | val, ind = torch.topk(array, k) |
| 108 | |
| 109 | fig = plt.Figure(figsize=figsize, facecolor="w", edgecolor="k") |
| 110 | |
| 111 | ax = fig.add_subplot(1, 1, 1) |
| 112 | |
| 113 | if class_names is None: |
| 114 | class_names = [str(i) for i in ind] |
| 115 | else: |
| 116 | class_names = [class_names[i] for i in ind] |
| 117 | |
| 118 | tick_marks = np.arange(k) |
| 119 | width = 0.75 |
| 120 | ax.bar( |
| 121 | tick_marks, |
| 122 | val, |
| 123 | width, |
| 124 | color="orange", |
| 125 | tick_label=class_names, |
| 126 | edgecolor="w", |
| 127 | linewidth=1, |
| 128 | ) |
| 129 | |
| 130 | ax.set_xlabel("Candidates") |
| 131 | ax.set_xticks(tick_marks) |
| 132 | ax.set_xticklabels(class_names, rotation=-45, ha="center") |
| 133 | ax.xaxis.set_label_position("bottom") |
| 134 | ax.xaxis.tick_bottom() |
| 135 | |
| 136 | y_tick = np.linspace(0, 1, num=10) |
| 137 | ax.set_ylabel("Frequency") |
| 138 | ax.set_yticks(y_tick) |
| 139 | y_labels = [format(i, ".1f") for i in y_tick] |
| 140 | ax.set_yticklabels(y_labels, ha="center") |
| 141 | |
| 142 | for i, v in enumerate(val.numpy()): |
| 143 | ax.text( |
| 144 | i - 0.1, |
| 145 | v + 0.03, |
| 146 | format(v, ".2f"), |
| 147 | color="orange", |
| 148 | fontweight="bold", |
| 149 | ) |
nothing calls this directly
no outgoing calls
no test coverage detected