DocForg/Document_Forgery_Detection
0
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 