Team Ai
Apppublic

David310/Detect_AI-generated_Image

sourceHugging Faceupdated 2y agoView on Hugging Face
4likes
detect_one_image.py333 linesDownload Raw Back to root
1import argparse2from ast import arg3import os4import csv5import torch6import torchvision.transforms as transforms7import torch.utils.data8import numpy as np9# from 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 shutil21# from 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 """67def png2jpg(img, quality):68    out = BytesIO()69    img.save(out, format='jpeg', quality=quality) # ranging from 0-95, 75 is default70    img = Image.open(out)71    # load from memory before ByteIO closes72    img = np.array(img)73    out.close()74    return Image.fromarray(img)75"""76def gaussian_blur(img, sigma):77    img = np.array(img)78 79    gaussian_filter(img[:,:,0], output=img[:,:,0], sigma=sigma)80    gaussian_filter(img[:,:,1], output=img[:,:,1], sigma=sigma)81    gaussian_filter(img[:,:,2], output=img[:,:,2], sigma=sigma)82 83    return Image.fromarray(img)84 85def calculate_acc(y_true, y_pred, thres):86    r_acc = accuracy_score(y_true[y_true==0], y_pred[y_true==0] > thres)87    f_acc = accuracy_score(y_true[y_true==1], y_pred[y_true==1] > thres)88    acc = accuracy_score(y_true, y_pred > thres)89    return r_acc, f_acc, acc    90"""91 92 93 94def validate(model, loader, find_thres=False):95 96    with torch.no_grad():97        y_true, y_pred = [], []98        print ("Length of dataset: %d" %(len(loader)))99        for img, label in loader:100            in_tens = img.cuda()101 102            y_pred.extend(model(in_tens).sigmoid().flatten().tolist())103            y_true.extend(label.flatten().tolist())104 105    y_true, y_pred = np.array(y_true), np.array(y_pred)106 107    # ================== save this if you want to plot the curves =========== # 108    # torch.save( torch.stack( [torch.tensor(y_true), torch.tensor(y_pred)] ),  'baseline_predication_for_pr_roc_curve.pth' )109    # exit()110    # =================================================================== #111    112    # Get AP 113    ap = average_precision_score(y_true, y_pred)114 115    # Acc based on 0.5116    r_acc0, f_acc0, acc0 = calculate_acc(y_true, y_pred, 0.5)117    if not find_thres:118        return ap, r_acc0, f_acc0, acc0119 120 121    # Acc based on the best thres122    best_thres = find_best_threshold(y_true, y_pred)123    r_acc1, f_acc1, acc1 = calculate_acc(y_true, y_pred, best_thres)124 125    return ap, r_acc0, f_acc0, acc0, r_acc1, f_acc1, acc1, best_thres126 127 128def detect_one_image(model, image_path):129 130    """131    model = get_model('CLIP:ViT-L/14')132    state_dict = torch.load(ckpt, map_location='cpu')133    model.fc.load_state_dict(state_dict)134    print ("Model loaded..")135    model.eval()136    model.cuda()137    """138    img = Image.open(image_path).convert("RGB")139    """140    if jpeg_quality is not None:141        img = png2jpg(img, jpeg_quality)142    """143    transform = transforms.Compose([144            transforms.CenterCrop(224),145            transforms.ToTensor(),146            transforms.Normalize( mean=MEAN['clip'], std=STD['clip'] ),147        ])148    img = transform(img)149    img = img.to('cuda:0')150 151    detection_output = model(img)152    output = torch.sigmoid(detection_output)153 154    return output155 156 157 158 159# = = = = = = = = = = = = = = = = = = = = = = = = = = = = = = = = = = = = # 160"""161def recursively_read(rootdir, must_contain, exts=["png", "jpg", "JPEG", "jpeg", "bmp"]):162    out = [] 163    for r, d, f in os.walk(rootdir):164        for file in f:165            if (file.split('.')[1] in exts)  and  (must_contain in os.path.join(r, file)):166                out.append(os.path.join(r, file))167    return out168 169def get_list(path, must_contain=''):170    if ".pickle" in path:171        with open(path, 'rb') as f:172            image_list = pickle.load(f)173        image_list = [ item for item in image_list if must_contain in item   ]174    else:175        image_list = recursively_read(path, must_contain)176    return image_list177 178class RealFakeDataset(Dataset):179    def __init__(self,  real_path, 180                        fake_path, 181                        data_mode, 182                        max_sample,183                        arch,184                        jpeg_quality=None,185                        gaussian_sigma=None):186 187        assert data_mode in ["wang2020", "ours"]188        self.jpeg_quality = jpeg_quality189        self.gaussian_sigma = gaussian_sigma190        191        # = = = = = = data path = = = = = = = = = # 192        if type(real_path) == str and type(fake_path) == str:193            real_list, fake_list = self.read_path(real_path, fake_path, data_mode, max_sample)194        else:195            real_list = []196            fake_list = []197            for real_p, fake_p in zip(real_path, fake_path):198                real_l, fake_l = self.read_path(real_p, fake_p, data_mode, max_sample)199                real_list += real_l200                fake_list += fake_l201 202        self.total_list = real_list + fake_list203 204 205        # = = = = = =  label = = = = = = = = = # 206 207        self.labels_dict = {}208        for i in real_list:209            self.labels_dict[i] = 0210        for i in fake_list:211            self.labels_dict[i] = 1212 213        stat_from = "imagenet" if arch.lower().startswith("imagenet") else "clip"214        self.transform = transforms.Compose([215            transforms.CenterCrop(224),216            transforms.ToTensor(),217            transforms.Normalize( mean=MEAN[stat_from], std=STD[stat_from] ),218        ])219 220 221    def read_path(self, real_path, fake_path, data_mode, max_sample):222 223        if data_mode == 'wang2020':224            real_list = get_list(real_path, must_contain='0_real')225            fake_list = get_list(fake_path, must_contain='1_fake')226        else:227            real_list = get_list(real_path)228            fake_list = get_list(fake_path)229 230 231        if max_sample is not None:232            if (max_sample > len(real_list)) or (max_sample > len(fake_list)):233                max_sample = 100234                print("not enough images, max_sample falling to 100")235            random.shuffle(real_list)236            random.shuffle(fake_list)237            real_list = real_list[0:max_sample]238            fake_list = fake_list[0:max_sample]239 240        assert len(real_list) == len(fake_list)  241 242        return real_list, fake_list243 244 245 246    def __len__(self):247        return len(self.total_list)248 249    def __getitem__(self, idx):250        251        img_path = self.total_list[idx]252 253        label = self.labels_dict[img_path]254        img = Image.open(img_path).convert("RGB")255 256        if self.gaussian_sigma is not None:257            img = gaussian_blur(img, self.gaussian_sigma) 258        if self.jpeg_quality is not None:259            img = png2jpg(img, self.jpeg_quality)260 261        img = self.transform(img)262        return img, label263"""264 265 266 267 268if __name__ == '__main__':269 270 271    parser = argparse.ArgumentParser(formatter_class=argparse.ArgumentDefaultsHelpFormatter)272    parser.add_argument('--image_path', type=str, default=None, help='path of the image for detection')273    """274    parser.add_argument('--real_path', type=str, default=None, help='dir name or a pickle')275    parser.add_argument('--fake_path', type=str, default=None, help='dir name or a pickle')276    parser.add_argument('--data_mode', type=str, default=None, help='wang2020 or ours')277    parser.add_argument('--max_sample', type=int, default=1000, help='only check this number of images for both fake/real')278    """279    parser.add_argument('--arch', type=str, default='CLIP:ViT-L/14')280    parser.add_argument('--ckpt', type=str, default='./pretrained_weights/fc_weights.pth')281    """282    parser.add_argument('--result_folder', type=str, default='result', help='')283    parser.add_argument('--batch_size', type=int, default=128)284    """285    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")286    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")287 288 289    opt = parser.parse_args()290 291    """292    if os.path.exists(opt.result_folder):293        shutil.rmtree(opt.result_folder)294    os.makedirs(opt.result_folder)295    """296    model = get_model(opt.arch)297    state_dict = torch.load(opt.ckpt, map_location='cpu')298    model.fc.load_state_dict(state_dict)299    # model.load_state_dict(state_dict)300    print ("Model loaded..")301    model.eval()302    model.cuda()303    """304    if (opt.real_path == None) or (opt.fake_path == None) or (opt.data_mode == None):305        dataset_paths = DATASET_PATHS306    else:307        dataset_paths = [ dict(real_path=opt.real_path, fake_path=opt.fake_path, data_mode=opt.data_mode) ]308 309 310 311    for dataset_path in (dataset_paths):312        set_seed()313 314        dataset = RealFakeDataset(  dataset_path['real_path'], 315                                    dataset_path['fake_path'], 316                                    dataset_path['data_mode'], 317                                    opt.max_sample, 318                                    opt.arch,319                                    jpeg_quality=opt.jpeg_quality, 320                                    gaussian_sigma=opt.gaussian_sigma,321                                    )322 323        loader = torch.utils.data.DataLoader(dataset, batch_size=opt.batch_size, shuffle=False, num_workers=4)324        ap, r_acc0, f_acc0, acc0, r_acc1, f_acc1, acc1, best_thres = validate(model, loader, find_thres=True)325 326        with open( os.path.join(opt.result_folder,'ap.txt'), 'a') as f:327            f.write(dataset_path['key']+': ' + str(round(ap*100, 2))+'\n' )328 329        with open( os.path.join(opt.result_folder,'acc0.txt'), 'a') as f:330            f.write(dataset_path['key']+': ' + str(round(r_acc0*100, 2))+'  '+str(round(f_acc0*100, 2))+'  '+str(round(acc0*100, 2))+'\n' )331    """332    output = detect_one_image(model, opt.image_path)333    print(output)