knowledge_graph_agent.py 5.8 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153
  1. """知识图谱智能体 —— 从已保存论文中抽取主题/方法/引用关系并构建可视化图谱."""
  2. from __future__ import annotations
  3. import json
  4. import logging
  5. from typing import Any
  6. from hello_agents import SimpleAgent
  7. from ..utils import parse_llm_json
  8. from ..services.llm.agent_config import papergraph_agent_config
  9. from .base import BaseAgent
  10. from .prompts.knowledge_graph import REL_PROMPT
  11. logger = logging.getLogger(__name__)
  12. class KnowledgeGraphAgent(BaseAgent):
  13. """知识图谱构建智能体 —— 从论文集中抽取关系、生成节点与边数据."""
  14. def __init__(
  15. self,
  16. *,
  17. agent: SimpleAgent | None = None,
  18. min_score: float = 0.55,
  19. max_edges: int = 12,
  20. chunk_size: int = 20,
  21. ) -> None:
  22. super().__init__()
  23. self.min_score = min_score
  24. self.max_edges = max_edges
  25. self.chunk_size = chunk_size
  26. self._agent = agent or SimpleAgent(
  27. name="papergraph_kg_rel",
  28. llm=self.llm,
  29. system_prompt=REL_PROMPT,
  30. config=papergraph_agent_config(),
  31. )
  32. def _candidate_id(self, paper: dict[str, Any]) -> int | None:
  33. for key in ("paper_id", "id", "target_paper_id"):
  34. try:
  35. value = int(paper.get(key))
  36. if value > 0:
  37. return value
  38. except (TypeError, ValueError):
  39. continue
  40. return None
  41. def _compact_paper(self, paper: dict[str, Any]) -> dict[str, Any]:
  42. out = {
  43. "paper_id": self._candidate_id(paper),
  44. "title": self._clip(paper.get("title"), 300),
  45. "abstract": self._clip(paper.get("abstract"), 2000),
  46. "keywords": list((paper.get("keywords") or [])[:12]),
  47. "source": paper.get("source"),
  48. "year": paper.get("year"),
  49. "category": paper.get("category"),
  50. "pdf_excerpt": self._clip(paper.get("pdf_excerpt"), 1200),
  51. "related_work_excerpt": self._clip(paper.get("related_work_excerpt"), 1200),
  52. }
  53. return {k: v for k, v in out.items() if v not in (None, "", [], {})}
  54. def _validate_edges(self, edges: Any, allowed_ids: set[int]) -> list[dict[str, Any]]:
  55. if not isinstance(edges, list):
  56. raise ValueError("edges is not a list")
  57. best: dict[int, dict[str, Any]] = {}
  58. for e in edges:
  59. if not isinstance(e, dict):
  60. continue
  61. try:
  62. tid = int(e.get("target_paper_id"))
  63. score = float(e.get("score") or 0.0)
  64. except Exception:
  65. continue
  66. if tid not in allowed_ids or score < self.min_score:
  67. continue
  68. relation = str(e.get("relation") or "").strip()
  69. if not relation:
  70. continue
  71. edge = {"target_paper_id": tid, "relation": relation,
  72. "score": max(0.0, min(1.0, score)),
  73. "evidence": str(e.get("evidence") or "").strip()[:80]}
  74. if tid not in best or score > best[tid]["score"]:
  75. best[tid] = edge
  76. return sorted(best.values(), key=lambda x: x["score"], reverse=True)
  77. def _chunks(self, items: list[dict[str, Any]]) -> list[list[dict[str, Any]]]:
  78. return [items[i : i + self.chunk_size] for i in range(0, len(items), self.chunk_size)]
  79. def infer_edges(
  80. self, *, new_paper: dict[str, Any], candidates: list[dict[str, Any]]
  81. ) -> tuple[list[dict[str, Any]], str | None]:
  82. compact_new = self._compact_paper(new_paper)
  83. compact_candidates: list[dict[str, Any]] = []
  84. for c in candidates:
  85. tid = self._candidate_id(c)
  86. if tid is None:
  87. continue
  88. item = self._compact_paper(c)
  89. item["paper_id"] = tid
  90. compact_candidates.append(item)
  91. if not compact_candidates:
  92. return [], None
  93. merged: dict[int, dict[str, Any]] = {}
  94. for chunk in self._chunks(compact_candidates):
  95. payload = {"new_paper": compact_new, "candidates": chunk}
  96. try:
  97. raw = self._agent.run(json.dumps(payload, ensure_ascii=False))
  98. except Exception as exc:
  99. logger.exception("kg_llm_run_failed")
  100. raise RuntimeError("kg_llm_run_failed") from exc
  101. data = parse_llm_json(raw)
  102. if data is None:
  103. raise ValueError("kg_llm_parse_failed")
  104. try:
  105. valid_edges = self._validate_edges(
  106. data.get("edges"),
  107. allowed_ids={int(x["paper_id"]) for x in chunk},
  108. )
  109. except Exception as exc:
  110. raise ValueError("kg_edge_validation_failed") from exc
  111. for edge in valid_edges:
  112. tid = edge["target_paper_id"]
  113. if tid not in merged or edge["score"] > merged[tid]["score"]:
  114. merged[tid] = edge
  115. edges = sorted(merged.values(), key=lambda x: x["score"], reverse=True)[: self.max_edges]
  116. try:
  117. from ..services.memory.agent_memory import get_agent_memory
  118. am = get_agent_memory()
  119. title = str(new_paper.get("title") or "")[:120]
  120. am.add(agent_name="knowledge_graph", content=f"关系抽取:{title}", memory_type="working", importance=0.45, shared=False)
  121. if edges:
  122. am.add(
  123. agent_name="knowledge_graph",
  124. content=f"关系抽取要点:top_relation={edges[0].get('relation')} score={edges[0].get('score')}",
  125. memory_type="working",
  126. importance=0.5,
  127. shared=True,
  128. )
  129. except Exception:
  130. logger.debug("kg_memory_write_failed", exc_info=True)
  131. return edges, None