Team Ai
Apppublic

David310/Detect_AI-generated_Image

sourceHugging Faceupdated 2y agoView on Hugging Face
4likes
validate.py313 linesDownload Raw Back to root
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