CoolFace
Apppublic

DocForg/Document_Forgery_Detection

sourceHugging Facemitupdated 6mo agoView on Hugging Face
0likes
evaluate.py159 linesDownload Raw Back to scripts
1"""
2Model Evaluation Script
3
4Evaluate trained model on validation/test sets with comprehensive metrics.
5
6Usage:
7    python scripts/evaluate.py --model outputs/checkpoints/best_doctamper.pth --dataset doctamper
8"""
9
10import argparse
11import sys
12from pathlib import Path
13import json
14import numpy as np
15from tqdm import tqdm
16import torch
17
18# Add src to path
19sys.path.insert(0, str(Path(__file__).parent.parent))
20
21from src.config import get_config
22from src.models import get_model
23from src.data import get_dataset
24from src.training.metrics import SegmentationMetrics
25from src.utils import plot_training_curves
26
27
28def parse_args():
29    parser = argparse.ArgumentParser(description="Evaluate forgery detection model")
30    
31    parser.add_argument('--model', type=str, required=True,
32                       help='Path to model checkpoint')
33    
34    parser.add_argument('--dataset', type=str, required=True,
35                       choices=['doctamper', 'rtm', 'casia', 'receipts'],
36                       help='Dataset to evaluate on')
37    
38    parser.add_argument('--split', type=str, default='val',
39                       help='Data split (val/test)')
40    
41    parser.add_argument('--output', type=str, default='outputs/evaluation',
42                       help='Output directory')
43    
44    parser.add_argument('--config', type=str, default='config.yaml',
45                       help='Path to config file')
46    
47    return parser.parse_args()
48
49
50def main():
51    args = parse_args()
52    
53    # Load config
54    config = get_config(args.config)
55    
56    # Device
57    device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
58    
59    print("\n" + "="*60)
60    print("Model Evaluation")
61    print("="*60)
62    print(f"Model: {args.model}")
63    print(f"Dataset: {args.dataset}")
64    print(f"Split: {args.split}")
65    print(f"Device: {device}")
66    print("="*60)
67    
68    # Load model
69    model = get_model(config).to(device)
70    checkpoint = torch.load(args.model, map_location=device)
71    
72    if 'model_state_dict' in checkpoint:
73        model.load_state_dict(checkpoint['model_state_dict'])
74    else:
75        model.load_state_dict(checkpoint)
76    
77    model.eval()
78    print("Model loaded")
79    
80    # Load dataset
81    dataset = get_dataset(config, args.dataset, split=args.split)
82    print(f"Dataset loaded: {len(dataset)} samples")
83    
84    # Create output directory
85    output_dir = Path(args.output)
86    output_dir.mkdir(parents=True, exist_ok=True)
87    
88    # Evaluate
89    metrics = SegmentationMetrics()
90    has_pixel_mask = config.has_pixel_mask(args.dataset)
91    
92    print(f"\nEvaluating...")
93    
94    all_ious = []
95    all_dices = []
96    
97    with torch.no_grad():
98        for i in tqdm(range(len(dataset)), desc="Evaluating"):
99            try:
100                image, mask, metadata = dataset[i]
101                
102                # Move to device
103                image = image.unsqueeze(0).to(device)
104                mask = mask.unsqueeze(0).to(device)
105                
106                # Forward pass
107                logits, _ = model(image)
108                probs = torch.sigmoid(logits)
109                
110                # Update metrics
111                if has_pixel_mask:
112                    metrics.update(probs, mask, has_pixel_mask=True)
113                    
114                    # Per-sample metrics
115                    pred_binary = (probs > 0.5).float()
116                    intersection = (pred_binary * mask).sum().item()
117                    union = pred_binary.sum().item() + mask.sum().item() - intersection
118                    
119                    iou = intersection / (union + 1e-8)
120                    dice = (2 * intersection) / (pred_binary.sum().item() + mask.sum().item() + 1e-8)
121                    
122                    all_ious.append(iou)
123                    all_dices.append(dice)
124            
125            except Exception as e:
126                print(f"Error on sample {i}: {e}")
127                continue
128    
129    # Compute final metrics
130    results = metrics.compute()
131    
132    # Add per-sample statistics
133    if has_pixel_mask and all_ious:
134        results['iou_mean'] = np.mean(all_ious)
135        results['iou_std'] = np.std(all_ious)
136        results['dice_mean'] = np.mean(all_dices)
137        results['dice_std'] = np.std(all_dices)
138    
139    # Print results
140    print("\n" + "="*60)
141    print("Evaluation Results")
142    print("="*60)
143    
144    for key, value in results.items():
145        if isinstance(value, float):
146            print(f"  {key}: {value:.4f}")
147    
148    # Save results
149    results_path = output_dir / f'{args.dataset}_{args.split}_results.json'
150    with open(results_path, 'w') as f:
151        json.dump(results, f, indent=2)
152    
153    print(f"\nResults saved to: {results_path}")
154    print("="*60)
155
156
157if __name__ == '__main__':
158    main()
159