Download clean/image/cospy/dataSets/cospy_real.py from deepsafe/model-code: direct link, hf CLI and curl.
- Browser
- Download file 3.25 kB
-
https://huggingface.co/deepsafe/model-code/resolve/main/clean/image/cospy/dataSets/cospy_real.py
- Command line
-
hf download hf://deepsafe/model-code/clean/image/cospy/dataSets/cospy_real.py
-
curl -L -o cospy_real.py https://huggingface.co/deepsafe/model-code/resolve/main/clean/image/cospy/dataSets/cospy_real.py
3.25 kB
| import os | |
| import json | |
| import torch | |
| import numpy as np | |
| from PIL import Image | |
| import datasets as ds | |
| class MSCOCO2017(torch.utils.data.Dataset): | |
| def __init__(self, split='val', transform=None): | |
| # Split [train: 118287, val: 5000] | |
| image_dir = f"data/train/sd-v1_4/mscoco2017/{split}2017" | |
| self.image_list = [os.path.join(image_dir, f) for f in os.listdir(image_dir) if f.endswith(('.jpg'))] | |
| self.image_list.sort() | |
| caption_filepath = f"data/train/sd-v1_4/stable-diffusion-v1-4/{split}2017_gen.json" | |
| with open(caption_filepath, 'r') as f: | |
| self.captions = json.load(f) | |
| # Preprocess the images | |
| self.transform = transform | |
| def __len__(self): | |
| return len(self.image_list) | |
| def __getitem__(self, idx): | |
| image_filepath = self.image_list[idx] | |
| image = Image.open(image_filepath).convert('RGB') | |
| if self.transform: | |
| image = self.transform(image) | |
| # Get the captions | |
| image_id = os.path.basename(image_filepath) | |
| caption = self.captions[image_id] | |
| return image, caption | |
| class Flickr30k(torch.utils.data.Dataset): | |
| def __init__(self, split='test', transform=None): | |
| # Split [test: 31014] | |
| self.dataset = ds.load_dataset("nlphuji/flickr30k")[split] | |
| # Preprocess the images | |
| self.transform = transform | |
| def __len__(self): | |
| return len(self.dataset) | |
| def __getitem__(self, idx): | |
| example = self.dataset[idx] | |
| # PIL RGB image | |
| image = example['image'] | |
| if self.transform: | |
| image = self.transform(image) | |
| # A list of valid captions | |
| caption_list = example['caption'] | |
| # Randomly select a caption | |
| caption = np.random.choice(caption_list) | |
| return image, caption | |
| class OtherReal(torch.utils.data.Dataset): | |
| def __init__(self, dataset_name, transform=None): | |
| # Load the test samples from other real datasets (CC3M, SBU, TextCaps) | |
| # Note: there are only 2,000 examples for each dataset | |
| root_dir = "data/test/Co-Spy-Bench/real_image_examples" | |
| if dataset_name == "cc3m": | |
| self.image_dir = os.path.join(root_dir, "CC3M") | |
| elif dataset_name == "sbu": | |
| self.image_dir = os.path.join(root_dir, "SBU") | |
| elif dataset_name == "textcaps": | |
| self.image_dir = os.path.join(root_dir, "TextCaps") | |
| else: | |
| raise ValueError(f"Dataset {dataset_name} not supported in OtherReal class.") | |
| self.image_list = [f for f in os.listdir(self.image_dir) if f.endswith(('.jpg'))] | |
| self.image_list.sort() | |
| # Preprocess the images | |
| self.transform = transform | |
| def __len__(self): | |
| return len(self.image_list) | |
| def __getitem__(self, idx): | |
| img_name = self.image_list[idx] | |
| img_path = os.path.join(self.image_dir, img_name) | |
| image = Image.open(img_path).convert('RGB') | |
| if self.transform: | |
| image = self.transform(image) | |
| caption_filepath = img_path.replace('.jpg', '.json') | |
| with open(caption_filepath, 'r') as f: | |
| caption_data = json.load(f) | |
| caption = caption_data['caption'] | |
| return image, caption | |