DocForg/Document_Forgery_Detection
0
1"""
2Inference Script for Document Forgery Detection
3
4Run inference on single images or entire directories.
5
6Usage:
7 python scripts/run_inference.py --input path/to/image.jpg --model outputs/checkpoints/best_doctamper.pth
8 python scripts/run_inference.py --input path/to/folder/ --model outputs/checkpoints/best_doctamper.pth
9"""
10
11import argparse
12import sys
13from pathlib import Path
14import json
15
16# Add src to path
17sys.path.insert(0, str(Path(__file__).parent.parent))
18
19from src.config import get_config
20from src.inference import get_pipeline
21
22
23def parse_args():
24 parser = argparse.ArgumentParser(description="Run forgery detection inference")
25
26 parser.add_argument('--input', type=str, required=True,
27 help='Input image or directory path')
28
29 parser.add_argument('--model', type=str, required=True,
30 help='Path to localization model checkpoint')
31
32 parser.add_argument('--classifier', type=str, default=None,
33 help='Path to classifier directory (optional)')
34
35 parser.add_argument('--output', type=str, default='outputs/results',
36 help='Output directory')
37
38 parser.add_argument('--is_text', action='store_true',
39 help='Enable OCR features for text documents')
40
41 parser.add_argument('--config', type=str, default='config.yaml',
42 help='Path to config file')
43
44 return parser.parse_args()
45
46
47def process_file(pipeline, input_path: str, output_dir: str):
48 """Process a single file"""
49 try:
50 result = pipeline.run(input_path, output_dir)
51 return result
52 except Exception as e:
53 print(f"Error processing {input_path}: {e}")
54 return None
55
56
57def main():
58 args = parse_args()
59
60 # Load config
61 config = get_config(args.config)
62
63 print("\n" + "="*60)
64 print("Hybrid Document Forgery Detection - Inference")
65 print("="*60)
66 print(f"Input: {args.input}")
67 print(f"Model: {args.model}")
68 print(f"Classifier: {args.classifier or 'None'}")
69 print(f"Output: {args.output}")
70 print("="*60)
71
72 # Create pipeline
73 pipeline = get_pipeline(
74 config,
75 model_path=args.model,
76 classifier_path=args.classifier,
77 is_text_document=args.is_text
78 )
79
80 # Create output directory
81 output_dir = Path(args.output)
82 output_dir.mkdir(parents=True, exist_ok=True)
83
84 # Get input files
85 input_path = Path(args.input)
86
87 if input_path.is_file():
88 files = [input_path]
89 elif input_path.is_dir():
90 extensions = ['.jpg', '.jpeg', '.png', '.pdf', '.bmp', '.tiff']
91 files = [f for f in input_path.iterdir()
92 if f.suffix.lower() in extensions]
93 else:
94 print(f"Invalid input path: {input_path}")
95 return
96
97 print(f"\nProcessing {len(files)} file(s)...")
98
99 # Process files
100 all_results = []
101
102 for file_path in files:
103 result = process_file(pipeline, str(file_path), str(output_dir))
104 if result:
105 all_results.append(result)
106
107 # Print summary
108 status = "TAMPERED" if result['is_tampered'] else "AUTHENTIC"
109 print(f"\n {file_path.name}: {status}")
110 if result['is_tampered']:
111 print(f" Regions detected: {result['num_regions']}")
112 for region in result['regions'][:3]: # Show first 3
113 print(f" - {region['forgery_type']} (conf: {region['confidence']:.2f})")
114
115 # Save summary
116 summary_path = output_dir / 'inference_summary.json'
117 summary = {
118 'total_files': len(files),
119 'processed': len(all_results),
120 'tampered': sum(1 for r in all_results if r['is_tampered']),
121 'authentic': sum(1 for r in all_results if not r['is_tampered']),
122 'results': all_results
123 }
124
125 with open(summary_path, 'w') as f:
126 json.dump(summary, f, indent=2, default=str)
127
128 print("\n" + "="*60)
129 print("Inference Complete!")
130 print(f"Total: {summary['total_files']}, "
131 f"Tampered: {summary['tampered']}, "
132 f"Authentic: {summary['authentic']}")
133 print(f"Results saved to: {output_dir}")
134 print("="*60)
135
136
137if __name__ == '__main__':
138 main()
139 