File size: 11,705 Bytes
d04fea8
 
 
 
 
9b4263f
 
d04fea8
 
 
 
 
 
9b4263f
 
 
d04fea8
 
 
 
 
 
 
 
 
 
 
 
9b4263f
 
 
 
 
d04fea8
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
9b4263f
d04fea8
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
#!/usr/bin/env python3
"""
关键词归因评测:只打 /api/analyze-semantic-keywords(默认本地)。

真值约定:
  - 磁盘上每 chunk 的 query 为数组,每项自带 expect_relevant / expect_keywords
  - 加载后展平;仅 expect_relevant=true 的用例参与本脚本;无关例跳过(expect_keywords 无意义)
  - 相关例对照 expect_keywords 与 top scored raw(粗匹配)

相关性门控请用 scripts/eval_semantic_relevance_remote.py(云端 Chat)。

用法(项目根目录):
  python scripts/eval_semantic_keywords.py \\
    -c scripts/cases/红楼-第3回.json \\
    -o scripts/results/红楼-第3回_kw.jsonl \\
    --review-md scripts/results/红楼-第3回_kw_review.md
"""

from __future__ import annotations

import argparse
import json
import os
import sys
import time
from pathlib import Path
from typing import Any, Dict, List, Optional

_SCRIPTS_DIR = Path(__file__).resolve().parent
if str(_SCRIPTS_DIR) not in sys.path:
    sys.path.insert(0, str(_SCRIPTS_DIR))
from semantic_case_load import load_relevant_cases

HF_TOKEN_ENV = "HF_TOKEN"

try:
    import requests
except ImportError:
    print("错误: 需要安装 requests 库")
    print("请运行: pip install requests")
    sys.exit(1)

DEFAULT_API_BASE = "http://localhost:5001"
PATH_KEYWORDS = "/api/analyze-semantic-keywords"


def analyze_keywords_http(
    api_base: str,
    query: str,
    text: str,
    token: Optional[str] = None,
    timeout: int = 300,
) -> dict:
    url = f"{api_base.rstrip('/')}{PATH_KEYWORDS}"
    headers = {"Content-Type": "application/json"}
    if token:
        headers["Authorization"] = f"Bearer {token}"
    resp = requests.post(
        url,
        json={"query": query, "text": text, "debug_info": True},
        headers=headers,
        timeout=timeout,
    )
    resp.raise_for_status()
    data = resp.json()
    if not data.get("success"):
        raise RuntimeError(data.get("message", "分析失败"))
    return data


def _load_jsonl(path: Path) -> list:
    if not path.exists():
        return []
    results = []
    for line in path.read_text(encoding="utf-8").strip().split("\n"):
        if not line:
            continue
        try:
            results.append(json.loads(line))
        except json.JSONDecodeError:
            pass
    return results


def _append_record(path: Path, record: dict) -> None:
    with path.open("a", encoding="utf-8") as f:
        f.write(json.dumps(record, ensure_ascii=False) + "\n")


def _top_scored_from_attention(token_attention: list, k: int = 10) -> List[dict]:
    if not token_attention:
        return []
    score_max = max(a["score"] for a in token_attention)
    denom = score_max if score_max > 0 else 1
    sorted_attn = sorted(token_attention, key=lambda x: x["score"], reverse=True)[:k]
    out = []
    for a in sorted_attn:
        out.append({
            "raw": a.get("raw"),
            "score": round(a["score"], 6),
            "score_norm": round(a["score"] / denom, 6),
            "offset": a.get("offset"),
        })
    return out


def _keyword_hits(expect_keywords: List[str], top_scored: List[dict]) -> List[str]:
    if not expect_keywords:
        return []
    tops = [str(x.get("raw") or "") for x in (top_scored or [])[:10]]
    blob = " ".join(tops).lower()
    hits = []
    for kw in expect_keywords:
        kl = kw.lower()
        if kl in blob or any(kl in t.lower() or t.lower() in kl for t in tops if t.strip()):
            hits.append(kw)
            continue
        parts = [
            p for p in kw.replace("(", " ").replace(")", " ").replace("(", " ").replace(")", " ").split()
            if len(p) >= 2
        ]
        if parts and any(p.lower() in blob for p in parts):
            hits.append(kw)
    return hits


