InvarRAG / opencompass /scripts /generate_query_variants_deepseek.py
Ryanlal's picture
Upload 1923 files
258adb2 verified
Raw
History Blame Contribute Delete
9.98 kB
import json
import os
import time
import argparse
from pathlib import Path
from typing import List, Dict, Optional
from tqdm import tqdm
import hashlib
import sys
sys.path.append(os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
from opencompass.models.deepseek_api import DeepseekAPI
class QueryVariantGenerator:
def __init__(
self,
api_key: str,
api_url: str = "https://api.deepseek.com/v1/chat/completions",
model: str = "deepseek-chat",
cache_dir: str = "./cache/query_variants",
num_variants: int = 3,
):
self.api = DeepseekAPI(
path=model,
key=api_key,
url=api_url,
query_per_second=1,
retry=3,
)
self.cache_dir = Path(cache_dir)
self.cache_dir.mkdir(parents=True, exist_ok=True)
self.cache_file = self.cache_dir / "query_variants_cache.json"
self.num_variants = num_variants
self.cache = self._load_cache()
def _load_cache(self) -> Dict:
if self.cache_file.exists():
with open(self.cache_file, 'r', encoding='utf-8') as f:
return json.load(f)
return {}
def _save_cache(self):
with open(self.cache_file, 'w', encoding='utf-8') as f:
json.dump(self.cache, f, ensure_ascii=False, indent=2)
def _get_cache_key(self, query: str, variant_type: str) -> str:
content = f"{query}_{variant_type}_{self.num_variants}"
return hashlib.md5(content.encode()).hexdigest()
def _create_rewrite_prompts(self) -> List[Dict]:
prompts = [
{
"name": "paraphrase",
"instruction": """Please paraphrase the following question while keeping its exact meaning and intent.
Rules:
1. Maintain the same semantic meaning
2. Change the sentence structure and word choice
3. Keep the question type (what, who, when, etc.) if possible
4. Do not add or remove information
5. Only output the paraphrased question, nothing else
Original question: {query}
Paraphrased question:"""
},
{
"name": "rephrase_formal",
"instruction": """Please rephrase the following question in a more formal academic style while preserving its meaning.
Rules:
1. Use more formal vocabulary
2. Maintain the exact same meaning
3. Change sentence structure appropriately
4. Only output the rephrased question
Question: {query}
Formal version:"""
},
{
"name": "rephrase_conversational",
"instruction": """Please rewrite the following question in a conversational, natural language style while keeping the same meaning.
Rules:
1. Use everyday language
2. Make it sound natural and conversational
3. Preserve the core question
4. Only output the conversational version
Question: {query}
Conversational version:"""
},
{
"name": "restructure",
"instruction": """Please restructure the following question using different syntax while maintaining identical meaning.
Rules:
1. Change word order or sentence structure
2. Use synonyms where appropriate
3. Keep the same information content
4. Only output the restructured question
Original: {query}
Restructured:"""
},
{
"name": "simplify",
"instruction": """Please simplify the following question using simpler words while keeping the exact same meaning.
Rules:
1. Use simpler, more common words
2. Maintain the same question intent
3. Don't lose any information
4. Only output the simplified question
Question: {query}
Simplified:"""
},
]
return prompts
def generate_variants(
self,
query: str,
use_cache: bool = True
) -> List[str]:
"""
生成query的变体
Args:
query: 原始查询
use_cache: 是否使用缓存
Returns:
query变体列表 (不包含原始query)
"""
variants = []
prompts = self._create_rewrite_prompts()
selected_prompts = prompts[:self.num_variants]
for prompt_config in selected_prompts:
variant_type = prompt_config["name"]
cache_key = self._get_cache_key(query, variant_type)
if use_cache and cache_key in self.cache:
variant = self.cache[cache_key]
print(f"✓ 从缓存加载: {variant_type}")
else:
prompt = prompt_config["instruction"].format(query=query)
try:
variant = self.api._generate(prompt, max_out_len=256)
variant = variant.strip()
prefixes_to_remove = [
"Paraphrased question:",
"Formal version:",
"Conversational version:",
"Restructured:",
"Simplified:",
]
for prefix in prefixes_to_remove:
if variant.startswith(prefix):
variant = variant[len(prefix):].strip()
self.cache[cache_key] = variant
self._save_cache()
print(f"✓ API生成: {variant_type}")
time.sleep(1)
except Exception as e:
print(f"✗ 生成失败 ({variant_type}): {e}")
variant = query # 失败时使用原始query
if variant and variant != query:
variants.append(variant)
return variants
def process_dataset(
self,
input_path: str,
output_path: str,
max_samples: Optional[int] = None,
):
"""
处理整个数据集
Args:
input_path: 输入数据文件路径
output_path: 输出文件路径
max_samples: 最大处理样本数
"""
print(f"\n{'='*60}")
print(f"开始处理数据集: {input_path}")
print(f"生成 {self.num_variants} 个query变体")
print(f"{'='*60}\n")
with open(input_path, 'r', encoding='utf-8') as f:
if input_path.endswith('.jsonl'):
data = [json.loads(line) for line in f]
else:
data = json.load(f)
if max_samples:
data = data[:max_samples]
augmented_data = []
for idx, item in enumerate(tqdm(data, desc="生成query variants")):
query = item.get('question', item.get('query', ''))
if not query:
print(f"警告: 样本 {idx} 缺少query,跳过")
continue
print(f"\n[{idx+1}/{len(data)}] 原始Query: {query[:100]}...")
variants = self.generate_variants(query, use_cache=True)
print(f" 生成了 {len(variants)} 个变体:")
for i, var in enumerate(variants, 1):
print(f" {i}. {var[:80]}...")
augmented_item = {
**item,
'query_original': query,
'query_variants': variants,
'num_variants': len(variants),
}
augmented_data.append(augmented_item)
output_dir = os.path.dirname(output_path)
if output_dir:
os.makedirs(output_dir, exist_ok=True)
with open(output_path, 'w', encoding='utf-8') as f:
json.dump(augmented_data, f, ensure_ascii=False, indent=2)
print(f"\n{'='*60}")
print(f"✓ 处理完成!")
print(f" 输入: {len(data)} 个样本")
print(f" 输出: {len(augmented_data)} 个样本")
print(f" 保存到: {output_path}")
print(f" 缓存: {len(self.cache)} 个entries")
print(f"{'='*60}\n")
return augmented_data
def main():
parser = argparse.ArgumentParser(
description="使用DeepSeek API生成Query Variants"
)
parser.add_argument(
"--api_key",
type=str,
required=True,
help="DeepSeek API Key"
)
parser.add_argument(
"--api_url",
type=str,
default="https://api.deepseek.com/v1/chat/completions",
help="DeepSeek API URL"
)
parser.add_argument(
"--model",
type=str,
default="deepseek-chat",
help="模型名称"
)
parser.add_argument(
"--input",
type=str,
required=True,
help="输入数据文件 (JSON/JSONL)"
)
parser.add_argument(
"--output",
type=str,
required=True,
help="输出文件路径"
)
parser.add_argument(
"--max_samples",
type=int,
default=None,
help="最大处理样本数"
)
parser.add_argument(
"--num_variants",
type=int,
default=3,
help="每个query生成的变体数量"
)
parser.add_argument(
"--cache_dir",
type=str,
default="./cache/query_variants",
help="缓存目录"
)
args = parser.parse_args()
generator = QueryVariantGenerator(
api_key=args.api_key,
api_url=args.api_url,
model=args.model,
cache_dir=args.cache_dir,
num_variants=args.num_variants,
)
generator.process_dataset(
input_path=args.input,
output_path=args.output,
max_samples=args.max_samples,
)
if __name__ == "__main__":
main()