CoolFace
Apppublic

marclelarge/knn_encoder_decoder

sourceHugging Faceapache-2.0updated 4y agoView on Hugging Face
1likes
utils.py89 linesDownload Raw Back to root
1import numpy as np2import scipy3from PIL import Image4 5VALUE = 5126 7def resize(value,img):8  img = Image.open(img)9  #img = img.resize((value,value), Image.Resampling.LANCZOS)10  img.thumbnail((VALUE,VALUE), Image.Resampling.LANCZOS)11  return img12 13def get_mask(img,p):14    w,h=img.size15    return np.random.choice(a=[False, True], size=(w, h), p=[p, 1-p])16 17def generate_points(mask):18    (w,h) = mask.shape19    noise_points = []20    color_points = []21    for x in range(w):22        for y in range(h):23            if mask[x,y]:24                color_points.append(np.array([x,y]))25            else:26                noise_points.append(np.array([x,y]))27    return color_points, noise_points28 29def encoder_cp(img,color_points):30    w,h=img.size31    img2=Image.new('RGB',(w,h))32    for p in color_points:33        t = img.getpixel((p[0],p[1]))34        img2.putpixel((p[0],p[1]),(t[0],t[1],t[2]))35    return img236 37def encoder(img,p=0.95):38    img = resize(VALUE,img)39    mask = get_mask(img,p)40    c_p, n_p = generate_points(mask)41    return encoder_cp(img, c_p)42 43 44def get_points(img):45    w,h=img.size46    noise_points = []47    color_points = []48    for x in range(w):49        for y in range(h):50            t = img.getpixel((x,y))51            if np.sum(t[:3]) > 0 :52                color_points.append(np.array([x,y]))53            else:54                noise_points.append(np.array([x,y]))55    return color_points, noise_points56 57def restore(img, k, color_points, noise_points):58    kdtree = scipy.spatial.KDTree(color_points)59    for p in noise_points:60        _, knn_p = kdtree.query(p, k)61        r_m = []62        v_m = []63        b_m = []64        if k == 1:65            c_p = color_points[knn_p]66            t = img.getpixel((c_p[0],c_p[1]))67            img.putpixel((p[0],p[1]),(t[0],t[1],t[2]))68        else:69            for c_p in [color_points[j] for j in list(knn_p)]:70                t = img.getpixel((c_p[0],c_p[1]))71                r_m.append(t[0])72                v_m.append(t[1])73                b_m.append(t[2])74            r_m = int(sum(r_m)/k)75            v_m = int(sum(v_m)/k)76            b_m = int(sum(b_m)/k)77            img.putpixel((p[0],p[1]),(r_m,v_m,b_m))78    return img79 80def decoder(img,k=1):81    img = resize(VALUE,img)82    c_p, n_p = get_points(img)83    return restore(img,int(k),c_p,n_p)84 85def decoder_noise(img,k=1):86    img = Image.fromarray(img)87    #img = resize(VALUE,img)88    c_p, n_p = get_points(img)89    return restore(img,int(k),c_p,n_p)