def img_flip_vertical(img): return tf.image.flip_left_right(img) test_flip = (train_data.map(lambda img, label: img_flip_vertical(img))) plt.figure(figsize=(6,6)) for img in test_flip.take(1): for num in range(16): ax = plt.subplot(4,4,num+1) plt.imshow((img[num].numpy()).astype("uint8")) plt.axis("off") plt.suptitle("Flip Augmentation") plt.tight_layout() plt.show() __ __