CoolFace
Apppublic

DocForg/Document_Forgery_Detection

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