File size: 7,823 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 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 | import os
import json
import torch
import importlib
import numpy as np
from tqdm import tqdm
from PIL import Image
from loguru import logger
from dataSets import *
from utils import seed_torch, evaluate
import warnings
warnings.filterwarnings("ignore")
# Load pre-trained models for evaluation
class Detector:
def __init__(self,
device: str,
mode: str = "fusion",
train_dataset: str = "sd-v1_4",
pretrain: bool = False,
ckpt: str = "ckpt",
batch_size: int = 32):
# Device
self.device = device
self.mode = mode
self.train_dataset = train_dataset
self.pretrain = pretrain
self.ckpt = ckpt
self.batch_size = batch_size
# Dynamically import the detector module from either "progan" or "sd-v1_4"
detector_module = importlib.import_module(f"detectors.{train_dataset}")
# Get the detector and load weights based on mode
if pretrain:
# Only provide pre-trained weights for fusion mode
# Hardcode to fusion mode
self.mode = "fusion"
# Load the fusion detector with pre-trained weights
semantic_weights_path = f"pretrained/{self.train_dataset}/semantic_weights.pth"
artifact_weights_path = f"pretrained/{self.train_dataset}/artifact_weights.pth"
fusion_weights_path = f"pretrained/{self.train_dataset}/fusion_weights.pth"
if not os.path.exists(semantic_weights_path) or not os.path.exists(artifact_weights_path) or not os.path.exists(fusion_weights_path):
raise ValueError("The pre-trained weights are not complete for evaluation")
CoSpyFusionDetector = getattr(detector_module, "CoSpyFusionDetector")
self.model = CoSpyFusionDetector(
semantic_weights_path=semantic_weights_path,
artifact_weights_path=artifact_weights_path)
self.model.load_weights(fusion_weights_path)
else:
if self.mode == "fusion":
semantic_weights_path = os.path.join(self.ckpt, self.train_dataset, "semantic", "best_model.pth")
artifact_weights_path = os.path.join(self.ckpt, self.train_dataset, "artifact", "best_model.pth")
fusion_weights_path = os.path.join(self.ckpt, self.train_dataset, "fusion", "best_model.pth")
if not os.path.exists(semantic_weights_path) or not os.path.exists(artifact_weights_path) or not os.path.exists(fusion_weights_path):
raise ValueError("Semantic, Artifact or Fusion weights path does not exist for fusion mode")
CoSpyFusionDetector = getattr(detector_module, "CoSpyFusionDetector")
self.model = CoSpyFusionDetector(
semantic_weights_path=semantic_weights_path,
artifact_weights_path=artifact_weights_path)
self.model.load_weights(fusion_weights_path)
elif self.mode == "end2end":
End2EndDetector = getattr(detector_module, "End2EndDetector")
self.model = End2EndDetector()
end2end_weights_path = os.path.join(self.ckpt, self.train_dataset, "end2end", "best_model.pth")
if not os.path.exists(end2end_weights_path):
raise ValueError("End2End weights path does not exist for end2end mode")
self.model.load_weights(end2end_weights_path)
else:
raise ValueError(f"Unknown mode: {self.mode}")
# Put the model on the device and set to eval
self.model.to(self.device)
self.model.eval()
def evaluate_benchmark(self):
# Select the appropriate test dataset and evaluation lists
if self.train_dataset == "progan":
benchmark_name = "AIGCDetectionBenchMark"
TestDataset = AIGCDetectTestDataset
eval_dataset_list = AIGCDetectionBenchMark_DATASET_LIST
eval_model_list = AIGCDetectionBenchMark_MODEL_LIST
elif self.train_dataset == "sd-v1_4":
benchmark_name = "Co-Spy-Bench"
TestDataset = CoSpyBenchTestDataset
eval_dataset_list = CoSpyBench_DATASET_LIST
eval_model_list = CoSpyBench_MODEL_LIST
else:
raise ValueError(f"Unknown train dataset: {self.train_dataset}")
# Set the saving directory
if self.pretrain:
save_dir = os.path.join(self.ckpt, self.train_dataset, self.mode, f"pretrain_{benchmark_name}")
else:
save_dir = os.path.join(self.ckpt, self.train_dataset, self.mode, f"eval_{benchmark_name}")
if not os.path.exists(save_dir):
os.makedirs(save_dir)
# Setup logger
log_path = f"{save_dir}/evaluation.log"
if os.path.exists(log_path):
os.remove(log_path)
logger_id = logger.add(
log_path,
format="{time:MM-DD at HH:mm:ss} | {level} | {module}:{line} | {message}",
level="DEBUG",
)
# Save raw model prediction
save_output_path = os.path.join(save_dir, "output.json")
# Save summarized evaluation result
save_result_path = os.path.join(save_dir, "result.json")
# Begin the evaluation
result_all = {}
output_all = {}
for dataset_name in eval_dataset_list:
result_all[dataset_name] = {}
output_all[dataset_name] = {}
for model_name in eval_model_list:
test_dataset = TestDataset(dataset=dataset_name, model=model_name, transform=self.model.test_transform)
test_loader = torch.utils.data.DataLoader(test_dataset,
batch_size=self.batch_size,
shuffle=False,
num_workers=4,
pin_memory=True)
# Evaluate the model
y_pred, y_true = [], []
for (images, labels) in tqdm(test_loader, desc=f"Evaluating {benchmark_name} - {dataset_name} - {model_name}"):
y_pred.extend(self.model.predict(images))
y_true.extend(labels.tolist())
ap, accuracy = evaluate(y_pred, y_true)
logger.info(f"Evaluate on {benchmark_name} - {dataset_name} - {model_name} | Size {len(y_true)} | AP {ap*100:.2f}% | Accuracy {accuracy*100:.2f}%")
result_all[dataset_name][model_name] = {"size": len(y_true), "AP": ap, "Accuracy": accuracy}
output_all[dataset_name][model_name] = {"y_pred": y_pred, "y_true": y_true}
# Save the results
with open(save_result_path, "w") as f:
json.dump(result_all, f, indent=4)
with open(save_output_path, "w") as f:
json.dump(output_all, f, indent=4)
def scan(self):
# Load the image
image_filepath = input("Please enter the image filepath for scanning: ")
if not os.path.exists(image_filepath):
print(f"Image file not found: {image_filepath}")
image_filepath = input("Please enter the image filepath for scanning: ")
image = Image.open(image_filepath).convert("RGB")
image = self.model.test_transform(image)
image = image.unsqueeze(0)
image = image.to(self.device)
# Make the prediction
prediction = self.model.predict(image)[0]
if prediction > 0.5:
print(f"Co-Spy Prediction: {prediction:.3f} - AI-Generated")
else:
print(f"Co-Spy Prediction: {prediction:.3f} - Real")
|