CoolFace
Apppublic

DocForg/Document_Forgery_Detection

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