def _top5_raw(top_scored: List[dict]) -> str:
    if not top_scored:
        return "(无)"
    return ", ".join(repr(x.get("raw", "")) for x in top_scored[:5])


def load_cases(path: Path) -> tuple[List[dict], int]:
    """返回 (相关用例, 跳过的无关条数)。无关例的 expect_keywords 忽略。"""
    return load_relevant_cases(path)


def run_one(
    api_base: str,
    case: dict,
    token: Optional[str],
    timeout: int,
    max_retries: int,
) -> dict:
    name = case["name"]
    query = case["query"]
    text = case["text"]
    expect_kw = case.get("expect_keywords") or []
    disputed = bool(case.get("disputed"))
    dispute_note = case.get("dispute_note") or ""

    def _base(**extra: Any) -> dict:
        rec: Dict[str, Any] = {
            "case": name,
            "chunk_index": case.get("chunk_index"),
            "query": query,
            "expect_relevant": True,
            "expect_keywords": expect_kw,
            "source": case.get("source"),
            **extra,
        }
        if disputed:
            rec["disputed"] = True
            if dispute_note:
                rec["dispute_note"] = dispute_note
        return rec

    last_error: Optional[BaseException] = None
    r2 = None
    for attempt in range(max_retries + 1):
        try:
            r2 = analyze_keywords_http(api_base, query, text, token=token, timeout=timeout)
            break
        except Exception as e:
            last_error = e
            if attempt < max_retries:
                wait = 3 * (attempt + 1)
                print(f"  keywords 重试 {attempt + 1}/{max_retries}{wait}s… {e}", flush=True)
                time.sleep(wait)
    if r2 is None:
        return _base(error=f"keywords: {last_error}")

    top = _top_scored_from_attention(r2.get("token_attention") or [])
    hits = _keyword_hits(expect_kw, top)
    return _base(
        model=r2.get("model", ""),
        top10_scored_raw=top,
        keyword_hits=hits,
        keywords_ok=(not expect_kw) or bool(hits),
    )


def enrich_results_from_cases(results: List[dict], cases: List[dict]) -> None:
    by_name = {c["name"]: c for c in cases}
    for r in results:
        c = by_name.get(r.get("case") or "")
        if not c:
            continue
        if c.get("disputed"):
            r["disputed"] = True
            if c.get("dispute_note"):
                r["dispute_note"] = c["dispute_note"]
        else:
            r.pop("disputed", None)
            r.pop("dispute_note", None)


def write_review_markdown(results: List[dict], path: Path) -> None:
    lines = [
        "# 关键词归因对照表(仅相关例)",
        "",
        "只评 `/api/analyze-semantic-keywords`;无关例不跑本主题。",
        "`disputed=true`:边界争议,汇总仍按 expect 计;报告里点评实测。",
        "",
        "| case | chunk | disputed | hits | verdict |",
        "|---|---:|---|---|---|",
    ]
    ok = miss = err = 0
    disputed_rows: List[dict] = []

    for r in results:
        name = r.get("case", "?")
        ci = r.get("chunk_index", "")
        disp = "yes" if r.get("disputed") else ""
        if r.get("disputed"):
            disputed_rows.append(r)
        if r.get("error"):
            lines.append(f"| {name} | {ci} | {disp} | — | **error**: {r['error']} |")
            err += 1
            continue
        expect_kw = r.get("expect_keywords") or []
        hits = r.get("keyword_hits") or []
        top = _top5_raw(r.get("top10_scored_raw") or [])
        if expect_kw and not hits:
            lines.append(
                f"| {name} | {ci} | {disp} | [] top5={top} | **词未命中** expect={expect_kw} |"
            )
            miss += 1
        else:
            lines.append(
                f"| {name} | {ci} | {disp} | "
                f"hits={hits or '(无 expect_kw)'} top5={top} | **OK** |"
            )
            ok += 1

    lines.extend([
        "",
        "# 汇总",
        "",
        "| 词命中OK | 词未命中 | error |",
        "|---:|---:|---:|",
        f"| {ok} | {miss} | {err} |",
        "",
    ])
    if disputed_rows:
        lines.extend([
            "# 争议 case(disputed)",
            "",
            "| case | note | hits |",
            "|---|---|---|",
        ])
        for r in disputed_rows:
            note = (r.get("dispute_note") or "").replace("|", "\\|")
            hits = r.get("keyword_hits", "—")
            lines.append(f"| {r.get('case')} | {note} | {hits} |")
        lines.append("")
    path.parent.mkdir(parents=True, exist_ok=True)
    path.write_text("\n".join(lines), encoding="utf-8")
    print(f"✅ 对照表已写入 {path}")


