marclelarge/knn_encoder_decoder
1
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)