CoolFace
Apppublic

DocForg/Document_Forgery_Detection

sourceHugging Facemitupdated 6mo agoView on Hugging Face
0likes
run_inference.py139 linesDownload Raw Back to scripts
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