MCPcopy Create free account
hub / github.com/OpenGVLab/UniFormerV2 / plot_topk_histogram

Function plot_topk_histogram

slowfast/visualization/utils.py:92–155  ·  view source on GitHub ↗

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)

Source from the content-addressed store, hash-verified

90
91
92def 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 )

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected