-
Notifications
You must be signed in to change notification settings - Fork 1
Expand file tree
/
Copy pathtrain.py
More file actions
125 lines (98 loc) · 3.38 KB
/
Copy pathtrain.py
File metadata and controls
125 lines (98 loc) · 3.38 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
"""Main training script for noise correction pipeline."""
import argparse
import numpy as np
from sklearn.metrics import accuracy_score
from src.pipeline import NoiseCorrectionPipeline
from src.utils import set_seed
from config import get_config, get_small_config, get_large_config
def parse_args():
"""Parse command line arguments."""
6880
parser = argparse.ArgumentParser(
description='Train noise correction pipeline'
)
parser.add_argument(
'--config',
type=str,
default='default',
choices=['default', 'small', 'large'],
help='Configuration preset to use'
)
parser.add_argument(
'--seed',
type=int,
default=None,
help='Random seed (overrides config)'
)
parser.add_argument(
'--ground_truth_path',
type=str,
default=None,
help='Path to ground truth CSV file (overrides config)'
)
parser.add_argument(
'--features_path',
type=str,
default=None,
help='Path to features feather file (overrides config)'
)
parser.add_argument(
'--batch_size',
type=int,
default=None,
help='Batch size (overrides config)'
)
return parser.parse_args()
def main():
"""Main execution function."""
args = parse_args()
# Load configuration
if args.config == 'small':
config = get_small_config()
elif args.config == 'large':
config = get_large_config()
else:
config = get_config()
# Override with command line arguments
if args.seed is not None:
config['SEED'] = args.seed
if args.ground_truth_path is not None:
config['GROUND_TRUTH_PATH'] = args.ground_truth_path
if args.features_path is not None:
config['FEATURES_PATH'] = args.features_path
if args.batch_size is not None:
config['BATCH_SIZE'] = args.batch_size
# Set random seed
set_seed(config['SEED'])
# Initialize and run pipeline
print("="*60)
print("STARTING LABEL CORRECTION PIPELINE")
print("="*60)
pipeline = NoiseCorrectionPipeline(config)
corrected_labels, true_labels, noisy_labels = pipeline.run()
# Evaluate results
print("\n" + "="*60)
print("FINAL EVALUATION RESULTS")
print("="*60 + "\n")
# Original dataset quality
acc_original = accuracy_score(true_labels, noisy_labels)
noise_rate_original = 1 - acc_original
# Corrected dataset quality
acc_corrected = accuracy_score(true_labels, corrected_labels)
noise_rate_corrected = 1 - acc_corrected
print("--- Dataset Quality BEFORE and AFTER Correction ---")
print(f"📈 ORIGINAL Dataset Accuracy: {acc_original*100:.2f}%")
print(f"🔥 ORIGINAL Noise Rate: {noise_rate_original*100:.2f}%")
print("-" * 40)
print(f"📉 CORRECTED Dataset Accuracy: {acc_corrected*100:.2f}%")
print(f"💧 REMAINING Noise Rate: {noise_rate_corrected*100:.2f}%")
print("-" * 60)
# Calculate improvement
improvement = (acc_corrected - acc_original) * 100
noise_reduction = (noise_rate_original - noise_rate_corrected) * 100
print(f"\n✨ Accuracy Improvement: +{improvement:.2f}%")
print(f"✨ Noise Reduction: -{noise_reduction:.2f}%")
print("\n" + "="*60)
print("CO
32B9
MPLETED!")
print("="*60)
if __name__ == "__main__":
main()