David310/Detect_AI-generated_Image
4
1import argparse2from ast import arg3import os4import csv5import torch6import torchvision.transforms as transforms7import torch.utils.data8import numpy as np9from sklearn.metrics import average_precision_score, precision_recall_curve, accuracy_score10from torch.utils.data import Dataset11import sys12from models import get_model13from PIL import Image 14import pickle15from tqdm import tqdm16from io import BytesIO17from copy import deepcopy18from dataset_paths import DATASET_PATHS19import random20import shutil21from scipy.ndimage.filters import gaussian_filter22 23SEED = 024def set_seed():25 torch.manual_seed(SEED)26 torch.cuda.manual_seed(SEED)27 np.random.seed(SEED)28 random.seed(SEED)29 30 31MEAN = {32 "imagenet":[0.485, 0.456, 0.406],33 "clip":[0.48145466, 0.4578275, 0.40821073]34}35 36STD = {37 "imagenet":[0.229, 0.224, 0.225],38 "clip":[0.26862954, 0.26130258, 0.27577711]39}40 41 42 43 44 45def find_best_threshold(y_true, y_pred):46 "We assume first half is real 0, and the second half is fake 1"47 48 N = y_true.shape[0]49 50 if y_pred[0:N//2].max() <= y_pred[N//2:N].min(): # perfectly separable case51 return (y_pred[0:N//2].max() + y_pred[N//2:N].min()) / 2 52 53 best_acc = 0 54 best_thres = 0 55 for thres in y_pred:56 temp = deepcopy(y_pred)57 temp[temp>=thres] = 1 58 temp[temp<thres] = 0 59 60 acc = (temp == y_true).sum() / N 61 if acc >= best_acc:62 best_thres = thres63 best_acc = acc 64 65 return best_thres66 67 68 69def png2jpg(img, quality):70 out = BytesIO()71 img.save(out, format='jpeg', quality=quality) # ranging from 0-95, 75 is default72 img = Image.open(out)73 # load from memory before ByteIO closes74 img = np.array(img)75 out.close()76 return Image.fromarray(img)77 78 79def gaussian_blur(img, sigma):80 img = np.array(img)81 82 gaussian_filter(img[:,:,0], output=img[:,:,0], sigma=sigma)83 gaussian_filter(img[:,:,1], output=img[:,:,1], sigma=sigma)84 gaussian_filter(img[:,:,2], output=img[:,:,2], sigma=sigma)85 86 return Image.fromarray(img)87 88 89 90def calculate_acc(y_true, y_pred, thres):91 r_acc = accuracy_score(y_true[y_true==0], y_pred[y_true==0] > thres)92 f_acc = accuracy_score(y_true[y_true==1], y_pred[y_true==1] > thres)93 acc = accuracy_score(y_true, y_pred > thres)94 return r_acc, f_acc, acc 95 96 97def validate(model, loader, find_thres=False):98 99 with torch.no_grad():100 y_true, y_pred = [], []101 print ("Length of dataset: %d" %(len(loader)))102 for img, label in loader:103 in_tens = img.cuda()104 105 y_pred.extend(model(in_tens).sigmoid().flatten().tolist())106 y_true.extend(label.flatten().tolist())107 108 y_true, y_pred = np.array(y_true), np.array(y_pred)109 110 # ================== save this if you want to plot the curves =========== # 111 # torch.save( torch.stack( [torch.tensor(y_true), torch.tensor(y_pred)] ), 'baseline_predication_for_pr_roc_curve.pth' )112 # exit()113 # =================================================================== #114 115 # Get AP 116 ap = average_precision_score(y_true, y_pred)117 118 # Acc based on 0.5119 r_acc0, f_acc0, acc0 = calculate_acc(y_true, y_pred, 0.5)120 if not find_thres:121 return ap, r_acc0, f_acc0, acc0122 123 124 # Acc based on the best thres125 best_thres = find_best_threshold(y_true, y_pred)126 r_acc1, f_acc1, acc1 = calculate_acc(y_true, y_pred, best_thres)127 128 return ap, r_acc0, f_acc0, acc0, r_acc1, f_acc1, acc1, best_thres129 130 131 132 133 134 135# = = = = = = = = = = = = = = = = = = = = = = = = = = = = = = = = = = = = # 136 137 138 139 140def recursively_read(rootdir, must_contain, exts=["png", "jpg", "JPEG", "jpeg", "bmp"]):141 out = [] 142 for r, d, f in os.walk(rootdir):143 for file in f:144 if (file.split('.')[1] in exts) and (must_contain in os.path.join(r, file)):145 out.append(os.path.join(r, file))146 return out147 148 149def get_list(path, must_contain=''):150 if ".pickle" in path:151 with open(path, 'rb') as f:152 image_list = pickle.load(f)153 image_list = [ item for item in image_list if must_contain in item ]154 else:155 image_list = recursively_read(path, must_contain)156 return image_list157 158 159 160 161 162class RealFakeDataset(Dataset):163 def __init__(self, real_path, 164 fake_path, 165 data_mode, 166 max_sample,167 arch,168 jpeg_quality=None,169 gaussian_sigma=None):170 171 assert data_mode in ["wang2020", "ours"]172 self.jpeg_quality = jpeg_quality173 self.gaussian_sigma = gaussian_sigma174 175 # = = = = = = data path = = = = = = = = = # 176 if type(real_path) == str and type(fake_path) == str:177 real_list, fake_list = self.read_path(real_path, fake_path, data_mode, max_sample)178 else:179 real_list = []180 fake_list = []181 for real_p, fake_p in zip(real_path, fake_path):182 real_l, fake_l = self.read_path(real_p, fake_p, data_mode, max_sample)183 real_list += real_l184 fake_list += fake_l185 186 self.total_list = real_list + fake_list187 188 189 # = = = = = = label = = = = = = = = = # 190 191 self.labels_dict = {}192 for i in real_list:193 self.labels_dict[i] = 0194 for i in fake_list:195 self.labels_dict[i] = 1196 197 stat_from = "imagenet" if arch.lower().startswith("imagenet") else "clip"198 self.transform = transforms.Compose([199 transforms.CenterCrop(224),200 transforms.ToTensor(),201 transforms.Normalize( mean=MEAN[stat_from], std=STD[stat_from] ),202 ])203 204 205 def read_path(self, real_path, fake_path, data_mode, max_sample):206 207 if data_mode == 'wang2020':208 real_list = get_list(real_path, must_contain='0_real')209 fake_list = get_list(fake_path, must_contain='1_fake')210 else:211 real_list = get_list(real_path)212 fake_list = get_list(fake_path)213 214 215 if max_sample is not None:216 if (max_sample > len(real_list)) or (max_sample > len(fake_list)):217 max_sample = 100218 print("not enough images, max_sample falling to 100")219 random.shuffle(real_list)220 random.shuffle(fake_list)221 real_list = real_list[0:max_sample]222 fake_list = fake_list[0:max_sample]223 224 assert len(real_list) == len(fake_list) 225 226 return real_list, fake_list227 228 229 230 def __len__(self):231 return len(self.total_list)232 233 def __getitem__(self, idx):234 235 img_path = self.total_list[idx]236 237 label = self.labels_dict[img_path]238 img = Image.open(img_path).convert("RGB")239 240 if self.gaussian_sigma is not None:241 img = gaussian_blur(img, self.gaussian_sigma) 242 if self.jpeg_quality is not None:243 img = png2jpg(img, self.jpeg_quality)244 245 img = self.transform(img)246 return img, label247 248 249 250 251 252if __name__ == '__main__':253 254 255 parser = argparse.ArgumentParser(formatter_class=argparse.ArgumentDefaultsHelpFormatter)256 parser.add_argument('--real_path', type=str, default=None, help='dir name or a pickle')257 parser.add_argument('--fake_path', type=str, default=None, help='dir name or a pickle')258 parser.add_argument('--data_mode', type=str, default=None, help='wang2020 or ours')259 parser.add_argument('--max_sample', type=int, default=1000, help='only check this number of images for both fake/real')260 261 parser.add_argument('--arch', type=str, default='res50')262 parser.add_argument('--ckpt', type=str, default='./pretrained_weights/fc_weights.pth')263 264 parser.add_argument('--result_folder', type=str, default='result', help='')265 parser.add_argument('--batch_size', type=int, default=128)266 267 parser.add_argument('--jpeg_quality', type=int, default=None, help="100, 90, 80, ... 30. Used to test robustness of our model. Not apply if None")268 parser.add_argument('--gaussian_sigma', type=int, default=None, help="0,1,2,3,4. Used to test robustness of our model. Not apply if None")269 270 271 opt = parser.parse_args()272 273 274 if os.path.exists(opt.result_folder):275 shutil.rmtree(opt.result_folder)276 os.makedirs(opt.result_folder)277 278 model = get_model(opt.arch)279 state_dict = torch.load(opt.ckpt, map_location='cpu')280 model.fc.load_state_dict(state_dict)281 print ("Model loaded..")282 model.eval()283 model.cuda()284 285 if (opt.real_path == None) or (opt.fake_path == None) or (opt.data_mode == None):286 dataset_paths = DATASET_PATHS287 else:288 dataset_paths = [ dict(real_path=opt.real_path, fake_path=opt.fake_path, data_mode=opt.data_mode) ]289 290 291 292 for dataset_path in (dataset_paths):293 set_seed()294 295 dataset = RealFakeDataset( dataset_path['real_path'], 296 dataset_path['fake_path'], 297 dataset_path['data_mode'], 298 opt.max_sample, 299 opt.arch,300 jpeg_quality=opt.jpeg_quality, 301 gaussian_sigma=opt.gaussian_sigma,302 )303 304 loader = torch.utils.data.DataLoader(dataset, batch_size=opt.batch_size, shuffle=False, num_workers=4)305 ap, r_acc0, f_acc0, acc0, r_acc1, f_acc1, acc1, best_thres = validate(model, loader, find_thres=True)306 307 with open( os.path.join(opt.result_folder,'ap.txt'), 'a') as f:308 f.write(dataset_path['key']+': ' + str(round(ap*100, 2))+'\n' )309 310 with open( os.path.join(opt.result_folder,'acc0.txt'), 'a') as f:311 f.write(dataset_path['key']+': ' + str(round(r_acc0*100, 2))+' '+str(round(f_acc0*100, 2))+' '+str(round(acc0*100, 2))+'\n' )312 313 