File size: 7,321 Bytes
1e05592
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
#!/usr/bin/env python3
"""
分析annotation质量
找出少于20字符的简单标注,以便调整prompt策略
"""

import json
from pathlib import Path
from collections import defaultdict


class AnnotationAnalyzer:
    """标注分析器"""
    
    def __init__(self, threshold_chars=20):
        self.threshold = threshold_chars
        self.stats = defaultdict(list)
    
    def analyze_annotation(self, anno_file: Path):
        """分析单个annotation文件"""
        with open(anno_file, 'r') as f:
            data = json.load(f)
        
        case_id = data.get('id', anno_file.parent.name)
        accident_type = data.get('accident_type', '')
        
        if not accident_type or accident_type.lower() in ['null', 'none', 'unknown', '']:
            return None, 'empty'
        
        # 清理标注
        accident_type = accident_type.strip()
        char_count = len(accident_type)
        word_count = len(accident_type.split())
        
        # 分类
        if char_count < self.threshold:
            category = 'short'  # 简单标注
        else:
            category = 'detailed'  # 详细标注
        
        return {
            'case_id': case_id,
            'dataset': data.get('dataset', 'unknown'),
            'accident_type': accident_type,
            'char_count': char_count,
            'word_count': word_count,
            'category': category,
            'file': str(anno_file)
        }, category
    
    def analyze_dataset(self, dataset_root: Path, dataset_name: str):
        """分析整个数据集"""
        print(f"\n分析 {dataset_name}...")
        
        anno_files = list(dataset_root.rglob("annotation.json"))
        print(f"找到 {len(anno_files)} 个annotation文件")
        
        results = {
            'short': [],
            'detailed': [],
            'empty': []
        }
        
        for anno_file in anno_files:
            try:
                info, category = self.analyze_annotation(anno_file)
                if info:
                    results[category].append(info)
            except Exception as e:
                print(f"处理失败 {anno_file}: {e}")
        
        return results
    
    def print_summary(self, results: dict, dataset_name: str):
        """打印统计摘要"""
        total = sum(len(results[cat]) for cat in ['short', 'detailed', 'empty'])
        
        print(f"\n{'='*70}")
        print(f"{dataset_name} - 标注质量统计")
        print("=" * 70)
        print(f"总计: {total} 案例")
        print(f"  简单标注 (<{self.threshold}字符): {len(results['short'])} ({len(results['short'])/total*100:.1f}%)")
        print(f"  详细标注 (>={self.threshold}字符): {len(results['detailed'])} ({len(results['detailed'])/total*100:.1f}%)")
        print(f"  空标注: {len(results['empty'])} ({len(results['empty'])/total*100:.1f}%)")
    
    def print_examples(self, results: dict, n=10):
        """打印示例"""
        print(f"\n{'='*70}")
        print("简单标注示例 (前{}):".format(min(n, len(results['short']))))
        print("=" * 70)
        
        # 按字符数排序
        short_sorted = sorted(results['short'], key=lambda x: x['char_count'])
        
        for i, item in enumerate(short_sorted[:n], 1):
            print(f"\n{i}. [{item['char_count']}字符, {item['word_count']}词]")
            print(f"   案例: {item['case_id']}")
            print(f"   标注: \"{item['accident_type']}\"")
        
        print(f"\n{'='*70}")
        print("详细标注示例 (前5):")
        print("=" * 70)
        
        # 按字符数排序 (降序)
        detailed_sorted = sorted(results['detailed'], key=lambda x: x['char_count'], reverse=True)
        
        for i, item in enumerate(detailed_sorted[:5], 1):
            print(f"\n{i}. [{item['char_count']}字符, {item['word_count']}词]")
            print(f"   案例: {item['case_id']}")
            print(f"   标注: \"{item['accident_type'][:100]}...\"" if len(item['accident_type']) > 100 
                  else f"   标注: \"{item['accident_type']}\"")
    
    def save_analysis(self, results: dict, output_file: Path):
        """保存分析结果"""
        analysis = {
            'threshold': self.threshold,
            'short_annotations': results['short'],
            'detailed_annotations': results['detailed'],
            'empty_annotations': results['empty'],
            'statistics': {
                'total': sum(len(results[cat]) for cat in ['short', 'detailed', 'empty']),
                'short_count': len(results['short']),
                'detailed_count': len(results['detailed']),
                'empty_count': len(results['empty'])
            }
        }
        
        with open(output_file, 'w') as f:
            json.dump(analysis, f, indent=2)
        
        print(f"\n✓ 分析结果保存到: {output_file}")


def main():
    """主函数"""
    print("=" * 70)
    print("Annotation质量分析")
    print("阈值: 20字符")
    print("=" * 70)
    
    analyzer = AnnotationAnalyzer(threshold_chars=20)
    
    all_results = {
        'short': [],
        'detailed': [],
        'empty': []
    }
    
    # 分析DADA-2000
    dada_root = Path("PROJECT_ROOT/data/dataset/pretrain/DADA-2000")
    if dada_root.exists():
        dada_results = analyzer.analyze_dataset(dada_root, "DADA-2000")
        analyzer.print_summary(dada_results, "DADA-2000")
        
        for cat in ['short', 'detailed', 'empty']:
            all_results[cat].extend(dada_results[cat])
    
    # 分析NEXAR
    nexar_root = Path("PROJECT_ROOT/data/dataset/pretrain/nexar")
    if nexar_root.exists():
        nexar_results = analyzer.analyze_dataset(nexar_root, "NEXAR")
        analyzer.print_summary(nexar_results, "NEXAR")
        
        for cat in ['short', 'detailed', 'empty']:
            all_results[cat].extend(nexar_results[cat])
    
    # 总体统计
    analyzer.print_summary(all_results, "总体")
    
    # 打印示例
    analyzer.print_examples(all_results, n=15)
    
    # 保存分析结果
    output_dir = Path("PROJECT_ROOT/data/dataset/pretrain/train")
    output_dir.mkdir(parents=True, exist_ok=True)
    analyzer.save_analysis(all_results, output_dir / "annotation_analysis.json")
    
    # 生成prompt策略建议
    print("\n" + "=" * 70)
    print("建议的Prompt策略")
    print("=" * 70)
    
    print("\n简单标注 (<20字符) - 使用简单prompt:")
    print("  - 'What object or vehicle was involved in this accident?'")
    print("  - 'Identify the main entity in this traffic incident.'")
    print("  - 'What type of collision is shown? (e.g., vehicle, pedestrian, bicycle)'")
    
    print("\n详细标注 (>=20字符) - 使用详细prompt:")
    print("  - 'Describe the accident in this image. What happened and why?'")
    print("  - 'Provide a detailed description of the traffic incident.'")
    print("  - 'Explain what led to this accident and what occurred.'")
    
    print("\n" + "=" * 70)
    print("✅ 分析完成!")
    print("=" * 70)
    print("\n下一步:")
    print("1. 查看 annotation_analysis.json 了解详细情况")
    print("2. 运行 prepare_pretrain_data_adaptive.py 生成自适应prompt数据")


if __name__ == "__main__":
    main()