def main() -> None:
    parser = argparse.ArgumentParser(
        description="关键词归因评测(仅 expect_relevant=true;打本地 keywords API)"
    )
    parser.add_argument("-c", "--cases", type=Path, required=True, help="用例 JSON 数组")
    parser.add_argument("-o", "--output", type=Path, default=None, help="结果 JSONL(可续跑)")
    parser.add_argument("--review-md", type=Path, default=None, help="对照表 Markdown")
    parser.add_argument("--review-only", action="store_true", help="仅从 JSONL 生成对照表")
    parser.add_argument("--url", default=DEFAULT_API_BASE, help=f"keywords API 根,默认 {DEFAULT_API_BASE}")
    parser.add_argument("--hf-token", default=None)
    parser.add_argument("--retries", type=int, default=3)
    parser.add_argument("--timeout", type=int, default=300)
    args = parser.parse_args()

    if args.review_only:
        if not args.output or not args.review_md:
            print("错误: --review-only 需要 -o 与 --review-md")
            sys.exit(1)
        results = _load_jsonl(args.output)
        if args.cases:
            cases, _ = load_cases(args.cases)
            enrich_results_from_cases(results, cases)
        write_review_markdown(results, args.review_md)
        return

    cases, skipped = load_cases(args.cases)
    print(f"已加载 {len(cases)} 个相关用例(跳过无关 {skipped})")
    if not cases:
        print("无相关用例可跑")
        sys.exit(0)

    api_base = args.url.rstrip("/")
    token = args.hf_token or os.environ.get(HF_TOKEN_ENV)

    completed = set()
    all_results: list = []
    if args.output and args.output.exists():
        all_results = _load_jsonl(args.output)
        completed = {r["case"] for r in all_results if "case" in r}
        print(f"已加载 {len(all_results)} 条历史,跳过 {len(completed)} 个 case")

    for i, case in enumerate(cases):
        name = case["name"]
        prog = f"[{i + 1}/{len(cases)}]"
        if name in completed:
            print(f"{prog}{name}", flush=True)
            continue
        print(f"{prog} 执行 {name}", flush=True)
        record = run_one(api_base, case, token, args.timeout, args.retries)
        if record.get("error"):
            print(f"{prog}{name}: {record['error']}", flush=True)
            all_results.append(record)
            if args.output:
                args.output.parent.mkdir(parents=True, exist_ok=True)
                _append_record(args.output, record)
            print("⚠ 失败中断后续", flush=True)
            break
        print(
            f"{prog}{name} hits={record['keyword_hits']} "
            f"ok={record['keywords_ok']}",
            flush=True,
        )
        all_results.append(record)
        completed.add(name)
        if args.output:
            args.output.parent.mkdir(parents=True, exist_ok=True)
            _append_record(args.output, record)

    if args.output:
        print(f"\n✅ 结果已写入 {args.output}(共 {len(all_results)} 条)")
    if args.review_md and all_results:
        enrich_results_from_cases(all_results, cases)
        write_review_markdown(all_results, args.review_md)


if __name__ == "__main__":
    main()