encoder = keras.Model(input_img, encoded) encoded_imgs = encoder.predict(x_test) n = 10 plt.figure(figsize=(20, 8)) for i in range(1, n + 1): ax = plt.subplot(1, n, i) plt.imshow(encoded_imgs[i].reshape((4, 4 * 8)).T) plt.gray() ax.get_xaxis().set_visible(False) ax.get_yaxis().set_visible(False) plt.show()