DocForg/Document_Forgery_Detection
0
1"""
2Model Export Script
3
4Export trained model to ONNX format for deployment.
5
6Usage:
7 python scripts/export_model.py --model outputs/checkpoints/best_doctamper.pth --format onnx
8"""
9
10import argparse
11import sys
12from pathlib import Path
13
14# Add src to path
15sys.path.insert(0, str(Path(__file__).parent.parent))
16
17import torch
18
19from src.config import get_config
20from src.models import get_model
21from src.utils import export_to_onnx, export_to_torchscript
22
23
24def parse_args():
25 parser = argparse.ArgumentParser(description="Export model for deployment")
26
27 parser.add_argument('--model', type=str, required=True,
28 help='Path to model checkpoint')
29
30 parser.add_argument('--format', type=str, default='onnx',
31 choices=['onnx', 'torchscript', 'both'],
32 help='Export format')
33
34 parser.add_argument('--output', type=str, default='outputs/exported',
35 help='Output directory')
36
37 parser.add_argument('--config', type=str, default='config.yaml',
38 help='Path to config file')
39
40 return parser.parse_args()
41
42
43def main():
44 args = parse_args()
45
46 # Load config
47 config = get_config(args.config)
48
49 print("\n" + "="*60)
50 print("Model Export")
51 print("="*60)
52 print(f"Model: {args.model}")
53 print(f"Format: {args.format}")
54 print("="*60)
55
56 # Create output directory
57 output_dir = Path(args.output)
58 output_dir.mkdir(parents=True, exist_ok=True)
59
60 # Load model
61 device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
62 model = get_model(config).to(device)
63
64 checkpoint = torch.load(args.model, map_location=device)
65 if 'model_state_dict' in checkpoint:
66 model.load_state_dict(checkpoint['model_state_dict'])
67 else:
68 model.load_state_dict(checkpoint)
69
70 model.eval()
71 print("Model loaded")
72
73 # Get image size
74 image_size = config.get('data.image_size', 384)
75
76 # Export
77 if args.format in ['onnx', 'both']:
78 onnx_path = output_dir / 'model.onnx'
79 export_to_onnx(model, str(onnx_path), input_size=(image_size, image_size))
80
81 if args.format in ['torchscript', 'both']:
82 ts_path = output_dir / 'model.pt'
83 export_to_torchscript(model, str(ts_path), input_size=(image_size, image_size))
84
85 print("\n" + "="*60)
86 print("Export Complete!")
87 print(f"Output: {output_dir}")
88 print("="*60)
89
90
91if __name__ == '__main__':
92 main()
93 