# Codeblock 27 fig, axes = plt.subplots(ncols=10, figsize=(24, 8)) sample_no = 0 timestep_no = 0 for i in range(10): axes[i].imshow(denoised_images[timestep_no][sample_no].squeeze().detach().cpu().numpy(), cmap='gray') axes[i].get_xaxis().set_visible(False) axes[i].get_yaxis().set_visible(False) timestep_no += 1 plt.show()