Download app.py from LoliRimuru/AAL-Plus_Image_Quality_Assessment: direct link, hf CLI and curl.
- Browser
- Download file 10.3 kB
-
https://huggingface.co/spaces/LoliRimuru/AAL-Plus_Image_Quality_Assessment/resolve/205d3917fcf1e1e104b35e1c7fb8454ae151083e/app.py
- Command line
-
hf download hf://spaces/LoliRimuru/AAL-Plus_Image_Quality_Assessment@205d3917fcf1e1e104b35e1c7fb8454ae151083e/app.py
-
curl -L -o app.py https://huggingface.co/spaces/LoliRimuru/AAL-Plus_Image_Quality_Assessment/resolve/205d3917fcf1e1e104b35e1c7fb8454ae151083e/app.py
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() |