snpushpi/boolean_interpretability
0
1import torch2import requests3from PIL import Image4from torchvision import transforms5import gradio as gr6import io7import requests8import torch 9import numpy as np10from torch.autograd import Variable11import torchvision.models as models12import torchvision.transforms as transforms13import math14import uuid15 16import numpy as np17import PIL.Image18import PIL.ImageDraw19import cv220 21 22model = torch.hub.load('pytorch/vision:v0.6.0', 'vgg16', pretrained=True).eval()23device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")24model = model.to(device)25# Download human-readable labels for ImageNet.26response = requests.get("https://git.io/JJkYN")27labels = response.text.split("\n")28 29import itertools30#compare two binary strings, check where there is one difference31def compBinary(s1,s2):32 count = 033 pos = 034 for i in range(len(s1)):35 if s1[i] != s2[i]:36 count+=137 pos = i38 if count == 1:39 return True, pos40 else:41 return False, None42 43 44#compare if the number is same as implicant term45#s1 should be the term46def compBinarySame(term,number):47 for i in range(len(term)):48 if term[i] != '-':49 if term[i] != number[i]:50 return False51 52 return True53 54 55#combine pairs and make new group56def combinePairs(group, unchecked):57 #define length58 l = len(group) -159 60 #check list61 check_list = []62 63 #create next group64 next_group = [[] for x in range(l)]65 66 #go through the groups67 for i in range(l):68 #first selected group69 for elem1 in group[i]:70 #next selected group71 for elem2 in group[i+1]:72 b, pos = compBinary(elem1, elem2)73 if b == True:74 #append the ones used in check list75 check_list.append(elem1)76 check_list.append(elem2)77 #replace the different bit with '-'78 new_elem = list(elem1)79 new_elem[pos] = '-'80 new_elem = "".join(new_elem)81 next_group[i].append(new_elem)82 for i in group:83 for j in i:84 if j not in check_list:85 unchecked.append(j)86 87 return next_group, unchecked88 89 90#remove redundant lists in 2d list91def remove_redundant(group):92 new_group = []93 for j in group:94 new=[]95 for i in j:96 if i not in new:97 new.append(i)98 new_group.append(new)99 return new_group100 101 102#remove redundant in 1d list103def remove_redundant_list(list):104 new_list = []105 for i in list:106 if i not in new_list:107 new_list.append(i)108 return new_list109 110 111#return True if empty112def check_empty(group):113 114 if len(group) == 0:115 return True116 117 else:118 count = 0119 for i in group:120 if i:121 count+=1122 if count == 0:123 return True124 return False125 126 127#find essential prime implicants ( col num of ones = 1)128def find_prime(Chart):129 prime = []130 for col in range(len(Chart[0])):131 count = 0132 pos = 0133 for row in range(len(Chart)):134 #find essential135 if Chart[row][col] == 1:136 count += 1137 pos = row138 139 if count == 1:140 prime.append(pos)141 142 return prime143 144def check_all_zero(Chart):145 for i in Chart:146 for j in i:147 if j != 0:148 return False149 return True150 151#find max value in list152def find_max(l):153 max = -1154 index = 0155 for i in range(len(l)):156 if l[i] > max:157 max = l[i]158 index = i159 return index160 161#multiply two terms (ex. (p1 + p2)(p1+p4+p5) )..it returns the product162def multiplication(list1, list2):163 list_result = []164 #if empty165 if len(list1) == 0 and len(list2)== 0:166 return list_result167 #if one is empty168 elif len(list1)==0:169 return list2170 #if another is empty171 elif len(list2)==0:172 return list1173 174 #both not empty175 else:176 for i in list1:177 for j in list2:178 #if two term same179 if i == j:180 #list_result.append(sorted(i))181 list_result.append(i)182 else:183 #list_result.append(sorted(list(set(i+j))))184 list_result.append(list(set(i+j)))185 186 #sort and remove redundant lists and return this list187 list_result.sort()188 return list(list_result for list_result,_ in itertools.groupby(list_result))189 190#petrick's method191def petrick_method(Chart):192 #initial P193 P = []194 for col in range(len(Chart[0])):195 p =[]196 for row in range(len(Chart)):197 if Chart[row][col] == 1:198 p.append([row])199 P.append(p)200 #do multiplication201 for l in range(len(P)-1):202 P[l+1] = multiplication(P[l],P[l+1])203 204 P = sorted(P[len(P)-1],key=len)205 final = []206 #find the terms with min length = this is the one with lowest cost (optimized result)207 min=len(P[0])208 for i in P:209 if len(i) == min:210 final.append(i)211 else:212 break213 #final is the result of petrick's method214 return final215 216#chart = n*n list217def find_minimum_cost(Chart, unchecked):218 P_final = []219 #essential_prime = list with terms with only one 1 (Essential Prime Implicants)220 essential_prime = find_prime(Chart)221 essential_prime = remove_redundant_list(essential_prime)222 223 #print out the essential primes224 if len(essential_prime)>0:225 s = "\nEssential Prime Implicants :\n"226 for i in range(len(unchecked)):227 for j in essential_prime:228 if j == i:229 s= s+binary_to_letter(unchecked[i])+' , '230 #print(s[:(len(s)-3)])231 232 #modifiy the chart to exclude the covered terms233 for i in range(len(essential_prime)):234 for col in range(len(Chart[0])):235 if Chart[essential_prime[i]][col] == 1:236 for row in range(len(Chart)):237 Chart[row][col] = 0238 239 #if all zero, no need for petrick method240 if check_all_zero(Chart) == True:241 P_final = [essential_prime]242 else:243 #petrick's method244 P = petrick_method(Chart)245 246 #find the one with minimum cost247 #see "Introduction to Logic Design" - Alan B.Marcovitz Example 4.6 pg 213248 '''249 Although Petrick's method gives the minimum terms that cover all,250 it does not mean that it is the solution for minimum cost!251 '''252 253 P_cost = []254 for prime in P:255 count = 0256 for i in range(len(unchecked)):257 for j in prime:258 if j == i:259 count = count+ cal_efficient(unchecked[i])260 P_cost.append(count)261 262 263 for i in range(len(P_cost)):264 if P_cost[i] == min(P_cost):265 P_final.append(P[i])266 267 #append prime implicants to the solution of Petrick's method268 for i in P_final:269 for j in essential_prime:270 if j not in i:271 i.append(j)272 273 return P_final274 275#calculate the number of literals276def cal_efficient(s):277 count = 0278 for i in range(len(s)):279 if s[i] != '-':280 count+=1281 282 return count283 284#print the binary code to letter285def binary_to_letter(s):286 out = ''287 c = 'a'288 more = False289 n = 0290 for i in range(len(s)):291 #if it is a range a-zA-Z292 if more == False:293 if s[i] == '1':294 out = out + c295 elif s[i] == '0':296 out = out + c+'\''297 298 if more == True:299 if s[i] == '1':300 out = out + c + str(n)301 elif s[i] == '0':302 out = out + c + str(n) + '\''303 n+=1304 #conditions for next operations305 if c=='z' and more == False:306 c = 'A'307 elif c=='Z':308 c = 'a'309 more = True310 311 elif more == False:312 c = chr(ord(c)+1)313 return out314 315def binary_to_letter_final(s, dictionary):316 out = ''317 c = 'a'318 more = False319 n = 0320 for i in range(len(s)):321 #if it is a range a-zA-Z322 if more == False:323 if s[i] == '1':324 out = out + c325 elif s[i] == '0':326 out = out + c+'\''327 328 if more == True:329 if s[i] == '1':330 out = out + c + str(n)331 elif s[i] == '0':332 out = out + c + str(n) + '\''333 n+=1334 #conditions for next operations335 if c=='z' and more == False:336 c = 'A'337 elif c=='Z':338 c = 'a'339 more = True340 341 elif more == False:342 c = chr(ord(c)+1)343 return_string = ''344 for char in out:345 if char == "'":346 return_string += "'"347 else:348 return_string+='('+dictionary[char]+')'349 return return_string350 351 352 353#main function354def quin_macluskey(n_var,minterms, dictionary):355 a = minterms356 #make a group list357 group = [[] for x in range(n_var+1)]358 359 for i in range(len(a)):360 #convert to binary361 a[i] = bin(a[i])[2:]362 if len(a[i]) < n_var:363 #add zeros to fill the n-bits364 for j in range(n_var - len(a[i])):365 a[i] = '0'+ a[i]366 #if incorrect input367 elif len(a[i]) > n_var:368 print('\nError : Choose the correct number of variables(bits)\n')369 return370 #count the num of 1371 index = a[i].count('1')372 #group by num of 1 separately373 group[index].append(a[i])374 375 376 all_group=[]377 unchecked = []378 #combine the pairs in series until nothing new can be combined379 while check_empty(group) == False:380 all_group.append(group)381 next_group, unchecked = combinePairs(group,unchecked)382 group = remove_redundant(next_group)383 384 s = "\nPrime Implicants :\n"385 for i in unchecked:386 s= s + binary_to_letter(i) + " , "387 #print(s[:(len(s)-3)])388 389 #make the prime implicant chart390 Chart = [[0 for x in range(len(a))] for x in range(len(unchecked))]391 392 for i in range(len(a)):393 for j in range (len(unchecked)):394 #term is same as number395 if compBinarySame(unchecked[j], a[i]):396 Chart[j][i] = 1397 398 #prime contains the index of the prime implicant terms399 #prime = remove_redundant_list(find_minimum_cost(Chart))400 primes = find_minimum_cost(Chart, unchecked)401 primes = remove_redundant(primes)402 403 404 print("\n-- Answers --\n")405 result = ''406 for prime in primes:407 s=''408 for i in range(len(unchecked)):409 for j in prime:410 if j == i:411 s= s+binary_to_letter_final(unchecked[i], dictionary)+' + '412 result += s[:(len(s)-3)]413 #print(result)414 return result415 416 #This part shortens boolean formulas even more than the last step, applicable only in the context of image prediction boolean formulas417def pre_image(res_image,model):418 model.eval()419 res_image = res_image[:,:,::-1]420 img = Image.fromarray(res_image)421 normalize = transforms.Normalize(mean=[0.485, 0.456, 0.406],422 std=[0.229, 0.224, 0.225])423 transform_pipeline = transforms.Compose([424 transforms.Resize(256),425 transforms.CenterCrop(224),426 transforms.ToTensor(),427 normalize,428 ])429 img = transform_pipeline(img)430 431 # PyTorch pretrained models expect the Tensor dims to be (num input imgs, num color channels, height, width).432 # Currently however, we have (num color channels, height, width); let's fix this by inserting a new axis.433 img = img.unsqueeze(0) # Insert the new axis at index 0 i.e. in front of the other axes/dims. 434 435 # Now that we have preprocessed our img, we need to convert it into a 436 # Variable; PyTorch models expect inputs to be Variables. A PyTorch Variable is a 437 # wrapper around a PyTorch Tensor.438 img = Variable(img)439 img = img.to(device)440 441 output = model(img) # Returns a Tensor of shape (batch, num class labels)442 _, preds = torch.max(output, 1)443 preds = preds.detach().cpu().numpy()444 return preds[0]445 446def pre_image2(img,model):447 model.eval()448 #img = Image.open(img)449 normalize = transforms.Normalize(mean=[0.485, 0.456, 0.406],450 std=[0.229, 0.224, 0.225])451 transform_pipeline = transforms.Compose([452 transforms.Resize(256),453 transforms.CenterCrop(224),454 transforms.ToTensor(),455 normalize,456 ])457 img = transform_pipeline(img)458 459 # PyTorch pretrained models expect the Tensor dims to be (num input imgs, num color channels, height, width).460 # Currently however, we have (num color channels, height, width); let's fix this by inserting a new axis.461 img = img.unsqueeze(0) # Insert the new axis at index 0 i.e. in front of the other axes/dims. 462 463 # Now that we have preprocessed our img, we need to convert it into a 464 # Variable; PyTorch models expect inputs to be Variables. A PyTorch Variable is a 465 # wrapper around a PyTorch Tensor.466 img = Variable(img)467 img = img.to(device)468 469 output = model(img) # Returns a Tensor of shape (batch, num class labels)470 _, preds = torch.max(output, 1)471 preds = preds.detach().cpu().numpy()472 return preds[0]473 474def findsubsets(s, n):475 return list(itertools.combinations(s, n))476def find_all_subsets(s):477 res = []478 for i in range(1,len(s)+1):479 res+= findsubsets(s,i)480 return res+[()] 481 482def bitwise_general_or(comb,elt_mask_list):483 res_mask = elt_mask_list[comb[0]]484 for i in range(1,len(comb)):485 new_mask = elt_mask_list[comb[i]]486 res_mask = cv2.bitwise_and(res_mask,new_mask)487 return res_mask488 489def overlap(res_mask,main_img):490 res_mask = np.mean(res_mask,axis=2)491 bw_mask = np.array(res_mask,dtype=np.uint8)492 img = np.array(main_img)493 rows,cols,channels = img.shape494 roi = img[0:rows,0:cols]495 img1_bg = cv2.bitwise_and(roi,roi,mask=bw_mask)496 return img1_bg 497 498def mask_generation(main_img,painted_img): 499 img1 = cv2.imread(main_img)500 img2 = cv2.imread(painted_img)501 k,l,m = img1.shape502 mat_3 = np.full((k,l,m),255)503 mat_3[(img1[:,:,0]!=img2[:,:,0]) | (img1[:,:,1]!=img2[:,:,1]) | (img1[:,:,2]!=img2[:,:,2])]=0504 return mat_3505 506def calculate_weight(a_list):507 running_sum = 0508 for e in a_list:509 running_sum+=2**e510 return running_sum511 512def shape_to_mask(513 img_shape, points, shape_type=None, line_width=10, point_size=5514):515 mask = np.zeros(img_shape[:2], dtype=np.uint8)516 mask = PIL.Image.fromarray(mask)517 draw = PIL.ImageDraw.Draw(mask)518 xy = [tuple(point) for point in points]519 if shape_type == "circle":520 assert len(xy) == 2, "Shape of shape_type=circle must have 2 points"521 (cx, cy), (px, py) = xy522 d = math.sqrt((cx - px) ** 2 + (cy - py) ** 2)523 draw.ellipse([cx - d, cy - d, cx + d, cy + d], outline=1, fill=1)524 elif shape_type == "rectangle":525 assert len(xy) == 2, "Shape of shape_type=rectangle must have 2 points"526 draw.rectangle(xy, outline=1, fill=1)527 elif shape_type == "line":528 assert len(xy) == 2, "Shape of shape_type=line must have 2 points"529 draw.line(xy=xy, fill=1, width=line_width)530 elif shape_type == "linestrip":531 draw.line(xy=xy, fill=1, width=line_width)532 elif shape_type == "point":533 assert len(xy) == 1, "Shape of shape_type=point must have 1 points"534 cx, cy = xy[0]535 r = point_size536 draw.ellipse([cx - r, cy - r, cx + r, cy + r], outline=1, fill=1)537 else:538 assert len(xy) > 2, "Polygon must have points more than 2"539 draw.polygon(xy=xy, outline=1, fill=1)540 mask = np.array(mask, dtype=bool)541 return mask542 543def combine_similar_masks(dj,shape):544 j = 0545 for i in range(len(dj['shapes'])):546 if dj['shapes'][i]["label"]==shape:547 j+=1548 curr_mask = shape_to_mask((dj['imageHeight'],dj['imageWidth']), dj['shapes'][i]['points'], shape_type=None,line_width=1, point_size=1)549 mask_img = curr_mask.astype(np.int)#boolean to 0,Convert to 1550 mask_img = 255*mask_img551 mask_img = np.stack((mask_img, mask_img, mask_img), axis=2)552 mask_img = 255 - mask_img553 if j==1:554 result_mask= mask_img555 else:556 result_mask = cv2.bitwise_and(mask_img, result_mask)557 return result_mask 558 559def combine_all_masks(dj):560 j = 0561 for i in range(len(dj['shapes'])):562 j+=1563 curr_mask = shape_to_mask((dj['imageHeight'],dj['imageWidth']), dj['shapes'][i]['points'], shape_type=None,line_width=1, point_size=1)564 mask_img = curr_mask.astype(np.int)#boolean to 0,Convert to 1565 mask_img = 255*mask_img566 mask_img = np.stack((mask_img, mask_img, mask_img), axis=2)567 mask_img = 255 - mask_img568 if j==1:569 result_mask= mask_img570 else:571 result_mask = cv2.bitwise_and(mask_img, result_mask)572 573 return result_mask574 575def get_key(my_dict,val):576 for key, value in my_dict.items():577 if val == value:578 return key579 580import json581import tempfile582#TODO: Add a region for bg, so add a shape in unique_shape_set, write a function forcombiningall masks and doing the not to found the 583#bg mask and then the rest is the same584def generate_predictions(path, main_img):585 print(main_img)586 #generate all elemantary masks587 #if flag=True then combine similar label masks, else no588 #with open(path, "r",encoding="utf-8") as f:589 path[0].seek(0)590 my_bytes = path[0].read()591 string_data = my_bytes.decode('utf8')592 dj = json.loads(string_data)593 bit_assignment = {}594 elt_mask_dict = {}595 unique_shape_set = set()596 for elt in dj["shapes"]:597 unique_shape_set.add(elt["label"])598 #unique_shape_set.add('bg')599 l = len(unique_shape_set)600 i = l-1601 for shape in unique_shape_set:602 if shape=='bg':603 intermed_mask = combine_all_masks(dj)604 mask_img = 255 - intermed_mask605 bit_assignment[shape]=i606 elt_mask_dict[i] = mask_img607 i = i-1608 else:609 mask_img = combine_similar_masks(dj,shape)610 bit_assignment[shape]= i611 elt_mask_dict[i]= mask_img612 i=i-1613 614 true_prediction = pre_image2(main_img, model)615 all_file_comb_list = find_all_subsets(elt_mask_dict)616 minterms = []617 bit_list = list(elt_mask_dict.keys())618 print(len(all_file_comb_list))619 for comb in all_file_comb_list:620 print('ho')621 if comb == ():622 res_image = np.array(main_img)623 prediction = true_prediction624 else:625 res_mask = bitwise_general_or(comb,elt_mask_dict)626 res_image = overlap(res_mask,main_img)627 prediction = pre_image(res_image,model)628 if prediction==true_prediction:629 region_present = [e for e in bit_list if not e in list(comb)]630 minterms.append(calculate_weight(region_present))631 dictionary = {}632 ch = 'a'633 for i in range(l):634 dictionary[chr(ord(ch) -i+l-1)]=get_key(bit_assignment, i)635 #print(dictionary)636 print('no') 637 result = quin_macluskey(l,minterms, dictionary)638 print(result)639 return result640 641gr.Interface(fn=generate_predictions, 642 inputs=["files", "pil"],643 outputs="text").launch(share= True, debug= True)644 