baulab/Erasing-Concepts-In-Diffusion
49
1from PIL import Image2from matplotlib import pyplot as plt3import textwrap4 5 6def to_gif(images, path):7 8 images[0].save(path, save_all=True,9 append_images=images[1:], loop=0, duration=len(images) * 20)10 11 12def figure_to_image(figure):13 14 figure.set_dpi(300)15 16 figure.canvas.draw()17 18 return Image.frombytes('RGB', figure.canvas.get_width_height(), figure.canvas.tostring_rgb())19 20 21def image_grid(images, outpath=None, column_titles=None, row_titles=None):22 23 n_rows = len(images)24 n_cols = len(images[0])25 26 fig, axs = plt.subplots(nrows=n_rows, ncols=n_cols,27 figsize=(n_cols, n_rows), squeeze=False)28 29 for row, _images in enumerate(images):30 31 for column, image in enumerate(_images):32 ax = axs[row][column]33 ax.imshow(image)34 if column_titles and row == 0:35 ax.set_title(textwrap.fill(36 column_titles[column], width=12), fontsize='x-small')37 if row_titles and column == 0:38 ax.set_ylabel(row_titles[row], rotation=0, fontsize='x-small', labelpad=1.6 * len(row_titles[row]))39 ax.set_xticks([])40 ax.set_yticks([])41 42 plt.subplots_adjust(wspace=0, hspace=0)43 44 if outpath is not None:45 plt.savefig(outpath, bbox_inches='tight', dpi=300)46 plt.close()47 else:48 plt.tight_layout(pad=0)49 image = figure_to_image(plt.gcf())50 plt.close()51 return image52 53 54 55 56 57 58 59def get_module(module, module_name):60 61 if isinstance(module_name, str):62 module_name = module_name.split('.')63 64 if len(module_name) == 0:65 return module66 else:67 module = getattr(module, module_name[0])68 return get_module(module, module_name[1:])69 70 71def set_module(module, module_name, new_module):72 73 if isinstance(module_name, str):74 module_name = module_name.split('.')75 76 if len(module_name) == 1:77 return setattr(module, module_name[0], new_module)78 else:79 module = getattr(module, module_name[0])80 return set_module(module, module_name[1:], new_module)81 82 83def freeze(module):84 85 for parameter in module.parameters():86 87 parameter.requires_grad = False88 89 90def unfreeze(module):91 92 for parameter in module.parameters():93 94 parameter.requires_grad = True95 96 97def get_concat_h(im1, im2):98 dst = Image.new('RGB', (im1.width + im2.width, im1.height))99 dst.paste(im1, (0, 0))100 dst.paste(im2, (im1.width, 0))101 return dst102 103def get_concat_v(im1, im2):104 dst = Image.new('RGB', (im1.width, im1.height + im2.height))105 dst.paste(im1, (0, 0))106 dst.paste(im2, (0, im1.height))107 return dst