""" TFMF LoRA 评估脚本 - 提示词注入 + Tavily 联网 本脚本会加载基座模型和 LoRA 适配器, 根据输入问题进行推理。 如果模型输出中包含 ... , 则自动调用 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_PATTERN = re.compile(r"(.*?)", 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"\n{search_result}\n"}, ], 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()