Given all predictions and all true labels, plot histograms of top-k most frequently predicted classes for each true class. Args: writer (SummaryWriter object): a tensorboard SummaryWriter object. cmtx (ndarray): confusion matrix. num_classes (int): total number
(
writer,
cmtx,
num_classes,
k=10,
global_step=None,
subset_ids=None,
class_names=None,
figsize=None,
)
| 281 | |
| 282 | |
| 283 | def plot_hist( |
| 284 | writer, |
| 285 | cmtx, |
| 286 | num_classes, |
| 287 | k=10, |
| 288 | global_step=None, |
| 289 | subset_ids=None, |
| 290 | class_names=None, |
| 291 | figsize=None, |
| 292 | ): |
| 293 | """ |
| 294 | Given all predictions and all true labels, plot histograms of top-k most |
| 295 | frequently predicted classes for each true class. |
| 296 | |
| 297 | Args: |
| 298 | writer (SummaryWriter object): a tensorboard SummaryWriter object. |
| 299 | cmtx (ndarray): confusion matrix. |
| 300 | num_classes (int): total number of classes. |
| 301 | k (int): top k to plot histograms. |
| 302 | global_step (Optional[int]): current step. |
| 303 | subset_ids (list of ints, optional): class indices to plot histogram. |
| 304 | mapping (list of strings): names of all classes. |
| 305 | figsize (Optional[float, float]): the figure size of the confusion matrix. |
| 306 | If None, default to [6.4, 4.8]. |
| 307 | """ |
| 308 | if subset_ids is None or len(subset_ids) != 0: |
| 309 | if subset_ids is None: |
| 310 | subset_ids = set(range(num_classes)) |
| 311 | else: |
| 312 | subset_ids = set(subset_ids) |
| 313 | # If class names are not provided, use their indices as names. |
| 314 | if class_names is None: |
| 315 | class_names = list(range(num_classes)) |
| 316 | |
| 317 | for i in subset_ids: |
| 318 | pred = cmtx[i] |
| 319 | hist = vis_utils.plot_topk_histogram( |
| 320 | class_names[i], |
| 321 | torch.Tensor(pred), |
| 322 | k, |
| 323 | class_names, |
| 324 | figsize=figsize, |
| 325 | ) |
| 326 | writer.add_figure( |
| 327 | tag="Top {} predictions by classes/{}".format( |
| 328 | k, class_names[i] |
| 329 | ), |
| 330 | figure=hist, |
| 331 | global_step=global_step, |
| 332 | ) |
| 333 | |
| 334 | |
| 335 | def add_ndim_array( |