Team Ai
Apppublic

huggingface-projects/stable-diffusion-multiplayer

sourceHugging Faceupdated 3y agoView on Hugging Face
351likes
patch_match.py192 linesDownload Raw Back to PyPatchMatch
1#! /usr/bin/env python32# -*- coding: utf-8 -*-3# File   : patch_match.py4# Author : Jiayuan Mao5# Email  : maojiayuan@gmail.com6# Date   : 01/09/20207#8# Distributed under terms of the MIT license.9 10import ctypes11import os.path as osp12from typing import Optional, Union13 14import numpy as np15from PIL import Image16 17 18__all__ = ['set_random_seed', 'set_verbose', 'inpaint', 'inpaint_regularity']19 20 21class CShapeT(ctypes.Structure):22    _fields_ = [23        ('width', ctypes.c_int),24        ('height', ctypes.c_int),25        ('channels', ctypes.c_int),26    ]27 28 29class CMatT(ctypes.Structure):30    _fields_ = [31        ('data_ptr', ctypes.c_void_p),32        ('shape', CShapeT),33        ('dtype', ctypes.c_int)34    ]35 36 37PMLIB = ctypes.CDLL(osp.join(osp.dirname(__file__), 'libpatchmatch.so'))38 39PMLIB.PM_set_random_seed.argtypes = [ctypes.c_uint]40PMLIB.PM_set_verbose.argtypes = [ctypes.c_int]41PMLIB.PM_free_pymat.argtypes = [CMatT]42PMLIB.PM_inpaint.argtypes = [CMatT, CMatT, ctypes.c_int]43PMLIB.PM_inpaint.restype = CMatT44PMLIB.PM_inpaint_regularity.argtypes = [CMatT, CMatT, CMatT, ctypes.c_int, ctypes.c_float]45PMLIB.PM_inpaint_regularity.restype = CMatT46PMLIB.PM_inpaint2.argtypes = [CMatT, CMatT, CMatT, ctypes.c_int]47PMLIB.PM_inpaint2.restype = CMatT48PMLIB.PM_inpaint2_regularity.argtypes = [CMatT, CMatT, CMatT, CMatT, ctypes.c_int, ctypes.c_float]49PMLIB.PM_inpaint2_regularity.restype = CMatT50 51 52def set_random_seed(seed: int):53    PMLIB.PM_set_random_seed(ctypes.c_uint(seed))54 55 56def set_verbose(verbose: bool):57    PMLIB.PM_set_verbose(ctypes.c_int(verbose))58 59 60def inpaint(61    image: Union[np.ndarray, Image.Image],62    mask: Optional[Union[np.ndarray, Image.Image]] = None,63    *,64    global_mask: Optional[Union[np.ndarray, Image.Image]] = None,65    patch_size: int = 1566) -> np.ndarray:67    """68    PatchMatch based inpainting proposed in:69 70        PatchMatch : A Randomized Correspondence Algorithm for Structural Image Editing71        C.Barnes, E.Shechtman, A.Finkelstein and Dan B.Goldman72        SIGGRAPH 200973 74    Args:75        image (Union[np.ndarray, Image.Image]): the input image, should be 3-channel RGB/BGR.76        mask (Union[np.array, Image.Image], optional): the mask of the hole(s) to be filled, should be 1-channel.77        If not provided (None), the algorithm will treat all purely white pixels as the holes (255, 255, 255).78        global_mask (Union[np.array, Image.Image], optional): the target mask of the output image.79        patch_size (int): the patch size for the inpainting algorithm.80 81    Return:82        result (np.ndarray): the repaired image, of the same size as the input image.83    """84 85    if isinstance(image, Image.Image):86        image = np.array(image)87    image = np.ascontiguousarray(image)88    assert image.ndim == 3 and image.shape[2] == 3 and image.dtype == 'uint8'89 90    if mask is None:91        mask = (image == (255, 255, 255)).all(axis=2, keepdims=True).astype('uint8')92        mask = np.ascontiguousarray(mask)93    else:94        mask = _canonize_mask_array(mask)95 96    if global_mask is None:97        ret_pymat = PMLIB.PM_inpaint(np_to_pymat(image), np_to_pymat(mask), ctypes.c_int(patch_size))98    else:99        global_mask = _canonize_mask_array(global_mask)100        ret_pymat = PMLIB.PM_inpaint2(np_to_pymat(image), np_to_pymat(mask), np_to_pymat(global_mask), ctypes.c_int(patch_size))101 102    ret_npmat = pymat_to_np(ret_pymat)103    PMLIB.PM_free_pymat(ret_pymat)104 105    return ret_npmat106 107 108def inpaint_regularity(109    image: Union[np.ndarray, Image.Image],110    mask: Optional[Union[np.ndarray, Image.Image]],111    ijmap: np.ndarray,112    *,113    global_mask: Optional[Union[np.ndarray, Image.Image]] = None,114    patch_size: int = 15, guide_weight: float = 0.25115) -> np.ndarray:116    if isinstance(image, Image.Image):117        image = np.array(image)118    image = np.ascontiguousarray(image)119 120    assert isinstance(ijmap, np.ndarray) and ijmap.ndim == 3 and ijmap.shape[2] == 3 and ijmap.dtype == 'float32'121    ijmap = np.ascontiguousarray(ijmap)122 123    assert image.ndim == 3 and image.shape[2] == 3 and image.dtype == 'uint8'124    if mask is None:125        mask = (image == (255, 255, 255)).all(axis=2, keepdims=True).astype('uint8')126        mask = np.ascontiguousarray(mask)127    else:128        mask = _canonize_mask_array(mask)129 130 131    if global_mask is None:132        ret_pymat = PMLIB.PM_inpaint_regularity(np_to_pymat(image), np_to_pymat(mask), np_to_pymat(ijmap), ctypes.c_int(patch_size), ctypes.c_float(guide_weight))133    else:134        global_mask = _canonize_mask_array(global_mask)135        ret_pymat = PMLIB.PM_inpaint2_regularity(np_to_pymat(image), np_to_pymat(mask), np_to_pymat(global_mask), np_to_pymat(ijmap), ctypes.c_int(patch_size), ctypes.c_float(guide_weight))136 137    ret_npmat = pymat_to_np(ret_pymat)138    PMLIB.PM_free_pymat(ret_pymat)139 140    return ret_npmat141 142 143def _canonize_mask_array(mask):144    if isinstance(mask, Image.Image):145        mask = np.array(mask)146    if mask.ndim == 2 and mask.dtype == 'uint8':147        mask = mask[..., np.newaxis]148    assert mask.ndim == 3 and mask.shape[2] == 1 and mask.dtype == 'uint8'149    return np.ascontiguousarray(mask)150 151 152dtype_pymat_to_ctypes = [153    ctypes.c_uint8,154    ctypes.c_int8,155    ctypes.c_uint16,156    ctypes.c_int16,157    ctypes.c_int32,158    ctypes.c_float,159    ctypes.c_double,160]161 162 163dtype_np_to_pymat = {164    'uint8': 0,165    'int8': 1,166    'uint16': 2,167    'int16': 3,168    'int32': 4,169    'float32': 5,170    'float64': 6,171}172 173 174def np_to_pymat(npmat):175    assert npmat.ndim == 3176    return CMatT(177        ctypes.cast(npmat.ctypes.data, ctypes.c_void_p),178        CShapeT(npmat.shape[1], npmat.shape[0], npmat.shape[2]),179        dtype_np_to_pymat[str(npmat.dtype)]180    )181 182 183def pymat_to_np(pymat):184    npmat = np.ctypeslib.as_array(185        ctypes.cast(pymat.data_ptr, ctypes.POINTER(dtype_pymat_to_ctypes[pymat.dtype])),186        (pymat.shape.height, pymat.shape.width, pymat.shape.channels)187    )188    ret = np.empty(npmat.shape, npmat.dtype)189    ret[:] = npmat190    return ret191 192