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()