DocForg/Document_Forgery_Detection
0
1"""
2Chunked Training Script for Document Forgery Detection
3
4Supports training on large datasets (DocTamper) in chunks to manage RAM constraints.
5Usage:
6 python scripts/train_chunked.py --dataset doctamper --chunk 1
7 python scripts/train_chunked.py --dataset rtm
8 python scripts/train_chunked.py --dataset casia
9 python scripts/train_chunked.py --dataset receipts
10"""
11
12import argparse
13import os
14import sys
15from pathlib import Path
16
17# Add src to path
18sys.path.insert(0, str(Path(__file__).parent.parent))
19
20import torch
21import gc
22
23from src.config import get_config
24from src.training import get_trainer
25from src.utils import plot_training_curves, plot_chunked_training_progress, generate_training_report
26
27
28def parse_args():
29 parser = argparse.ArgumentParser(description="Train forgery detection model")
30
31 parser.add_argument('--dataset', type=str, default='doctamper',
32 choices=['doctamper', 'rtm', 'casia', 'receipts', 'fcd', 'scd'],
33 help='Dataset to train on')
34
35 parser.add_argument('--chunk', type=int, default=None,
36 help='Chunk number (1-4) for DocTamper chunked training')
37
38 parser.add_argument('--epochs', type=int, default=None,
39 help='Number of epochs (overrides config)')
40
41 parser.add_argument('--resume', type=str, default=None,
42 help='Checkpoint to resume from')
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 train_chunk(config, dataset_name: str, chunk_id: int, epochs: int = None, resume: str = None):
51 """Train a single chunk"""
52
53 # Calculate chunk boundaries
54 chunks = config.get('data.chunked_training.chunks', [])
55
56 if chunk_id > len(chunks):
57 raise ValueError(f"Invalid chunk ID: {chunk_id}. Max: {len(chunks)}")
58
59 chunk_config = chunks[chunk_id - 1]
60 chunk_start = chunk_config['start']
61 chunk_end = chunk_config['end']
62 chunk_name = chunk_config['name']
63
64 print(f"\n{'='*60}")
65 print(f"Training Chunk {chunk_id}: {chunk_name}")
66 print(f"Range: {chunk_start*100:.0f}% - {chunk_end*100:.0f}%")
67 print(f"{'='*60}")
68
69 # Create trainer
70 trainer = get_trainer(config, dataset_name)
71
72 # Resume from previous chunk if applicable
73 if resume:
74 # For chunked training, reset epoch counter to train full epochs on new data
75 trainer.load_checkpoint(resume, reset_epoch=True)
76 elif chunk_id > 1:
77 # Auto-resume from previous chunk
78 prev_checkpoint = f'{dataset_name}_chunk{chunk_id-1}_final.pth'
79 if (Path(config.get('outputs.checkpoints')) / prev_checkpoint).exists():
80 print(f"Auto-resuming from previous chunk: {prev_checkpoint}")
81 trainer.load_checkpoint(prev_checkpoint, reset_epoch=True)
82
83 # Train
84 history = trainer.train(
85 epochs=epochs,
86 chunk_start=chunk_start,
87 chunk_end=chunk_end,
88 chunk_id=chunk_id,
89 resume_from=None # Already loaded above
90 )
91
92 # Plot training curves
93 plot_dir = Path(config.get('outputs.plots', 'outputs/plots'))
94 plot_dir.mkdir(parents=True, exist_ok=True)
95
96 plot_path = plot_dir / f'{dataset_name}_chunk{chunk_id}_curves.png'
97 plot_training_curves(
98 history,
99 str(plot_path),
100 title=f"{dataset_name.upper()} Chunk {chunk_id} Training"
101 )
102
103 # Generate report
104 report_path = plot_dir / f'{dataset_name}_chunk{chunk_id}_report.txt'
105 generate_training_report(history, str(report_path), f"{dataset_name} Chunk {chunk_id}")
106
107 # Clear memory
108 del trainer
109 gc.collect()
110 torch.cuda.empty_cache()
111
112 return history
113
114
115def train_full_dataset(config, dataset_name: str, epochs: int = None, resume: str = None):
116 """Train on full dataset (for smaller datasets)"""
117
118 print(f"\n{'='*60}")
119 print(f"Training on: {dataset_name.upper()}")
120 print(f"{'='*60}")
121
122 # Create trainer
123 trainer = get_trainer(config, dataset_name)
124
125 # Load checkpoint if resuming (reset epoch counter for new dataset)
126 if resume:
127 print(f"Loading weights from: {resume}")
128 trainer.load_checkpoint(resume, reset_epoch=True)
129 print("Epoch counter reset to 0 for new dataset training")
130
131 # Train
132 history = trainer.train(
133 epochs=epochs,
134 chunk_id=0,
135 resume_from=None # Already loaded above
136 )
137
138 # Plot training curves
139 plot_dir = Path(config.get('outputs.plots', 'outputs/plots'))
140 plot_dir.mkdir(parents=True, exist_ok=True)
141
142 plot_path = plot_dir / f'{dataset_name}_training_curves.png'
143 plot_training_curves(
144 history,
145 str(plot_path),
146 title=f"{dataset_name.upper()} Training"
147 )
148
149 # Generate report
150 report_path = plot_dir / f'{dataset_name}_report.txt'
151 generate_training_report(history, str(report_path), dataset_name)
152
153 return history
154
155
156def main():
157 args = parse_args()
158
159 # Load config
160 config = get_config(args.config)
161
162 print("\n" + "="*60)
163 print("Hybrid Document Forgery Detection - Training")
164 print("="*60)
165 print(f"Dataset: {args.dataset}")
166 print(f"Device: {config.get('system.device')}")
167 print(f"CUDA Available: {torch.cuda.is_available()}")
168 if torch.cuda.is_available():
169 print(f"GPU: {torch.cuda.get_device_name(0)}")
170 print(f"GPU Memory: {torch.cuda.get_device_properties(0).total_memory / 1e9:.2f} GB")
171 print("="*60)
172
173 # DocTamper: chunked training
174 if args.dataset == 'doctamper' and args.chunk is not None:
175 history = train_chunk(
176 config,
177 args.dataset,
178 args.chunk,
179 epochs=args.epochs,
180 resume=args.resume
181 )
182
183 # DocTamper: all chunks sequentially
184 elif args.dataset == 'doctamper' and args.chunk is None:
185 print("Training DocTamper in 4 chunks...")
186
187 all_histories = []
188 for chunk_id in range(1, 5):
189 history = train_chunk(
190 config,
191 args.dataset,
192 chunk_id,
193 epochs=args.epochs,
194 resume=None if chunk_id == 1 else None # Auto-resume from prev chunk
195 )
196 all_histories.append(history)
197
198 # Plot combined progress
199 plot_dir = Path(config.get('outputs.plots', 'outputs/plots'))
200 combined_path = plot_dir / 'doctamper_all_chunks_progress.png'
201 plot_chunked_training_progress(
202 all_histories,
203 str(combined_path),
204 title="DocTamper Chunked Training Progress"
205 )
206
207 # Other datasets: full training
208 else:
209 history = train_full_dataset(
210 config,
211 args.dataset,
212 epochs=args.epochs,
213 resume=args.resume
214 )
215
216 print("\n" + "="*60)
217 print("Training Complete!")
218 print("="*60)
219
220
221if __name__ == '__main__':
222 main()
223 