| """ |
| TFMF LoRA 评估脚本 - 提示词注入 + Tavily 联网 |
| |
| 本脚本会加载基座模型和 LoRA 适配器, |
| 根据输入问题进行推理。 |
| 如果模型输出中包含 <search> ... </search>, |
| 则自动调用 Tavily 搜索并把结果注入到下一轮生成中。 |
| """ |
|
|
| import argparse |
| import json |
| import os |
| import re |
| import time |
| import torch |
| from tavily import TavilyClient |
| from peft import PeftModel |
| from tavily import TavilyClient |
| from transformers import AutoModelForCausalLM, AutoTokenizer |
|
|
| DEFAULT_MODEL_PATH = "/data/coding/TFMF" |
| DEFAULT_ADAPTER_PATH = "./TFMF/lora_adapters/teacher_chinese_auto/final" |
| DEFAULT_DATA_PATH = "./TFMF/dataset/语文教师_语文_高二_v1.jsonl" |
|
|
| SYSTEM_PROMPT = '''你是一名能够联网搜索的高中语文教师助手。 |
| 你在回答时要保持: |
| - 直白浅近,善于用生活化例子讲清楚抽象概念; |
| - 喜欢通过师生问答、追问和归纳总结来引导学生; |
| - 不要编造事实,仅当你真的需要最新信息时才使用网络搜索; |
| - 注意用户问题中的“最新”“最近”“现在”“目前”以及近义词往往需要联网搜索 |
| - 如果需要搜索,请输出 <search>你的搜索查询</search>; |
| - 收到搜索结果后,请用 <observation> ... </observation> 包裹结果,并基于结果继续回答。 |
| |
| 回答中不得包含搜索语句本身,最终回答只保留正常回答内容。 |
| ''' |
|
|
| SEARCH_PATTERN = re.compile(r"<search>(.*?)</search>", re.DOTALL | re.IGNORECASE) |
|
|
|
|
| def parse_args(): |
| parser = argparse.ArgumentParser(description="TFMF LoRA 评估脚本(提示词注入 + Tavily 联网)") |
| parser.add_argument("--model-path", type=str, default=DEFAULT_MODEL_PATH, |
| help="基座模型目录,例如 /data/coding/TFMF") |
| parser.add_argument("--adapter-path", type=str, default=DEFAULT_ADAPTER_PATH, |
| help="LoRA 适配器 final 目录") |
| parser.add_argument("--tavily-api-key", type=str, default=None, |
| help="Tavily API Key,若不传则读取环境变量 TAVILY_API_KEY") |
| parser.add_argument("--data-path", type=str, default=DEFAULT_DATA_PATH, |
| help="测试数据文件路径,可选") |
| parser.add_argument("--n-samples", type=int, default=3, |
| help="评估时读取的测试样本数量") |
| parser.add_argument("--interactive", action="store_true", |
| help="进入交互模式,手动输入问题") |
| parser.add_argument("--max-new-tokens", type=int, default=512, |
| help="生成最大新 token 数") |
| parser.add_argument("--temperature", type=float, default=0.7, |
| help="采样温度") |
| parser.add_argument("--top-p", type=float, default=0.9, |
| help="top-p 采样") |
| parser.add_argument("--max-search", type=int, default=2, |
| help="最多进行几轮网络搜索") |
| return parser.parse_args() |
|
|
|
|
| def load_model(model_path: str, adapter_path: str): |
| print("加载 tokenizer 和基座模型...") |
| tokenizer = AutoTokenizer.from_pretrained(model_path, trust_remote_code=True) |
| if tokenizer.pad_token is None: |
| tokenizer.pad_token = tokenizer.eos_token |
| tokenizer.padding_side = "right" |
|
|
| model = AutoModelForCausalLM.from_pretrained( |
| model_path, |
| device_map="auto", |
| torch_dtype=torch.bfloat16, |
| trust_remote_code=True, |
| low_cpu_mem_usage=True, |
| ) |
| model.config.use_cache = False |
|
|
| print("挂载 LoRA 适配器...") |
| model = PeftModel.from_pretrained(model, adapter_path) |
| model.eval() |
| return tokenizer, model |
|
|
|
|
| def search_web(tavily: TavilyClient, query: str) -> str: |
| try: |
| result = tavily.search( |
| query=query, |
| search_depth="basic", |
| max_results=5, |
| include_answer=False, |
| include_raw_content=True, |
| ) |
| hits = result.get("results", []) |
| if not hits: |
| return "没有搜索到结果。" |
|
|
| texts = [] |
| for idx, item in enumerate(hits, start=1): |
| title = item.get("title", "") |
| url = item.get("url", "") |
| content = item.get("raw_content") or item.get("content", "") |
| if len(content) > 400: |
| content = content[:400] + "..." |
| texts.append(f"[{idx}] {title}\n{url}\n{content}") |
| return "\n\n".join(texts) |
| except Exception as exc: |
| return f"搜索失败: {exc}" |
|
|
|
|
| def extract_search(text: str): |
| match = SEARCH_PATTERN.search(text) |
| return match.group(1).strip() if match else None |
|
|
|
|
| def remove_search_tags(text: str) -> str: |
| return SEARCH_PATTERN.sub("", text).strip() |
|
|
|
|
| def generate_reply(tokenizer, model, prompt: str, max_new_tokens: int, temperature: float, top_p: float) -> str: |
| inputs = tokenizer(prompt, return_tensors="pt").to(model.device) |
| with torch.no_grad(): |
| outputs = model.generate( |
| **inputs, |
| max_new_tokens=max_new_tokens, |
| temperature=temperature, |
| top_p=top_p, |
| do_sample=True, |
| pad_token_id=tokenizer.eos_token_id, |
| eos_token_id=tokenizer.eos_token_id, |
| repetition_penalty=1.1, |
| ) |
| answer = tokenizer.decode(outputs[0][inputs["input_ids"].shape[1]:], skip_special_tokens=True) |
| return answer.strip() |
|
|
|
|
| def build_prompt(system: str, user: str): |
| messages = [ |
| {"role": "system", "content": SYSTEM_PROMPT}, |
| {"role": "user", "content": user}, |
| ] |
| return tokenizer.apply_chat_template(messages, tokenize=False, add_generation_prompt=True, enable_thinking=False) |
|
|
|
|
| def eval_sample(sample: dict, tokenizer, model, tavily, args): |
| system = sample.get("system", "") |
| user = sample.get("user", "") |
| prompt = tokenizer.apply_chat_template( |
| [{"role": "system", "content": SYSTEM_PROMPT}, {"role": "user", "content": user}], |
| tokenize=False, |
| add_generation_prompt=True, |
| enable_thinking=False, |
| ) |
|
|
| response = generate_reply(tokenizer, model, prompt, args.max_new_tokens, args.temperature, args.top_p) |
| print("\n=== 初次生成 ===") |
| print(response) |
|
|
| for step in range(args.max_search): |
| query = extract_search(response) |
| if not query: |
| break |
| print(f"\n[搜索意图] {query}") |
| search_result = search_web(tavily, query) |
| print(f"\n[搜索结果]\n{search_result[:1200]}") |
|
|
| prompt = tokenizer.apply_chat_template( |
| [ |
| {"role": "system", "content": SYSTEM_PROMPT}, |
| {"role": "user", "content": user}, |
| {"role": "assistant", "content": response}, |
| {"role": "user", "content": f"<observation>\n{search_result}\n</observation>"}, |
| ], |
| tokenize=False, |
| add_generation_prompt=True, |
| enable_thinking=False, |
| ) |
| response = generate_reply(tokenizer, model, prompt, args.max_new_tokens, args.temperature, args.top_p) |
| print(f"\n=== 第 {step + 1} 轮搜索后生成 ===") |
| print(response) |
|
|
| response = remove_search_tags(response) |
| return response |
|
|
|
|
| def load_samples(data_path: str, limit: int = 3): |
| with open(data_path, "r", encoding="utf-8") as f: |
| samples = [json.loads(line) for line in f if line.strip()] |
| return samples[:limit] |
|
|
|
|
| def main(): |
| args = parse_args() |
| api_key = args.tavily_api_key or os.getenv("TAVILY_API_KEY") |
| if not api_key: |
| raise ValueError("请设置 Tavily API Key:export TAVILY_API_KEY='tvly-xxx' 或使用 --tavily-api-key 参数") |
|
|
| tokenizer, model = load_model(args.model_path, args.adapter_path) |
| tavily = TavilyClient(api_key=api_key) |
|
|
| print("模型加载完成。") |
| print(f"模型: {args.model_path}") |
| print(f"LoRA: {args.adapter_path}") |
|
|
| if args.interactive: |
| print("进入交互模式,输入空行退出。") |
| while True: |
| user_input = input("问题: ").strip() |
| if not user_input: |
| break |
| sample = {"system": SYSTEM_PROMPT, "user": user_input} |
| response = eval_sample(sample, tokenizer, model, tavily, args) |
| print(f"\n最终回答:\n{response}\n") |
| return |
|
|
| print(f"加载测试数据: {args.data_path}") |
| samples = load_samples(args.data_path, limit=args.n_samples) |
| for idx, sample in enumerate(samples, start=1): |
| print("\n" + "#" * 60) |
| print(f"测试样本 {idx}") |
| print(f"用户提问: {sample.get('user', '')}") |
| print(f"参考回答: {sample.get('assistant', '')[:200]}...") |
| response = eval_sample(sample, tokenizer, model, tavily, args) |
| print(f"\n最终回答:\n{response}") |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|