File size: 4,771 Bytes
9e14838 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 | import os
import numpy as np
from PIL import Image
from tqdm import tqdm
from torch.utils.data import Dataset
from utils import get_list, png_to_jpeg
from .cospy_real import MSCOCO2017, Flickr30k, OtherReal
class TrainDataset(Dataset):
def __init__(self, train_dataset, split="train", add_jpeg=False, transform=None):
assert split in ["train", "val"]
# Root directory of the training datasets
root_dir = f"data/train/{train_dataset}"
# Load the dataset for training
if train_dataset == "progan":
real_list = get_list(os.path.join(root_dir, split), must_contain='0_real')
fake_list = get_list(os.path.join(root_dir, split), must_contain='1_fake')
elif train_dataset == "sd-v1_4":
real_list = get_list(os.path.join(root_dir, "mscoco2017", f"{split}2017"))
fake_list = get_list(os.path.join(root_dir, "stable-diffusion-v1-4", f"{split}2017"))
# Setting the labels for the dataset
self.labels_dict = {}
for i in real_list:
self.labels_dict[i] = 0
for i in fake_list:
self.labels_dict[i] = 1
# Construct the entire dataset
self.total_list = real_list + fake_list
np.random.shuffle(self.total_list)
# JPEG compression
self.add_jpeg = add_jpeg
# Transformations
self.transform = transform
def __len__(self):
return len(self.total_list)
def __getitem__(self, idx):
img_path = self.total_list[idx]
label = self.labels_dict[img_path]
image = Image.open(img_path).convert("RGB")
# Add JPEG compression
if self.add_jpeg:
image = png_to_jpeg(image, quality=95)
# Apply the transformation
if self.transform is not None:
image = self.transform(image)
return image, label
class CoSpyBenchTestDataset(Dataset):
def __init__(self, dataset, model, num_real=2000, add_jpeg=True, transform=None):
# Root path of the Co-Spy-Bench dataset
root_path = "data/test/Co-Spy-Bench/synthetic"
# Load fake images
fake_dir = os.path.join(root_path, dataset, model)
fake_list = [i for i in os.listdir(fake_dir) if i.endswith(".png")]
fake_list.sort()
self.fake = [os.path.join(fake_dir, i) for i in fake_list]
# Take the real images from the dataset
if dataset == "mscoco":
self.real = MSCOCO2017()
elif dataset == "flickr":
self.real = Flickr30k()
else:
self.real = OtherReal(dataset)
# Ensure the number of real and fake images are the same
self.num_real = min(num_real, len(self.real), len(self.fake))
self.image_idx = list(range(self.num_real * 2))
# First half is real, second half is fake
self.labels = [0] * self.num_real + [1] * self.num_real
# JPEG compression
self.add_jpeg = add_jpeg
# Transformations
self.transform = transform
def __len__(self):
return len(self.image_idx)
def __getitem__(self, idx):
if idx < self.num_real:
image, _ = self.real[idx]
else:
image = Image.open(self.fake[idx - self.num_real]).convert("RGB")
# JPEG compression
if self.add_jpeg:
image = png_to_jpeg(image, quality=95)
# Transformations
if self.transform is not None:
image = self.transform(image)
label = self.labels[idx]
return image, label
class AIGCDetectTestDataset(Dataset):
def __init__(self, dataset, model, transform=None):
# Root path of the AIGCDetectionBenchMark dataset
root_path = "data/test/AIGCDetectionBenchMark"
# Load images
image_dir = os.path.join(root_path, dataset, model)
real_list = get_list(image_dir, must_contain='0_real')
fake_list = get_list(image_dir, must_contain='1_fake')
# Setting the labels for the dataset
self.labels_dict = {}
for i in real_list:
self.labels_dict[i] = 0
for i in fake_list:
self.labels_dict[i] = 1
# Construct the entire dataset
self.total_list = real_list + fake_list
np.random.shuffle(self.total_list)
# Transformations
self.transform = transform
def __len__(self):
return len(self.total_list)
def __getitem__(self, idx):
img_path = self.total_list[idx]
label = self.labels_dict[img_path]
image = Image.open(img_path).convert("RGB")
# Transformations
if self.transform is not None:
image = self.transform(image)
return image, label
|