CoolFace
Apppublic

baulab/Erasing-Concepts-In-Diffusion

sourceHugging Facemitupdated 3y agoView on Hugging Face
49likes
util.py107 linesDownload Raw Back to root
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