File size: 8,835 Bytes
34cc882 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 | """
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()
|