|
|
| echo "================================================" |
| echo "Invar-RAG Query Variants 生成示例" |
| echo "================================================" |
|
|
| |
| DEEPSEEK_API_KEY="${DEEPSEEK_API_KEY:-your_api_key_here}" |
| BASE_DIR="/Users/dwayneryan/Downloads/RAG-main/opencompass" |
| DATA_DIR="$BASE_DIR/data/my_datasets" |
| OUTPUT_DIR="$BASE_DIR/data/variants" |
| CACHE_DIR="$BASE_DIR/cache/query_variants" |
|
|
| |
| if [ "$DEEPSEEK_API_KEY" = "your_api_key_here" ]; then |
| echo "错误: 请设置DEEPSEEK_API_KEY环境变量" |
| echo "export DEEPSEEK_API_KEY='your_actual_key'" |
| exit 1 |
| fi |
|
|
| |
| mkdir -p "$OUTPUT_DIR" |
| mkdir -p "$CACHE_DIR" |
|
|
| echo "" |
| echo "示例 1: 处理100个样本(测试)" |
| echo "----------------------------------------" |
|
|
| python "$BASE_DIR/scripts/prepare_invarrag_data.py" convert_nq \ |
| --input "$DATA_DIR/nq/nq-test.qa.csv" \ |
| --output "$OUTPUT_DIR/nq-test.json" \ |
| --max_samples 100 |
|
|
| python "$BASE_DIR/scripts/generate_query_variants_deepseek.py" \ |
| --api_key "$DEEPSEEK_API_KEY" \ |
| --input "$OUTPUT_DIR/nq-test.json" \ |
| --output "$OUTPUT_DIR/nq-test-sample-with-variants.json" \ |
| --num_variants 3 \ |
| --max_samples 100 \ |
| --cache_dir "$CACHE_DIR/nq" |
|
|
| echo "✓ 示例 1 完成" |
|
|
| echo "" |
| echo "示例 2: 处理完整NQ训练集" |
| echo "----------------------------------------" |
|
|
| python "$BASE_DIR/scripts/prepare_invarrag_data.py" convert_nq \ |
| --input "$DATA_DIR/nq/nq-train.qa.csv" \ |
| --output "$OUTPUT_DIR/nq-train.json" |
|
|
| python "$BASE_DIR/scripts/generate_query_variants_deepseek.py" \ |
| --api_key "$DEEPSEEK_API_KEY" \ |
| --input "$OUTPUT_DIR/nq-train.json" \ |
| --output "$OUTPUT_DIR/nq-train-with-variants.json" \ |
| --num_variants 3 \ |
| --cache_dir "$CACHE_DIR/nq" |
|
|
| echo "✓ 示例 2 完成" |
|
|
| echo "" |
| echo "示例 3: 批量处理所有数据集" |
| echo "----------------------------------------" |
|
|
| DATASETS=( |
| "nq:nq-train.qa.csv:nq-train.json:nq-train-with-variants.json:convert_nq" |
| "tqa:trivia-train.qa.csv:tqa-train.json:tqa-train-with-variants.json:convert_triviaqa" |
| ) |
|
|
| for dataset in "${DATASETS[@]}"; do |
| IFS=':' read -r name input_file json_file output_file convert_cmd <<< "$dataset" |
| |
| echo "" |
| echo "处理: $name" |
| |
| python "$BASE_DIR/scripts/prepare_invarrag_data.py" $convert_cmd \ |
| --input "$DATA_DIR/$name/$input_file" \ |
| --output "$OUTPUT_DIR/$json_file" |
|
|
| python "$BASE_DIR/scripts/generate_query_variants_deepseek.py" \ |
| --api_key "$DEEPSEEK_API_KEY" \ |
| --input "$OUTPUT_DIR/$json_file" \ |
| --output "$OUTPUT_DIR/$output_file" \ |
| --num_variants 3 \ |
| --cache_dir "$CACHE_DIR/$name" |
| |
| echo "✓ $name 完成" |
| done |
|
|
| echo "" |
| echo "验证生成的数据" |
| echo "----------------------------------------" |
|
|
| python - << EOF |
| import json |
| import os |
| |
| output_dir = "$OUTPUT_DIR" |
| files = [f for f in os.listdir(output_dir) if f.endswith('.json')] |
| |
| print(f"生成的文件: {len(files)}") |
| print("") |
| |
| for file in files: |
| path = os.path.join(output_dir, file) |
| with open(path) as f: |
| data = json.load(f) |
| |
| num_samples = len(data) |
| num_with_variants = sum(1 for d in data if 'query_variants' in d and d['query_variants']) |
| avg_variants = sum(len(d.get('query_variants', [])) for d in data) / num_samples if num_samples > 0 else 0 |
| |
| print(f"文件: {file}") |
| print(f" 样本数: {num_samples}") |
| print(f" 包含variants: {num_with_variants} ({num_with_variants/num_samples*100:.1f}%)") |
| print(f" 平均variants数: {avg_variants:.2f}") |
| print("") |
| EOF |
|
|
| echo "" |
| echo "缓存统计" |
| echo "----------------------------------------" |
|
|
| total_cache_size=$(du -sh "$CACHE_DIR" | cut -f1) |
| cache_files=$(find "$CACHE_DIR" -name "*.json" | wc -l) |
|
|
| echo "缓存目录: $CACHE_DIR" |
| echo "缓存文件数: $cache_files" |
| echo "缓存大小: $total_cache_size" |
|
|
| echo "" |
| echo "================================================" |
| echo "✓ 所有示例完成" |
| echo "================================================" |
| echo "" |
| echo "生成的文件位于: $OUTPUT_DIR" |
| echo "缓存文件位于: $CACHE_DIR" |
| echo "" |
| echo "下一步:" |
| echo "1. 检查生成的variants质量" |
| echo "2. 使用variants数据训练模型:" |
| echo " python invarrag_train_retriever.py \\" |
| echo " --train_data $OUTPUT_DIR/nq-train-with-variants.json \\" |
| echo " --corpus data/corpus.jsonl \\" |
| echo " ..." |
| echo "" |
|
|
|
|