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()