query_filter.py 3.0 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960
  1. """创作 query 筛选器:调用 acquisition/query_filter.txt,批量判一组 query。
  2. 模型对每条 query 打两项:语义合法性 valid(0-10) + 创作相关性 relevant(bool)。
  3. 最终保留 keep = valid ≥ VALID_MIN 且 relevant —— 阈值放代码里(好调,不改提示词)。
  4. · valid:词组合本身说不说得通(机械正交易产生的废话短语在此淘汰)。
  5. · relevant:搜回来是否和创作知识相关(三把尺=设计决策vs工艺/可迁移/业务阶段;再排 制作/题材本身/应试政企/作品)。
  6. 返回与输入等长的 [{keep, valid, relevant, reason}]。供 build_creation_demo、filter_multiaxis 等共用。
  7. """
  8. from __future__ import annotations
  9. import json
  10. from pathlib import Path
  11. import httpx
  12. from core.config import Settings
  13. PROMPT = Path(__file__).resolve().parent / "query_filter.txt"
  14. VALID_MIN = 6 # 语义合法性阈值:valid ≥ 此值才算说得通;调这里即可松紧
  15. def filter_queries(queries: list[str], settings: Settings, *, batch: int = 40) -> list[dict]:
  16. """批量过滤(分批调用,避免一次喂太多)。返回 [{keep, reason}],与 queries 对齐。"""
  17. out: list[dict] = []
  18. for i in range(0, len(queries), batch):
  19. out.extend(_filter_batch(queries[i:i + batch], settings))
  20. return out
  21. def _filter_batch(queries: list[str], settings: Settings) -> list[dict]:
  22. if not queries:
  23. return []
  24. user = json.dumps([{"idx": i, "query": q} for i, q in enumerate(queries)], ensure_ascii=False)
  25. api = settings.openrouter_base_url.rstrip("/") + "/chat/completions"
  26. headers = {"Authorization": f"Bearer {settings.openrouter_api_key}", "Content-Type": "application/json"}
  27. body = {"model": settings.llm_model, "messages": [
  28. {"role": "system", "content": PROMPT.read_text("utf-8")},
  29. {"role": "user", "content": user}], "response_format": {"type": "json_object"}}
  30. try:
  31. resp = httpx.post(api, headers=headers, json=body, timeout=120)
  32. resp.raise_for_status()
  33. txt = resp.json()["choices"][0]["message"]["content"]
  34. # query_filter 要求输出数组;有的模型会包一层 {"result":[...]},都兜住
  35. data = json.loads(txt)
  36. arr = data if isinstance(data, list) else next((v for v in data.values() if isinstance(v, list)), [])
  37. by = {d.get("idx"): d for d in arr if isinstance(d, dict)}
  38. out = []
  39. for i in range(len(queries)):
  40. d = by.get(i, {})
  41. try:
  42. valid = int(d.get("valid", 10))
  43. except (TypeError, ValueError):
  44. valid = 10
  45. relevant = bool(d.get("relevant", True))
  46. out.append({"keep": valid >= VALID_MIN and relevant, # 阈值在代码里判
  47. "valid": valid, "relevant": relevant,
  48. "reason": str(d.get("reason", ""))[:50]})
  49. return out
  50. except Exception as exc:
  51. return [{"keep": True, "valid": 10, "relevant": True, "reason": f"筛选失败:{str(exc)[:30]}"} for _ in queries]