img_num = 0 IMG_PLOT = 16 plt.figure(figsize=(6,6)) table = wandb.Table(columns=["Image", "True Label", "Predicted Label"]) for img_batches, label in normalized_test_dataset.take(1): batches = img.shape[0] for b in range(batches): if img_num >= IMG_PLOT: break ax = plt.subplot(4,4,img_num+1) img = img_batches[b].numpy().squeeze() plt.imshow(img) plt.title(f"actual:{np.argmax(label[b].numpy())},pred :{y_pred[b]}",fontsize=6) img_num+=1 table.add_data(wandb.Image(img), np.argmax(label[b].numpy()), y_pred[b]) run.log({"predictions_table": table}) run.finish() __ __