LoliRimuru commited on
Commit
205d391
·
verified ·
1 Parent(s): 55edf0f

Create app.py

Browse files
Files changed (1) hide show
  1. app.py +274 -0
app.py ADDED
@@ -0,0 +1,274 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import os
2
+ import random
3
+ import json
4
+ from pathlib import Path
5
+ from typing import List, Dict
6
+ import numpy as np
7
+ import torch
8
+ import torch.nn as nn
9
+ from torch.utils.data import Dataset, DataLoader
10
+ from torchvision import transforms
11
+ from PIL import Image
12
+ import tqdm
13
+ from sklearn.metrics import (accuracy_score, precision_recall_fscore_support,
14
+ roc_auc_score, confusion_matrix)
15
+ import matplotlib.pyplot as plt
16
+
17
+ # ==================== CONFIGURATION ====================
18
+ class Config:
19
+ CROP_SIZE = 512
20
+ CHECKPOINT_DIR = "./checkpoints"
21
+ RESULTS_DIR = "./results"
22
+
23
+ # ==================== UTILITIES ====================
24
+ def ensure_dir(path: str):
25
+ Path(path).mkdir(parents=True, exist_ok=True)
26
+
27
+ # ==================== MODEL (copy from training) ====================
28
+ class LightweightCompressionNet(nn.Module):
29
+ def __init__(self):
30
+ super().__init__()
31
+ self.conv_blocks = nn.Sequential(
32
+ nn.Conv2d(3, 16, kernel_size=4, stride=1, padding=0), nn.GELU(),
33
+ nn.Conv2d(16, 32, kernel_size=4, stride=1, padding=0), nn.GELU(),
34
+ nn.Conv2d(32, 64, kernel_size=4, stride=2, padding=0), nn.GELU(),
35
+ nn.Conv2d(64, 128, kernel_size=4, stride=2, padding=0), nn.GELU(),
36
+ nn.Conv2d(128, 256, kernel_size=4, stride=4, padding=0), nn.GELU(),
37
+ nn.Conv2d(256, 256, kernel_size=4, stride=4, padding=0), nn.GELU(),
38
+ nn.Conv2d(256, 256, kernel_size=3, stride=2, padding=0), nn.GELU(),
39
+ nn.AdaptiveAvgPool2d(1)
40
+ )
41
+ self.head = nn.Sequential(
42
+ nn.Linear(256, 32), nn.GELU(),
43
+ nn.Linear(32, 1), nn.Sigmoid()
44
+ )
45
+
46
+ def forward(self, x):
47
+ features = self.conv_blocks(x)
48
+ features = features.view(features.size(0), -1)
49
+ return self.head(features).squeeze(1)
50
+
51
+ # ==================== INFERENCE DATASET ====================
52
+ class InferenceDataset(Dataset):
53
+ def __init__(self, image_paths: List[str]):
54
+ self.image_paths = image_paths
55
+ self.transform = transforms.Compose([
56
+ transforms.Resize(Config.CROP_SIZE),
57
+ transforms.CenterCrop(Config.CROP_SIZE),
58
+ transforms.ToTensor(),
59
+ ])
60
+
61
+ def __len__(self):
62
+ return len(self.image_paths)
63
+
64
+ def __getitem__(self, idx):
65
+ path = self.image_paths[idx]
66
+ try:
67
+ image = Image.open(path).convert('RGB')
68
+ image = self.transform(image)
69
+ # Apply same INT8 quantization as training
70
+ image = (image * 255).round().clamp(0, 255) / 255
71
+ return image
72
+ except Exception as e:
73
+ print(f"Warning: Failed to load {path}: {e}")
74
+ return torch.zeros(3, Config.CROP_SIZE, Config.CROP_SIZE)
75
+
76
+ # ==================== INFERENCE FUNCTIONS ====================
77
+ def load_model(checkpoint_path: str, device: torch.device) -> nn.Module:
78
+ """Load trained model from checkpoint"""
79
+ if not os.path.exists(checkpoint_path):
80
+ raise FileNotFoundError(f"Checkpoint not found: {checkpoint_path}")
81
+
82
+ model = LightweightCompressionNet().to(device)
83
+ checkpoint = torch.load(checkpoint_path, map_location=device)
84
+ model.load_state_dict(checkpoint['model_state_dict'])
85
+ model.eval()
86
+ print(f"Loaded model from epoch {checkpoint.get('epoch', 'unknown')} "
87
+ f"(Val loss: {checkpoint.get('val_loss', 'N/A'):.4f})")
88
+ return model
89
+
90
+ def run_inference(model: nn.Module, image_paths: List[str], device: torch.device,
91
+ batch_size: int = 4) -> np.ndarray:
92
+ """Run inference on a list of images"""
93
+ dataset = InferenceDataset(image_paths)
94
+ loader = DataLoader(dataset, batch_size=batch_size, shuffle=False,
95
+ num_workers=2, pin_memory=True)
96
+
97
+ all_predictions = []
98
+ with torch.no_grad():
99
+ pbar = tqdm.tqdm(loader, desc="Processing images", unit="batch")
100
+ for batch in pbar:
101
+ images = batch.to(device, non_blocking=True)
102
+ predictions = model(images)
103
+ all_predictions.extend(predictions.cpu().numpy())
104
+
105
+ return np.array(all_predictions)
106
+
107
+ def calculate_metrics(y_true: np.ndarray, y_scores: np.ndarray,
108
+ threshold: float = 0.5) -> Dict:
109
+ """Calculate classification metrics"""
110
+ y_pred = (y_scores >= threshold).astype(int)
111
+
112
+ accuracy = accuracy_score(y_true, y_pred)
113
+ precision, recall, f1, _ = precision_recall_fscore_support(
114
+ y_true, y_pred, average='binary', zero_division=0
115
+ )
116
+
117
+ try:
118
+ roc_auc = roc_auc_score(y_true, y_scores)
119
+ except ValueError:
120
+ roc_auc = None
121
+
122
+ conf_matrix = confusion_matrix(y_true, y_pred)
123
+ tn, fp, fn, tp = conf_matrix.ravel()
124
+
125
+ return {
126
+ 'accuracy': accuracy,
127
+ 'precision': precision,
128
+ 'recall': recall,
129
+ 'f1_score': f1,
130
+ 'roc_auc': roc_auc,
131
+ 'true_negatives': int(tn),
132
+ 'false_positives': int(fp),
133
+ 'false_negatives': int(fn),
134
+ 'true_positives': int(tp),
135
+ 'threshold': threshold,
136
+ 'confusion_matrix': conf_matrix.tolist()
137
+ }
138
+
139
+ def plot_score_distribution(ai_scores: np.ndarray, non_ai_scores: np.ndarray,
140
+ save_path: str = None):
141
+ """Plot histogram of prediction scores"""
142
+ plt.figure(figsize=(10, 6))
143
+ plt.hist(non_ai_scores, bins=20, alpha=0.7, label='Non-AI', color='blue', density=True)
144
+ plt.hist(ai_scores, bins=20, alpha=0.7, label='AI', color='red', density=True)
145
+ plt.axvline(x=0.5, color='black', linestyle='--', label='Threshold (0.5)')
146
+ plt.xlabel('AI Detection Score')
147
+ plt.ylabel('Density')
148
+ plt.title('Distribution of AI Detection Scores')
149
+ plt.legend()
150
+ plt.grid(True, alpha=0.3)
151
+
152
+ if save_path:
153
+ plt.savefig(save_path, dpi=150, bbox_inches='tight')
154
+ print(f"Saved plot: {save_path}")
155
+ plt.show()
156
+
157
+ # ==================== MAIN INFERENCE ====================
158
+ def main():
159
+ # =================================================
160
+ # HARDCODED CONFIGURATION - MODIFY THESE VALUES
161
+ # =================================================
162
+ NON_AI_FOLDER = "/home/pc/Dokumenty/test_non_ai" # CHANGE THIS to your non-AI folder
163
+ AI_FOLDER = "/home/pc/Dokumenty/test_ai" # CHANGE THIS to your AI folder
164
+ SAMPLE_SIZE = 10 # Number of images from each folder
165
+ CHECKPOINT_PATH = os.path.join(Config.CHECKPOINT_DIR, "best_model.pt")
166
+ BATCH_SIZE = 4
167
+ SEED = 42
168
+ OUTPUT_JSON = os.path.join(Config.RESULTS_DIR, "inference_results.json")
169
+ # =================================================
170
+
171
+ # Setup
172
+ device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
173
+ print(f"Using device: {device}")
174
+
175
+ ensure_dir(Config.RESULTS_DIR)
176
+
177
+ # Load model
178
+ model = load_model(CHECKPOINT_PATH, device)
179
+
180
+ # Get image paths (support PNG, JPG, JPEG)
181
+ non_ai_paths = []
182
+ for ext in ['*.png', '*.jpg', '*.jpeg']:
183
+ non_ai_paths.extend([str(p) for p in Path(NON_AI_FOLDER).rglob(ext)])
184
+
185
+ ai_paths = []
186
+ for ext in ['*.png', '*.jpg', '*.jpeg']:
187
+ ai_paths.extend([str(p) for p in Path(AI_FOLDER).rglob(ext)])
188
+
189
+ if not non_ai_paths:
190
+ raise ValueError(f"No images found in {NON_AI_FOLDER}")
191
+ if not ai_paths:
192
+ raise ValueError(f"No images found in {AI_FOLDER}")
193
+
194
+ # Random sampling
195
+ random.seed(SEED)
196
+ non_ai_sample = random.sample(non_ai_paths, min(SAMPLE_SIZE, len(non_ai_paths)))
197
+ ai_sample = random.sample(ai_paths, min(SAMPLE_SIZE, len(ai_paths)))
198
+
199
+ print(f"\nSampling {len(non_ai_sample)} non-AI images from {NON_AI_FOLDER}")
200
+ print(f"Sampling {len(ai_sample)} AI images from {AI_FOLDER}")
201
+
202
+ # Run inference
203
+ print("\n" + "="*50)
204
+ non_ai_scores = run_inference(model, non_ai_sample, device, BATCH_SIZE)
205
+ ai_scores = run_inference(model, ai_sample, device, BATCH_SIZE)
206
+
207
+ # Create labels (0=non-AI, 1=AI)
208
+ y_true = np.array([0] * len(non_ai_scores) + [1] * len(ai_scores))
209
+ y_scores = np.concatenate([non_ai_scores, ai_scores])
210
+
211
+ # Calculate metrics
212
+ metrics = calculate_metrics(y_true, y_scores)
213
+
214
+ # Print results
215
+ print("\n" + "="*50)
216
+ print("INFERENCE RESULTS")
217
+ print("="*50)
218
+ print(f"\nOverall Metrics:")
219
+ print(f" Accuracy: {metrics['accuracy']:.4f}")
220
+ print(f" Precision: {metrics['precision']:.4f}")
221
+ print(f" Recall: {metrics['recall']:.4f}")
222
+ print(f" F1-Score: {metrics['f1_score']:.4f}")
223
+ if metrics['roc_auc'] is not None:
224
+ print(f" ROC-AUC: {metrics['roc_auc']:.4f}")
225
+
226
+ print(f"\nConfusion Matrix:")
227
+ print(f" True Negatives (Non-AI correct): {metrics['true_negatives']}")
228
+ print(f" False Positives (Non-AI wrong): {metrics['false_positives']}")
229
+ print(f" False Negatives (AI wrong): {metrics['false_negatives']}")
230
+ print(f" True Positives (AI correct): {metrics['true_positives']}")
231
+
232
+ # Per-class accuracy
233
+ if len(non_ai_scores) > 0:
234
+ non_ai_acc = (non_ai_scores < 0.5).mean()
235
+ print(f"\nNon-AI Detection Accuracy: {non_ai_acc:.4f} ({non_ai_acc * 100:.1f}%)")
236
+
237
+ if len(ai_scores) > 0:
238
+ ai_acc = (ai_scores >= 0.5).mean()
239
+ print(f"AI Detection Accuracy: {ai_acc:.4f} ({ai_acc * 100:.1f}%)")
240
+
241
+ # Save detailed results
242
+ results = {
243
+ 'config': {
244
+ 'non_ai_folder': NON_AI_FOLDER,
245
+ 'ai_folder': AI_FOLDER,
246
+ 'sample_size': SAMPLE_SIZE,
247
+ 'seed': SEED,
248
+ 'checkpoint': CHECKPOINT_PATH,
249
+ 'threshold': metrics['threshold']
250
+ },
251
+ 'image_paths': {
252
+ 'non_ai': non_ai_sample,
253
+ 'ai': ai_sample
254
+ },
255
+ 'predictions': {
256
+ 'non_ai_scores': non_ai_scores.tolist(),
257
+ 'ai_scores': ai_scores.tolist()
258
+ },
259
+ 'metrics': metrics
260
+ }
261
+
262
+ with open(OUTPUT_JSON, 'w') as f:
263
+ json.dump(results, f, indent=2)
264
+ print(f"\nDetailed results saved to: {OUTPUT_JSON}")
265
+
266
+ # Plot distribution
267
+ plot_path = os.path.join(Config.RESULTS_DIR, "score_distribution.png")
268
+ plot_score_distribution(ai_scores, non_ai_scores, plot_path)
269
+
270
+ print("\nDone!")
271
+
272
+
273
+ if __name__ == "__main__":
274
+ main()