-
Notifications
You must be signed in to change notification settings - Fork 10
Expand file tree
/
Copy pathcheck_classes.py
More file actions
124 lines (104 loc) Β· 4.83 KB
/
Copy pathcheck_classes.py
File metadata and controls
124 lines (104 loc) Β· 4.83 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
#!/usr/bin/env python3
"""
Check if the classes used in training match the data.yaml configuration
"""
import os
import yaml
from collections import Counter
from pathlib import Path
def load_data_yaml(yaml_path):
"""Load data.yaml configuration"""
with open(yaml_path, 'r') as f:
data = yaml.safe_load(f)
return data
def check_label_files(labels_dir):
"""Check all label files and extract class IDs"""
label_files = list(Path(labels_dir).glob("*.txt"))
all_classes = []
for label_file in label_files:
with open(label_file, 'r') as f:
for line in f:
if line.strip():
class_id = int(float(line.strip().split()[0]))
all_classes.append(class_id)
return all_classes
def main():
"""Main function to check class consistency"""
print("π Checking class configuration consistency...")
# Load data.yaml
data_yaml_path = 'notebooks/data/processed/data.yaml'
data_config = load_data_yaml(data_yaml_path)
print(f"\nπ Data.yaml configuration:")
print(f" Number of classes (nc): {data_config['nc']}")
print(f" Class names: {data_config['names']}")
print(f" Expected class IDs: 0-{data_config['nc']-1}")
# Check training labels
train_labels_dir = 'notebooks/data/processed/train/labels'
val_labels_dir = 'notebooks/data/processed/val/labels'
test_labels_dir = 'notebooks/data/processed/test/labels'
datasets = {
'training': train_labels_dir,
'validation': val_labels_dir,
'test': test_labels_dir
}
for dataset_name, labels_dir in datasets.items():
if os.path.exists(labels_dir):
print(f"\nπ {dataset_name.title()} dataset:")
classes_found = check_label_files(labels_dir)
if classes_found:
class_counts = Counter(classes_found)
min_class = min(classes_found)
max_class = max(classes_found)
unique_classes = len(set(classes_found))
print(f" Total annotations: {len(classes_found)}")
print(f" Unique classes found: {unique_classes}")
print(f" Class ID range: {min_class} - {max_class}")
# Check if any class IDs are outside expected range
expected_range = set(range(data_config['nc']))
found_classes = set(classes_found)
if found_classes <= expected_range:
print(f" β
All class IDs are within expected range (0-{data_config['nc']-1})")
else:
unexpected = found_classes - expected_range
print(f" β Unexpected class IDs found: {unexpected}")
# Check if any expected classes are missing
missing = expected_range - found_classes
if missing:
print(f" β οΈ Missing class IDs: {sorted(missing)}")
# Map missing IDs to class names
missing_names = [data_config['names'][i] for i in sorted(missing) if i < len(data_config['names'])]
print(f" β οΈ Missing class names: {missing_names}")
else:
print(f" β
All expected classes are present")
# Show top 10 most frequent classes
print(f" π Top 10 most frequent classes:")
for class_id, count in class_counts.most_common(10):
class_name = data_config['names'][class_id] if class_id < len(data_config['names']) else f'unknown_{class_id}'
print(f" Class {class_id} ({class_name}): {count} instances")
else:
print(f" β No labels found")
else:
print(f"\nβ {dataset_name.title()} labels directory not found: {labels_dir}")
# Additional validation - check a trained model
print(f"\nπ€ Checking trained model class configuration...")
try:
from ultralytics import YOLO
model_path = 'trained_models_v2/yolo11m_best.pt'
if os.path.exists(model_path):
model = YOLO(model_path)
model_names = model.names
print(f" Model class count: {len(model_names)}")
print(f" Model class names: {list(model_names.values())}")
# Compare with data.yaml
if list(model_names.values()) == data_config['names']:
print(f" β
Model classes match data.yaml exactly")
else:
print(f" β Model classes differ from data.yaml")
print(f" Expected: {data_config['names']}")
print(f" Model has: {list(model_names.values())}")
else:
print(f" β οΈ Model not found: {model_path}")
except Exception as e:
print(f" β Error loading model: {e}")
if __name__ == "__main__":
main()