| import json |
| import argparse |
| from pathlib import Path |
| from typing import List, Dict |
| from difflib import SequenceMatcher |
| import numpy as np |
| from collections import defaultdict |
|
|
|
|
| class QueryVariantVerifier: |
| |
| def __init__(self, min_similarity: float = 0.2, max_similarity: float = 0.8): |
| self.min_similarity = min_similarity |
| self.max_similarity = max_similarity |
| |
| def compute_similarity(self, text1: str, text2: str) -> float: |
| return SequenceMatcher(None, text1.lower(), text2.lower()).ratio() |
| |
| def verify_single_item(self, item: Dict) -> Dict: |
| result = { |
| 'has_variants': False, |
| 'num_variants': 0, |
| 'similarities': [], |
| 'issues': [], |
| 'quality_score': 0.0, |
| } |
| |
| if 'query_variants' not in item or not item['query_variants']: |
| result['issues'].append("缺少query_variants字段") |
| return result |
| |
| result['has_variants'] = True |
| result['num_variants'] = len(item['query_variants']) |
| |
| original = item.get('query_original', item.get('question', '')) |
| if not original: |
| result['issues'].append("缺少原始query") |
| return result |
| |
| variants = item['query_variants'] |
| |
| for i, variant in enumerate(variants): |
| if not variant or not variant.strip(): |
| result['issues'].append(f"Variant {i+1} 为空") |
| continue |
| |
| sim = self.compute_similarity(original, variant) |
| result['similarities'].append(sim) |
| |
| if sim < self.min_similarity: |
| result['issues'].append( |
| f"Variant {i+1} 相似度过低 ({sim:.2f}): 可能偏离原意" |
| ) |
| elif sim > self.max_similarity: |
| result['issues'].append( |
| f"Variant {i+1} 相似度过高 ({sim:.2f}): 过于相似" |
| ) |
| |
| if variant.lower().strip() == original.lower().strip(): |
| result['issues'].append(f"Variant {i+1} 与原query完全相同") |
| |
| if result['similarities']: |
| avg_sim = np.mean(result['similarities']) |
| if 0.4 <= avg_sim <= 0.6: |
| result['quality_score'] = 100 |
| elif 0.3 <= avg_sim < 0.4 or 0.6 < avg_sim <= 0.7: |
| result['quality_score'] = 80 |
| elif 0.2 <= avg_sim < 0.3 or 0.7 < avg_sim <= 0.8: |
| result['quality_score'] = 60 |
| else: |
| result['quality_score'] = 40 |
| |
| return result |
| |
| def verify_dataset(self, data_path: str) -> Dict: |
| print(f"\n{'='*60}") |
| print(f"验证数据集: {data_path}") |
| print(f"{'='*60}\n") |
| |
| with open(data_path, 'r', encoding='utf-8') as f: |
| if data_path.endswith('.jsonl'): |
| data = [json.loads(line) for line in f] |
| else: |
| data = json.load(f) |
| |
| results = [] |
| all_issues = defaultdict(int) |
| |
| for idx, item in enumerate(data): |
| result = self.verify_single_item(item) |
| result['idx'] = idx |
| results.append(result) |
| |
| for issue in result['issues']: |
| all_issues[issue] += 1 |
| |
| stats = self._compute_stats(results) |
| |
| self._print_report(stats, all_issues, data_path) |
| |
| self._print_examples(data, results) |
| |
| return { |
| 'stats': stats, |
| 'issues': dict(all_issues), |
| 'results': results, |
| } |
| |
| def _compute_stats(self, results: List[Dict]) -> Dict: |
| total = len(results) |
| with_variants = sum(1 for r in results if r['has_variants']) |
| |
| all_similarities = [] |
| all_quality_scores = [] |
| |
| for r in results: |
| all_similarities.extend(r['similarities']) |
| if r['quality_score'] > 0: |
| all_quality_scores.append(r['quality_score']) |
| |
| return { |
| 'total_samples': total, |
| 'with_variants': with_variants, |
| 'without_variants': total - with_variants, |
| 'avg_variants_per_query': np.mean([r['num_variants'] for r in results]), |
| 'avg_similarity': np.mean(all_similarities) if all_similarities else 0, |
| 'similarity_std': np.std(all_similarities) if all_similarities else 0, |
| 'avg_quality_score': np.mean(all_quality_scores) if all_quality_scores else 0, |
| 'samples_with_issues': sum(1 for r in results if r['issues']), |
| } |
| |
| def _print_report(self, stats: Dict, issues: Dict, data_path: str): |
| print("📊 验证报告") |
| print("="*60) |
| |
| print("\n基本统计:") |
| print(f" 总样本数: {stats['total_samples']}") |
| print(f" 包含variants: {stats['with_variants']} " + |
| f"({stats['with_variants']/stats['total_samples']*100:.1f}%)") |
| print(f" 缺失variants: {stats['without_variants']}") |
| print(f" 平均variants数/query: {stats['avg_variants_per_query']:.2f}") |
| |
| print("\n质量统计:") |
| print(f" 平均相似度: {stats['avg_similarity']:.3f} " + |
| f"(± {stats['similarity_std']:.3f})") |
| print(f" 平均质量分数: {stats['avg_quality_score']:.1f}/100") |
| print(f" 存在问题的样本: {stats['samples_with_issues']} " + |
| f"({stats['samples_with_issues']/stats['total_samples']*100:.1f}%)") |
| |
| print("\n质量评估:") |
| avg_sim = stats['avg_similarity'] |
| if 0.4 <= avg_sim <= 0.6: |
| print(" ✅ 优秀 - 相似度在理想范围内") |
| elif 0.3 <= avg_sim < 0.4 or 0.6 < avg_sim <= 0.7: |
| print(" ✓ 良好 - 相似度可接受") |
| elif 0.2 <= avg_sim < 0.3 or 0.7 < avg_sim <= 0.8: |
| print(" ⚠️ 一般 - 相似度需要改进") |
| else: |
| print(" ❌ 差 - 相似度不合理,需要重新生成") |
| |
| if issues: |
| print("\n常见问题:") |
| sorted_issues = sorted(issues.items(), key=lambda x: x[1], reverse=True) |
| for issue, count in sorted_issues[:5]: |
| print(f" • {issue}: {count} 次") |
| |
| print("\n" + "="*60) |
| |
| def _print_examples(self, data: List[Dict], results: List[Dict], num_examples: int = 3): |
| print("\n📝 示例") |
| print("="*60) |
| |
| good_examples = [r for r in results if r['quality_score'] >= 80 and not r['issues']] |
| medium_examples = [r for r in results if 50 <= r['quality_score'] < 80] |
| bad_examples = [r for r in results if r['quality_score'] < 50 or r['issues']] |
| |
| if good_examples: |
| print("\n✅ 优质示例:") |
| for i, result in enumerate(good_examples[:2], 1): |
| idx = result['idx'] |
| item = data[idx] |
| original = item.get('query_original', item.get('question', '')) |
| |
| print(f"\n示例 {i}:") |
| print(f" 原始: {original}") |
| for j, (variant, sim) in enumerate(zip(item['query_variants'], result['similarities']), 1): |
| print(f" 变体{j}: {variant} (相似度: {sim:.2f})") |
| |
| if bad_examples: |
| print("\n⚠️ 需要改进的示例:") |
| for i, result in enumerate(bad_examples[:2], 1): |
| idx = result['idx'] |
| item = data[idx] |
| original = item.get('query_original', item.get('question', '')) |
| |
| print(f"\n示例 {i}:") |
| print(f" 原始: {original}") |
| if 'query_variants' in item: |
| for j, variant in enumerate(item['query_variants'], 1): |
| print(f" 变体{j}: {variant}") |
| print(f" 问题: {', '.join(result['issues'][:3])}") |
| |
| print("\n" + "="*60) |
|
|
|
|
| def main(): |
| parser = argparse.ArgumentParser( |
| description="验证生成的Query Variants质量" |
| ) |
| |
| parser.add_argument( |
| "--input", |
| type=str, |
| required=True, |
| help="输入数据文件路径" |
| ) |
| parser.add_argument( |
| "--min_similarity", |
| type=float, |
| default=0.2, |
| help="最小相似度阈值 (默认: 0.2)" |
| ) |
| parser.add_argument( |
| "--max_similarity", |
| type=float, |
| default=0.8, |
| help="最大相似度阈值 (默认: 0.8)" |
| ) |
| parser.add_argument( |
| "--output_report", |
| type=str, |
| default=None, |
| help="输出报告文件路径 (JSON格式)" |
| ) |
| |
| args = parser.parse_args() |
| |
| verifier = QueryVariantVerifier( |
| min_similarity=args.min_similarity, |
| max_similarity=args.max_similarity, |
| ) |
| |
| report = verifier.verify_dataset(args.input) |
| |
| if args.output_report: |
| output_dir = Path(args.output_report).parent |
| output_dir.mkdir(parents=True, exist_ok=True) |
| |
| with open(args.output_report, 'w', encoding='utf-8') as f: |
| json.dump(report, f, ensure_ascii=False, indent=2) |
| |
| print(f"\n✓ 报告已保存到: {args.output_report}") |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|
|
|