LoliRimuru's picture
Create app.py
205d391 verified
Raw History Blame
10.3 kB
import os
import random
import json
from pathlib import Path
from typing import List, Dict
import numpy as np
import torch
import torch.nn as nn
from torch.utils.data import Dataset, DataLoader
from torchvision import transforms
from PIL import Image
import tqdm
from sklearn.metrics import (accuracy_score, precision_recall_fscore_support,
roc_auc_score, confusion_matrix)
import matplotlib.pyplot as plt
# ==================== CONFIGURATION ====================
class Config:
CROP_SIZE = 512
CHECKPOINT_DIR = "./checkpoints"
RESULTS_DIR = "./results"
# ==================== UTILITIES ====================
def ensure_dir(path: str):
Path(path).mkdir(parents=True, exist_ok=True)
# ==================== MODEL (copy from training) ====================
class LightweightCompressionNet(nn.Module):
def __init__(self):
super().__init__()
self.conv_blocks = nn.Sequential(
nn.Conv2d(3, 16, kernel_size=4, stride=1, padding=0), nn.GELU(),
nn.Conv2d(16, 32, kernel_size=4, stride=1, padding=0), nn.GELU(),
nn.Conv2d(32, 64, kernel_size=4, stride=2, padding=0), nn.GELU(),
nn.Conv2d(64, 128, kernel_size=4, stride=2, padding=0), nn.GELU(),
nn.Conv2d(128, 256, kernel_size=4, stride=4, padding=0), nn.GELU(),
nn.Conv2d(256, 256, kernel_size=4, stride=4, padding=0), nn.GELU(),
nn.Conv2d(256, 256, kernel_size=3, stride=2, padding=0), nn.GELU(),
nn.AdaptiveAvgPool2d(1)
)
self.head = nn.Sequential(
nn.Linear(256, 32), nn.GELU(),
nn.Linear(32, 1), nn.Sigmoid()
)
def forward(self, x):
features = self.conv_blocks(x)
features = features.view(features.size(0), -1)
return self.head(features).squeeze(1)
# ==================== INFERENCE DATASET ====================
class InferenceDataset(Dataset):
def __init__(self, image_paths: List[str]):
self.image_paths = image_paths
self.transform = transforms.Compose([
transforms.Resize(Config.CROP_SIZE),
transforms.CenterCrop(Config.CROP_SIZE),
transforms.ToTensor(),
])
def __len__(self):
return len(self.image_paths)
def __getitem__(self, idx):
path = self.image_paths[idx]
try:
image = Image.open(path).convert('RGB')
image = self.transform(image)
# Apply same INT8 quantization as training
image = (image * 255).round().clamp(0, 255) / 255
return image
except Exception as e:
print(f"Warning: Failed to load {path}: {e}")
return torch.zeros(3, Config.CROP_SIZE, Config.CROP_SIZE)
# ==================== INFERENCE FUNCTIONS ====================
def load_model(checkpoint_path: str, device: torch.device) -> nn.Module:
"""Load trained model from checkpoint"""
if not os.path.exists(checkpoint_path):
raise FileNotFoundError(f"Checkpoint not found: {checkpoint_path}")
model = LightweightCompressionNet().to(device)
checkpoint = torch.load(checkpoint_path, map_location=device)
model.load_state_dict(checkpoint['model_state_dict'])
model.eval()
print(f"Loaded model from epoch {checkpoint.get('epoch', 'unknown')} "
f"(Val loss: {checkpoint.get('val_loss', 'N/A'):.4f})")
return model
def run_inference(model: nn.Module, image_paths: List[str], device: torch.device,
batch_size: int = 4) -> np.ndarray:
"""Run inference on a list of images"""
dataset = InferenceDataset(image_paths)
loader = DataLoader(dataset, batch_size=batch_size, shuffle=False,
num_workers=2, pin_memory=True)
all_predictions = []
with torch.no_grad():
pbar = tqdm.tqdm(loader, desc="Processing images", unit="batch")
for batch in pbar:
images = batch.to(device, non_blocking=True)
predictions = model(images)
all_predictions.extend(predictions.cpu().numpy())
return np.array(all_predictions)
def calculate_metrics(y_true: np.ndarray, y_scores: np.ndarray,
threshold: float = 0.5) -> Dict:
"""Calculate classification metrics"""
y_pred = (y_scores >= threshold).astype(int)
accuracy = accuracy_score(y_true, y_pred)
precision, recall, f1, _ = precision_recall_fscore_support(
y_true, y_pred, average='binary', zero_division=0
)
try:
roc_auc = roc_auc_score(y_true, y_scores)
except ValueError:
roc_auc = None
conf_matrix = confusion_matrix(y_true, y_pred)
tn, fp, fn, tp = conf_matrix.ravel()
return {
'accuracy': accuracy,
'precision': precision,
'recall': recall,
'f1_score': f1,
'roc_auc': roc_auc,
'true_negatives': int(tn),
'false_positives': int(fp),
'false_negatives': int(fn),
'true_positives': int(tp),
'threshold': threshold,
'confusion_matrix': conf_matrix.tolist()
}
def plot_score_distribution(ai_scores: np.ndarray, non_ai_scores: np.ndarray,
save_path: str = None):
"""Plot histogram of prediction scores"""
plt.figure(figsize=(10, 6))
plt.hist(non_ai_scores, bins=20, alpha=0.7, label='Non-AI', color='blue', density=True)
plt.hist(ai_scores, bins=20, alpha=0.7, label='AI', color='red', density=True)
plt.axvline(x=0.5, color='black', linestyle='--', label='Threshold (0.5)')
plt.xlabel('AI Detection Score')
plt.ylabel('Density')
plt.title('Distribution of AI Detection Scores')
plt.legend()
plt.grid(True, alpha=0.3)
if save_path:
plt.savefig(save_path, dpi=150, bbox_inches='tight')
print(f"Saved plot: {save_path}")
plt.show()
# ==================== MAIN INFERENCE ====================
def main():
# =================================================
# HARDCODED CONFIGURATION - MODIFY THESE VALUES
# =================================================
NON_AI_FOLDER = "/home/pc/Dokumenty/test_non_ai" # CHANGE THIS to your non-AI folder
AI_FOLDER = "/home/pc/Dokumenty/test_ai" # CHANGE THIS to your AI folder
SAMPLE_SIZE = 10 # Number of images from each folder
CHECKPOINT_PATH = os.path.join(Config.CHECKPOINT_DIR, "best_model.pt")
BATCH_SIZE = 4
SEED = 42
OUTPUT_JSON = os.path.join(Config.RESULTS_DIR, "inference_results.json")
# =================================================
# Setup
device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
print(f"Using device: {device}")
ensure_dir(Config.RESULTS_DIR)
# Load model
model = load_model(CHECKPOINT_PATH, device)
# Get image paths (support PNG, JPG, JPEG)
non_ai_paths = []
for ext in ['*.png', '*.jpg', '*.jpeg']:
non_ai_paths.extend([str(p) for p in Path(NON_AI_FOLDER).rglob(ext)])
ai_paths = []
for ext in ['*.png', '*.jpg', '*.jpeg']:
ai_paths.extend([str(p) for p in Path(AI_FOLDER).rglob(ext)])
if not non_ai_paths:
raise ValueError(f"No images found in {NON_AI_FOLDER}")
if not ai_paths:
raise ValueError(f"No images found in {AI_FOLDER}")
# Random sampling
random.seed(SEED)
non_ai_sample = random.sample(non_ai_paths, min(SAMPLE_SIZE, len(non_ai_paths)))
ai_sample = random.sample(ai_paths, min(SAMPLE_SIZE, len(ai_paths)))
print(f"\nSampling {len(non_ai_sample)} non-AI images from {NON_AI_FOLDER}")
print(f"Sampling {len(ai_sample)} AI images from {AI_FOLDER}")
# Run inference
print("\n" + "="*50)
non_ai_scores = run_inference(model, non_ai_sample, device, BATCH_SIZE)
ai_scores = run_inference(model, ai_sample, device, BATCH_SIZE)
# Create labels (0=non-AI, 1=AI)
y_true = np.array([0] * len(non_ai_scores) + [1] * len(ai_scores))
y_scores = np.concatenate([non_ai_scores, ai_scores])
# Calculate metrics
metrics = calculate_metrics(y_true, y_scores)
# Print results
print("\n" + "="*50)
print("INFERENCE RESULTS")
print("="*50)
print(f"\nOverall Metrics:")
print(f" Accuracy: {metrics['accuracy']:.4f}")
print(f" Precision: {metrics['precision']:.4f}")
print(f" Recall: {metrics['recall']:.4f}")
print(f" F1-Score: {metrics['f1_score']:.4f}")
if metrics['roc_auc'] is not None:
print(f" ROC-AUC: {metrics['roc_auc']:.4f}")
print(f"\nConfusion Matrix:")
print(f" True Negatives (Non-AI correct): {metrics['true_negatives']}")
print(f" False Positives (Non-AI wrong): {metrics['false_positives']}")
print(f" False Negatives (AI wrong): {metrics['false_negatives']}")
print(f" True Positives (AI correct): {metrics['true_positives']}")
# Per-class accuracy
if len(non_ai_scores) > 0:
non_ai_acc = (non_ai_scores < 0.5).mean()
print(f"\nNon-AI Detection Accuracy: {non_ai_acc:.4f} ({non_ai_acc * 100:.1f}%)")
if len(ai_scores) > 0:
ai_acc = (ai_scores >= 0.5).mean()
print(f"AI Detection Accuracy: {ai_acc:.4f} ({ai_acc * 100:.1f}%)")
# Save detailed results
results = {
'config': {
'non_ai_folder': NON_AI_FOLDER,
'ai_folder': AI_FOLDER,
'sample_size': SAMPLE_SIZE,
'seed': SEED,
'checkpoint': CHECKPOINT_PATH,
'threshold': metrics['threshold']
},
'image_paths': {
'non_ai': non_ai_sample,
'ai': ai_sample
},
'predictions': {
'non_ai_scores': non_ai_scores.tolist(),
'ai_scores': ai_scores.tolist()
},
'metrics': metrics
}
with open(OUTPUT_JSON, 'w') as f:
json.dump(results, f, indent=2)
print(f"\nDetailed results saved to: {OUTPUT_JSON}")
# Plot distribution
plot_path = os.path.join(Config.RESULTS_DIR, "score_distribution.png")
plot_score_distribution(ai_scores, non_ai_scores, plot_path)
print("\nDone!")
if __name__ == "__main__":
main()