pattern_dimension_analyze.py 36 KB

12345678910111213141516171819202122232425262728293031323334353637383940414243444546474849505152535455565758596061626364656667686970717273747576777879808182838485868788899091929394959697989910010110210310410510610710810911011111211311411511611711811912012112212312412512612712812913013113213313413513613713813914014114214314414514614714814915015115215315415515615715815916016116216316416516616716816917017117217317417517617717817918018118218318418518618718818919019119219319419519619719819920020120220320420520620720820921021121221321421521621721821922022122222322422522622722822923023123223323423523623723823924024124224324424524624724824925025125225325425525625725825926026126226326426526626726826927027127227327427527627727827928028128228328428528628728828929029129229329429529629729829930030130230330430530630730830931031131231331431531631731831932032132232332432532632732832933033133233333433533633733833934034134234334434534634734834935035135235335435535635735835936036136236336436536636736836937037137237337437537637737837938038138238338438538638738838939039139239339439539639739839940040140240340440540640740840941041141241341441541641741841942042142242342442542642742842943043143243343443543643743843944044144244344444544644744844945045145245345445545645745845946046146246346446546646746846947047147247347447547647747847948048148248348448548648748848949049149249349449549649749849950050150250350450550650750850951051151251351451551651751851952052152252352452552652752852953053153253353453553653753853954054154254354454554654754854955055155255355455555655755855956056156256356456556656756856957057157257357457557657757857958058158258358458558658758858959059159259359459559659759859960060160260360460560660760860961061161261361461561661761861962062162262362462562662762862963063163263363463563663763863964064164264364464564664764864965065165265365465565665765865966066166266366466566666766866967067167267367467567667767867968068168268368468568668768868969069169269369469569669769869970070170270370470570670770870971071171271371471571671771871972072172272372472572672772872973073173273373473573673773873974074174274374474574674774874975075175275375475575675775875976076176276376476576676776876977077177277377477577677777877978078178278378478578678778878979079179279379479579679779879980080180280380480580680780880981081181281381481581681781881982082182282382482582682782882983083183283383483583683783883984084184284384484584684784884985085185285385485585685785885986086186286386486586686786886987087187287387487587687787887988088188288388488588688788888989089189289389489589689789889990090190290390490590690790890991091191291391491591691791891992092192292392492592692792892993093193293393493593693793893994094194294394494594694794894995095195295395495595695795895996096196296396496596696796896997097197297397497597697797897998098198298398498598698798898999099199299399499599699799899910001001100210031004100510061007
  1. """
  2. Pattern 维度分析 Tool
  3. 功能概述:
  4. 1. 读取某次整体推导日志目录下各轮评估结果,累计 matched_post_point / derivation_output_point 等字段。
  5. 2. 每轮通过 derivation_output_point 在人设树中找到 cluster_level 层祖先节点(已推导维度节点集合)。
  6. 3. 从 deduped_patterns 中筛选包含已推导维度节点的 pattern,并对各元素标记是否已推导。
  7. 输入参数:
  8. - account_name: 账号名称
  9. - post_id: 帖子 ID
  10. - log_id: 推导日志目录名(形如 20260313210921)
  11. 已推导/未推导维度节点在结果中以对象列表表示,字段见 _analyze_single_round 返回说明。
  12. """
  13. import json
  14. import logging
  15. import sys
  16. from pathlib import Path
  17. from typing import Any, Dict, List, Optional, Tuple, Set
  18. logger = logging.getLogger(__name__)
  19. try:
  20. from agent.tools import tool, ToolResult, ToolContext
  21. except ImportError:
  22. def tool(*args, **kwargs):
  23. return lambda f: f
  24. ToolResult = None
  25. ToolContext = None
  26. # 保证直接运行或作为包加载时都能解析 utils / tools(IDE 可跳转)
  27. _root = Path(__file__).resolve().parent.parent
  28. if str(_root) not in sys.path:
  29. sys.path.insert(0, str(_root))
  30. from tools.find_tree_node import _load_trees # 加载三棵人设树
  31. _BASE_INPUT = Path(__file__).resolve().parent.parent / "input"
  32. _BASE_OUTPUT = Path(__file__).resolve().parent.parent / "output"
  33. # pattern 库 key 定义(与 find_pattern 中保持一致)
  34. TOP_KEYS = [
  35. "depth_4",
  36. ]
  37. SUB_KEYS = ["two_x", "one_x", "zero_x"]
  38. # 在人设树中查找祖先节点的目标深度(root 为 0 层)
  39. CLUSTER_LEVEL = 3
  40. # ---------------------------------------------------------------------------
  41. # 1. 读取推导日志:按轮次累计 matched_post_point
  42. # ---------------------------------------------------------------------------
  43. def _round_eval_dir(account_name: str, post_id: str, log_id: str) -> Path:
  44. """
  45. 推导日志目录:
  46. ../output/{account_name}/推导日志/{post_id}/{log_id}/
  47. """
  48. return _BASE_OUTPUT / account_name / "推导日志" / post_id / log_id
  49. def _load_round_matched_points(
  50. account_name: str,
  51. post_id: str,
  52. log_id: str,
  53. max_round: Optional[int] = None,
  54. ) -> List[Dict[str, Any]]:
  55. """
  56. 读取指定日志目录下所有 {轮次}.评估.json,按轮次排序,生成:
  57. [
  58. {
  59. "round": 1,
  60. "round_points": [
  61. {
  62. "matched_post_point": "叙事结构",
  63. "derivation_output_point": "叙事编排",
  64. "matched_score": 0.9151,
  65. "is_fully_derived": true,
  66. },
  67. ...
  68. ],
  69. "cumulative_points": [
  70. ... 累计到本轮的去重列表(以 derivation_output_point 为去重 key) ...
  71. ],
  72. },
  73. ...
  74. ]
  75. """
  76. base_dir = _round_eval_dir(account_name, post_id, log_id)
  77. if not base_dir.is_dir():
  78. return []
  79. eval_files: List[Tuple[int, Path]] = []
  80. for p in base_dir.glob("*.json"):
  81. name = p.name
  82. # 只处理 *_评估.json
  83. if not name.endswith("评估.json"):
  84. continue
  85. try:
  86. round_str = name.split("_", 1)[0]
  87. r = int(round_str)
  88. except Exception:
  89. continue
  90. eval_files.append((r, p))
  91. if max_round is not None:
  92. eval_files = [(r, p) for r, p in eval_files if r <= max_round]
  93. eval_files.sort(key=lambda x: x[0])
  94. results: List[Dict[str, Any]] = []
  95. cumulative: List[Dict[str, Any]] = []
  96. cumulative_set: Set[str] = set() # 以 derivation_output_point 去重
  97. for r, path in eval_files:
  98. try:
  99. with open(path, "r", encoding="utf-8") as f:
  100. data = json.load(f)
  101. except Exception:
  102. continue
  103. eval_results = data.get("eval_results") or []
  104. round_points: List[Dict[str, Any]] = []
  105. seen_in_round: Set[str] = set()
  106. for item in eval_results:
  107. if not isinstance(item, dict):
  108. continue
  109. if not item.get("is_matched"):
  110. continue
  111. dop = item.get("derivation_output_point")
  112. if dop is None:
  113. continue
  114. dop = str(dop).strip()
  115. if not dop:
  116. continue
  117. # 本轮内按 derivation_output_point 去重
  118. if dop in seen_in_round:
  119. continue
  120. seen_in_round.add(dop)
  121. mpp = item.get("matched_post_point")
  122. entry: Dict[str, Any] = {
  123. "matched_post_point": str(mpp).strip() if mpp is not None else None,
  124. "derivation_output_point": dop,
  125. "matched_score": item.get("matched_score"),
  126. "is_fully_derived": item.get("is_fully_derived"),
  127. }
  128. round_points.append(entry)
  129. # 累加到累计列表(按 derivation_output_point 去重)
  130. for entry in round_points:
  131. dop = entry["derivation_output_point"]
  132. if dop not in cumulative_set:
  133. cumulative_set.add(dop)
  134. cumulative.append(entry)
  135. results.append(
  136. {
  137. "round": r,
  138. "round_points": round_points,
  139. "cumulative_points": list(cumulative),
  140. }
  141. )
  142. return results
  143. # ---------------------------------------------------------------------------
  144. # 2. 读取 pattern 库并按 matched_post_point 打分
  145. # ---------------------------------------------------------------------------
  146. def _pattern_file(account_name: str) -> Path:
  147. """pattern 库文件:../input/{account_name}/原始数据/pattern/processed_edge_data.json"""
  148. return _BASE_INPUT / account_name / "原始数据" / "pattern" / "processed_edge_data.json"
  149. def _load_raw_patterns(account_name: str) -> List[Dict[str, Any]]:
  150. """
  151. 读取 pattern 库中所有原始 pattern(保留 items 结构,不做合并)。
  152. 返回列表中每个元素形如原始 JSON 中的 pattern(此处不关心 item 的 point / dimension 字段)。
  153. """
  154. path = _pattern_file(account_name)
  155. if not path.is_file():
  156. return []
  157. with open(path, "r", encoding="utf-8") as f:
  158. data = json.load(f)
  159. patterns: List[Dict[str, Any]] = []
  160. for top in TOP_KEYS:
  161. block = data.get(top)
  162. if not isinstance(block, dict):
  163. continue
  164. for sub in SUB_KEYS:
  165. items = block.get(sub) or []
  166. if isinstance(items, list):
  167. for p in items:
  168. if isinstance(p, dict):
  169. patterns.append(p)
  170. return patterns
  171. def _slim_pattern_for_dedupe(p: Dict[str, Any]) -> Tuple[float, List[str]]:
  172. """
  173. 提取 pattern 的 support 与去重后的 item name 列表(按名称合并,不关心顺序),
  174. 用于与 find_pattern.py 中的去重逻辑对齐。
  175. """
  176. items = p.get("items") or []
  177. names = [str(it.get("name") or "").strip() for it in items if isinstance(it, dict)]
  178. seen: Set[str] = set()
  179. unique: List[str] = []
  180. for n in names:
  181. if n and n not in seen:
  182. seen.add(n)
  183. unique.append(n)
  184. try:
  185. support = float(p.get("support", 0.0))
  186. except (TypeError, ValueError):
  187. support = 0.0
  188. return support, unique
  189. def _dedupe_patterns(raw_patterns: List[Dict[str, Any]]) -> List[Dict[str, Any]]:
  190. """
  191. 按 pattern 的 item name 集合去重(不区分顺序),与 find_pattern.py 的思路一致:
  192. - key 为 sorted(unique item names)
  193. - 同一个 key 仅保留 support 最大的 pattern(保留其原始 items 结构,方便后续打分)
  194. """
  195. key_to_best: Dict[Tuple[str, ...], Dict[str, Any]] = {}
  196. key_to_support: Dict[Tuple[str, ...], float] = {}
  197. for p in raw_patterns:
  198. support, unique = _slim_pattern_for_dedupe(p)
  199. if not unique:
  200. continue
  201. key = tuple(sorted(unique))
  202. best_support = key_to_support.get(key)
  203. if best_support is None or support > best_support:
  204. key_to_support[key] = support
  205. key_to_best[key] = p
  206. return list(key_to_best.values())
  207. # ---------------------------------------------------------------------------
  208. # 3. 人设树节点信息 & 聚类节点搜索
  209. # ---------------------------------------------------------------------------
  210. class TreeIndex:
  211. """
  212. 人设树索引:
  213. - node_info: 节点 -> { "parent": 父节点名称, "children": [子节点名称...], "depth": 深度, "dimension": 维度名 }
  214. - roots: 维度名 -> 根节点名称(即维度名本身)
  215. - merged_tree: 将实质/形式/意图三棵树合并后的单个 JSON(顶层 key 为实质/形式/意图)
  216. """
  217. def __init__(self, account_name: str) -> None:
  218. self.account_name = account_name
  219. self.node_info: Dict[str, Dict[str, Any]] = {}
  220. self.roots: Dict[str, str] = {}
  221. # 三棵树合并后的 JSON:{"实质": {...}, "形式": {...}, "意图": {...}}
  222. self.merged_tree: Dict[str, Dict[str, Any]] = {}
  223. self._build()
  224. def _build(self) -> None:
  225. trees = _load_trees(self.account_name)
  226. # 1)先将三棵树合并成一个 JSON:{"实质": {...}, "形式": {...}, "意图": {...}}
  227. merged: Dict[str, Dict[str, Any]] = {}
  228. for dim_name, root in trees:
  229. if isinstance(root, dict):
  230. merged[dim_name] = root
  231. self.merged_tree = merged
  232. # 2)基于合并后的 JSON 构建 parent/children 结构
  233. for dim_name, root in merged.items():
  234. root_name = dim_name
  235. self.roots[dim_name] = root_name
  236. if root_name not in self.node_info:
  237. self.node_info[root_name] = {
  238. "parent": None,
  239. "children": [],
  240. "dimension": dim_name,
  241. "depth": 0,
  242. }
  243. def walk(parent_name: str, node_dict: Dict[str, Any]):
  244. children = node_dict.get("children") or {}
  245. for name, child in children.items():
  246. if not isinstance(child, dict):
  247. continue
  248. if name not in self.node_info:
  249. self.node_info[name] = {
  250. "parent": parent_name,
  251. "children": [],
  252. "dimension": dim_name,
  253. "depth": None, # 稍后统一计算
  254. }
  255. else:
  256. # 仅当不会形成自引用时才更新 parent(树中可能存在同名的父子节点)
  257. if name != parent_name:
  258. self.node_info[name]["parent"] = parent_name
  259. self.node_info[name]["dimension"] = dim_name
  260. # 维护父节点的 children
  261. if parent_name not in self.node_info:
  262. self.node_info[parent_name] = {
  263. "parent": None,
  264. "children": [],
  265. "dimension": dim_name,
  266. "depth": 0,
  267. }
  268. if name not in self.node_info[parent_name]["children"]:
  269. self.node_info[parent_name]["children"].append(name)
  270. walk(name, child)
  271. walk(root_name, root)
  272. # 统一计算各节点深度(从根开始 BFS)
  273. from collections import deque
  274. q = deque()
  275. for dim_name, root_name in self.roots.items():
  276. if root_name not in self.node_info:
  277. continue
  278. self.node_info[root_name]["depth"] = 0
  279. q.append(root_name)
  280. while q:
  281. cur = q.popleft()
  282. cur_depth = self.node_info[cur].get("depth", 0) or 0
  283. for child in self.node_info[cur].get("children", []):
  284. self.node_info.setdefault(child, {})
  285. if self.node_info[child].get("depth") is None:
  286. self.node_info[child]["depth"] = cur_depth + 1
  287. q.append(child)
  288. def find_ancestor_at_level(self, node_name: str, level: int) -> Optional[str]:
  289. """
  290. 在人设树中找到 node_name 的 depth == level 的祖先节点。
  291. - 若 node_name 自身 depth == level,直接返回自身。
  292. - 若 node_name depth < level(比目标层浅),返回自身。
  293. - 否则沿 parent 链向上查找,返回第一个 depth == level 的祖先节点。
  294. """
  295. info = self.node_info.get(node_name)
  296. if not info:
  297. return None
  298. depth = info.get("depth")
  299. if depth is None:
  300. return None
  301. if depth <= level:
  302. return node_name
  303. cur = node_name
  304. visited: Set[str] = set()
  305. while cur and cur not in visited:
  306. visited.add(cur)
  307. cur_info = self.node_info.get(cur) or {}
  308. cur_depth = cur_info.get("depth") or 0
  309. if cur_depth == level:
  310. return cur
  311. if cur_depth < level:
  312. return cur
  313. parent = cur_info.get("parent")
  314. if parent is None:
  315. return cur
  316. cur = parent
  317. return None
  318. # 聚类搜索(不再区分维度)
  319. def find_clusters(
  320. self,
  321. elements: List[str],
  322. cluster_level: int,
  323. ) -> List[Dict[str, Any]]:
  324. """
  325. 在所有人设树中,为给定元素列表寻找聚类节点(不再要求 dimension 一致)。
  326. 规则(固定聚类层级 cluster_level):
  327. - 仅在 depth == cluster_level 的节点上做聚类判断:
  328. * 若某节点子树中包含的元素数量 >= 2,
  329. 且在该路径上尚未存在更高层(深度更小)的聚类节点,则将其视为一个聚类节点。
  330. - 对无法向上形成聚类的元素,为其寻找 depth == cluster_level 的祖先节点,
  331. 若存在则作为该元素的「单元素聚类」节点。
  332. - 返回:
  333. [
  334. {
  335. "cluster_node": "节点名",
  336. "from_elements": ["元素A", "元素B", ...]
  337. },
  338. ...
  339. ]
  340. """
  341. # 过滤出真实存在于人设树中的元素
  342. elem_set: Set[str] = set()
  343. for e in elements:
  344. e = str(e).strip()
  345. if not e:
  346. continue
  347. info = self.node_info.get(e)
  348. if not info:
  349. continue
  350. elem_set.add(e)
  351. if not elem_set:
  352. return []
  353. # 先计算每个节点子树中包含的元素数量(跨所有维度的根)
  354. # 注意:人设树数据中可能存在意外的环或重复引用,这里通过 visited 集合避免递归死循环。
  355. subtree_count: Dict[str, int] = {}
  356. def dfs_count(node: str, visited: Set[str]) -> int:
  357. if node in visited:
  358. # 检测到环,直接返回 0,避免无限递归
  359. return 0
  360. visited.add(node)
  361. cnt = 1 if node in elem_set else 0
  362. for ch in self.node_info.get(node, {}).get("children", []):
  363. cnt += dfs_count(ch, visited)
  364. subtree_count[node] = cnt
  365. return cnt
  366. for root_name in self.roots.values():
  367. dfs_count(root_name, set())
  368. # 再自上而下优先选择「更上层」聚类节点(但仅在 cluster_level 层):
  369. # - 若当前节点已作为聚类节点,则其子孙不再作为聚类节点(保证尽量向上聚类);
  370. # 同样需要防止意外的环导致递归过深,这里使用 visited 集合。
  371. clusters: Set[str] = set()
  372. def dfs_select(node: str, ancestor_selected: bool, visited: Set[str]) -> None:
  373. if node in visited:
  374. return
  375. visited.add(node)
  376. info = self.node_info.get(node) or {}
  377. depth = info.get("depth", 0) or 0
  378. cnt = subtree_count.get(node, 0)
  379. selected_here = False
  380. # 仅当祖先尚未被选中、当前节点位于 cluster_level 层且满足条件时,选当前节点为聚类节点
  381. if (not ancestor_selected) and depth == cluster_level and cnt >= 2:
  382. clusters.add(node)
  383. selected_here = True
  384. # 祖先已经被选中或当前节点被选中,则子孙不再作为聚类节点
  385. for ch in info.get("children", []):
  386. dfs_select(ch, ancestor_selected or selected_here, visited)
  387. for root_name in self.roots.values():
  388. dfs_select(root_name, False, set())
  389. if not clusters:
  390. return []
  391. # 统计每个聚类节点下真实覆盖的元素列表
  392. cluster_to_elements: Dict[str, Set[str]] = {c: set() for c in clusters}
  393. for e in elem_set:
  394. cur = e
  395. visited: Set[str] = set()
  396. while cur and cur not in visited:
  397. visited.add(cur)
  398. if cur in clusters:
  399. cluster_to_elements[cur].add(e)
  400. parent = self.node_info.get(cur, {}).get("parent")
  401. if parent is None:
  402. break
  403. cur = parent
  404. out: List[Dict[str, Any]] = []
  405. # 1)多元素聚类:仅统计真正输出的聚类节点所覆盖的元素,
  406. # 避免把「元素数不足 2 的节点」也算作已覆盖,从而导致元素丢失。
  407. covered_elems: Set[str] = set()
  408. for node in clusters:
  409. elems = sorted(cluster_to_elements.get(node) or [])
  410. if len(elems) < 2:
  411. # 主聚类逻辑只考虑覆盖至少 2 个元素的节点
  412. continue
  413. out.append(
  414. {
  415. "cluster_node": node,
  416. "from_elements": elems,
  417. }
  418. )
  419. for e in elems:
  420. covered_elems.add(e)
  421. # 2)对无法向上形成聚类的元素,给一个「单元素聚类」
  422. uncovered = elem_set - covered_elems
  423. # 将未覆盖元素按「cluster_level 层级的祖先节点」分组,确保同一个祖先节点下的
  424. # 多个元素合并为一个聚类,而不是多个单元素聚类。
  425. single_clusters: Dict[str, Set[str]] = {}
  426. for e in uncovered:
  427. # 单元素聚类时,cluster_node 应为「祖先节点」,不直接使用元素自身。
  428. # 这里固定选择 depth == cluster_level 的祖先节点。
  429. info_e = self.node_info.get(e) or {}
  430. parent = info_e.get("parent")
  431. cur = parent
  432. best_ancestor: Optional[str] = None
  433. visited_chain: Set[str] = set()
  434. while cur and cur not in visited_chain:
  435. visited_chain.add(cur)
  436. info = self.node_info.get(cur) or {}
  437. depth = info.get("depth", 0) or 0
  438. if depth == cluster_level:
  439. best_ancestor = cur
  440. break
  441. parent = info.get("parent")
  442. if parent is None:
  443. break
  444. cur = parent
  445. if best_ancestor:
  446. single_clusters.setdefault(best_ancestor, set()).add(e)
  447. for anc, elems in single_clusters.items():
  448. out.append(
  449. {
  450. "cluster_node": anc,
  451. "from_elements": sorted(elems),
  452. }
  453. )
  454. # 为了输出更稳定,按 from_elements 的元素数量从大到小排序,数量相同再按节点名排序
  455. out.sort(key=lambda x: (-len(x["from_elements"]), x["cluster_node"]))
  456. return out
  457. # ---------------------------------------------------------------------------
  458. # 4. 对单轮数据执行 pattern & 聚类分析
  459. # ---------------------------------------------------------------------------
  460. def _dim_obj(
  461. tree_node_name: str,
  462. tree_index: TreeIndex,
  463. matched_point: Optional[str] = None,
  464. ) -> Dict[str, Any]:
  465. dim = (tree_index.node_info.get(tree_node_name) or {}).get("dimension") or ""
  466. o: Dict[str, Any] = {
  467. "tree_node_name": tree_node_name,
  468. "dimension": dim,
  469. }
  470. if matched_point is not None:
  471. o["matched_point"] = matched_point
  472. return o
  473. def _entry_to_matched_point(entry: Dict[str, Any]) -> str:
  474. """is_fully_derived 为 true 时用 matched_post_point,否则用 derivation_output_point。"""
  475. dop = entry.get("derivation_output_point")
  476. dop_s = str(dop).strip() if dop is not None else ""
  477. if entry.get("is_fully_derived") is True:
  478. mpp = entry.get("matched_post_point")
  479. return str(mpp).strip() if mpp is not None else ""
  480. return dop_s
  481. def _analyze_single_round(
  482. patterns: List[Dict[str, Any]],
  483. tree_index: TreeIndex,
  484. cumulative_points: List[Dict[str, Any]],
  485. cluster_level: int = CLUSTER_LEVEL,
  486. ) -> Dict[str, Any]:
  487. """
  488. 对某一轮(给定累计 point 列表)执行维度分析:
  489. 1. 从 cumulative_points 中提取 derivation_output_point,
  490. 在人设树中找到每个节点的 cluster_level 层祖先 → derived_ancestor_set(已推导维度节点集合)。
  491. 2. 从 deduped_patterns 中筛选出包含 derived_ancestor_set 中节点的 pattern。
  492. 3. 对筛选出 pattern 的每个元素标记是否已推导:
  493. - 元素在 derived_ancestor_set 中 → is_derived=True(已推导维度)
  494. - 其他 → is_derived=False(未推导维度)
  495. 4. 汇总 derived_dims / underived_dims 对象列表。
  496. 返回结构(节选):
  497. - derived_ancestor_nodes: [{ tree_node_name, dimension, matched_point }, ...]
  498. - derived_dims: [{ tree_node_name, dimension, matched_point }, ...]
  499. - underived_dims: [{ tree_node_name, dimension }, ...](无 matched_point)
  500. """
  501. # 1. 收集 derived_ancestor_set,同时按规则累计每个祖先的 matched_point
  502. derived_ancestor_set: Set[str] = set()
  503. ancestor_to_matched: Dict[str, List[str]] = {}
  504. for entry in cumulative_points:
  505. if not isinstance(entry, dict):
  506. continue
  507. dop = entry.get("derivation_output_point")
  508. if not dop:
  509. continue
  510. ancestor = tree_index.find_ancestor_at_level(str(dop).strip(), cluster_level)
  511. if not ancestor:
  512. continue
  513. derived_ancestor_set.add(ancestor)
  514. pt = _entry_to_matched_point(entry)
  515. if pt and pt not in ancestor_to_matched.get(ancestor, []):
  516. ancestor_to_matched.setdefault(ancestor, []).append(pt)
  517. # 2. 筛选 pattern:已推导维度节点占所有元素的比例 >= 50%
  518. filtered_patterns: List[Dict[str, Any]] = []
  519. for p in patterns:
  520. items = p.get("items") or []
  521. item_names = [
  522. str(it.get("name") or "").strip()
  523. for it in items
  524. if isinstance(it, dict)
  525. ]
  526. if not item_names:
  527. continue
  528. if len(item_names) < 5:
  529. continue
  530. derived_count = sum(1 for name in item_names if name in derived_ancestor_set)
  531. if derived_count / len(item_names) >= 0.5:
  532. filtered_patterns.append(p)
  533. print(
  534. f"filtered_patterns: {len(filtered_patterns)}, "
  535. f"derived_ancestor_set: {len(derived_ancestor_set)}"
  536. )
  537. def _matched_join(name: str) -> str:
  538. pts = ancestor_to_matched.get(name) or []
  539. return ", ".join(pts) if pts else ""
  540. derived_ancestor_nodes: List[Dict[str, Any]] = []
  541. for anc in sorted(derived_ancestor_set):
  542. derived_ancestor_nodes.append(
  543. _dim_obj(anc, tree_index, matched_point=_matched_join(anc) or "")
  544. )
  545. # 3. 对筛选 pattern 元素分类并汇总维度列表
  546. derived_dims: List[Dict[str, Any]] = []
  547. underived_dims: List[Dict[str, Any]] = []
  548. derived_dims_seen: Set[str] = set()
  549. underived_dims_seen: Set[str] = set()
  550. scored_patterns: List[Dict[str, Any]] = []
  551. for p in filtered_patterns:
  552. items = p.get("items") or []
  553. tagged_items: List[Dict[str, Any]] = []
  554. for it in items:
  555. if not isinstance(it, dict):
  556. continue
  557. name = str(it.get("name") or "").strip()
  558. is_derived = name in derived_ancestor_set
  559. tagged_items.append(
  560. {
  561. "name": name,
  562. "is_derived": is_derived,
  563. }
  564. )
  565. if is_derived:
  566. if name and name not in derived_dims_seen:
  567. derived_dims_seen.add(name)
  568. derived_dims.append(
  569. _dim_obj(
  570. name,
  571. tree_index,
  572. matched_point=_matched_join(name) or "",
  573. )
  574. )
  575. else:
  576. if name and name not in underived_dims_seen:
  577. underived_dims_seen.add(name)
  578. underived_dims.append(_dim_obj(name, tree_index))
  579. scored_patterns.append(
  580. {
  581. "id": p.get("id"),
  582. "support": p.get("support"),
  583. "items": tagged_items,
  584. }
  585. )
  586. # 从 underived_dims 中排除与 derived_dims 重叠的节点
  587. underived_dims = [d for d in underived_dims if d["tree_node_name"] not in derived_dims_seen]
  588. # 按 is_derived=True 的元素数量从高到低排序,数量相同再按元素总数从高到低
  589. scored_patterns.sort(
  590. key=lambda x: (
  591. sum(1 for it in x.get("items", []) if it.get("is_derived")),
  592. len(x.get("items", [])),
  593. ),
  594. reverse=True,
  595. )
  596. return {
  597. "cumulative_points": list(cumulative_points),
  598. "derived_ancestor_nodes": derived_ancestor_nodes,
  599. "patterns": scored_patterns,
  600. "derived_dims": derived_dims,
  601. "underived_dims": underived_dims,
  602. "patterns_count": len(scored_patterns),
  603. "derived_dim_count": len(derived_dims),
  604. "underived_dim_count": len(underived_dims),
  605. }
  606. def _format_round_dimension_text(
  607. derived_dims: List[Dict[str, Any]],
  608. underived_dims: List[Dict[str, Any]],
  609. ) -> str:
  610. """已推导/未推导维度,每行:维度:tree_node_name,匹配点:matched_point"""
  611. lines: List[str] = ["【已推导的维度】"]
  612. for d in derived_dims:
  613. name = d.get("tree_node_name") or ""
  614. mp = d.get("matched_point") or "-"
  615. lines.append(f"维度:{name},匹配点:{mp}")
  616. if not derived_dims:
  617. lines.append("(无)")
  618. lines.append("")
  619. lines.append("【未推导的维度】")
  620. for d in underived_dims:
  621. name = d.get("tree_node_name") or ""
  622. lines.append(f"维度:{name}")
  623. if not underived_dims:
  624. lines.append("(无)")
  625. return "\n".join(lines)
  626. def pattern_dimension_analyze(
  627. account_name: str,
  628. post_id: str,
  629. log_id: str,
  630. ) -> Dict[str, Any]:
  631. """
  632. Pattern 维度分析主入口。
  633. 参数
  634. -------
  635. account_name : 账号名(用于定位 input / output 下的数据目录)
  636. post_id : 帖子 ID(用于定位推导日志)
  637. log_id : 推导日志目录名(../output/{account_name}/推导日志/{post_id}/{log_id}/)
  638. 逻辑概述
  639. --------
  640. 聚类层级固定为 CLUSTER_LEVEL(默认 3)。每一轮:
  641. 1. 从 derivation_output_point 在人设树中找到该层祖先节点 → 已推导维度节点集合。
  642. 2. 筛选包含已推导维度节点的 pattern。
  643. 3. 标记每个 pattern 元素是否已推导,汇总 derived_dims / underived_dims(对象列表)。
  644. """
  645. eval_dir = _round_eval_dir(account_name, post_id, log_id)
  646. if not eval_dir.is_dir():
  647. raise FileNotFoundError(f"推导日志目录不存在: {eval_dir}")
  648. round_infos = _load_round_matched_points(account_name, post_id, log_id)
  649. if not round_infos:
  650. return {
  651. "account_name": account_name,
  652. "post_id": post_id,
  653. "log_id": log_id,
  654. "cluster_level": CLUSTER_LEVEL,
  655. "rounds": [],
  656. "message": "未在指定日志目录下找到任何评估结果文件(*_评估.json)",
  657. }
  658. tree_index = TreeIndex(account_name)
  659. # pattern 库只在整体分析时读取 & 去重一次,避免每一轮重复 IO 与解析
  660. raw_patterns = _load_raw_patterns(account_name)
  661. deduped_patterns = _dedupe_patterns(raw_patterns)
  662. print(f"deduped_patterns len: {len(deduped_patterns)}")
  663. rounds_output: List[Dict[str, Any]] = []
  664. for info in round_infos:
  665. r = info["round"]
  666. cumulative_points = info["cumulative_points"]
  667. analyzed = _analyze_single_round(
  668. patterns=deduped_patterns,
  669. tree_index=tree_index,
  670. cumulative_points=cumulative_points,
  671. )
  672. analyzed["round"] = r
  673. rounds_output.append(analyzed)
  674. return {
  675. "account_name": account_name,
  676. "post_id": post_id,
  677. "log_id": log_id,
  678. "cluster_level": CLUSTER_LEVEL,
  679. "rounds": rounds_output,
  680. }
  681. def round_pattern_dimension_analyze_core(
  682. account_name: str,
  683. post_id: str,
  684. log_id: str,
  685. round: int,
  686. ) -> Dict[str, Any]:
  687. """
  688. 仅使用第 round 轮及之前的评估文件,得到该轮结束时的累计选题点状态并做维度分析。
  689. 返回 analyzed 单轮结构(含 derived_dims / underived_dims 等),失败时含 error 字段。
  690. """
  691. eval_dir = _round_eval_dir(account_name, post_id, log_id)
  692. if not eval_dir.is_dir():
  693. return {"error": f"推导日志目录不存在: {eval_dir}"}
  694. round_infos = _load_round_matched_points(
  695. account_name, post_id, log_id, max_round=round
  696. )
  697. if not round_infos:
  698. return {
  699. "error": f"在 {eval_dir} 下未找到第 {round} 轮及之前的 *_评估.json",
  700. }
  701. last = round_infos[-1]
  702. if last.get("round") != round:
  703. return {
  704. "error": (
  705. f"指定轮次 {round} 的评估文件不存在;"
  706. f"当前仅加载到第 {last.get('round')} 轮"
  707. ),
  708. }
  709. tree_index = TreeIndex(account_name)
  710. raw_patterns = _load_raw_patterns(account_name)
  711. deduped_patterns = _dedupe_patterns(raw_patterns)
  712. analyzed = _analyze_single_round(
  713. patterns=deduped_patterns,
  714. tree_index=tree_index,
  715. cumulative_points=last["cumulative_points"],
  716. )
  717. analyzed["round"] = round
  718. return analyzed
  719. @tool()
  720. async def round_pattern_dimension_analyze(
  721. account_name: str,
  722. post_id: str,
  723. log_id: str,
  724. round: int,
  725. ) -> Any:
  726. """
  727. 推导维度分析,返回当前轮次已推导的维度和可能的未推导维度数据
  728. Args:
  729. account_name: 账号名称
  730. post_id: 帖子 ID
  731. log_id: 推导日志目录名
  732. round: 推导轮次(正整数)
  733. Returns:
  734. ToolResult:output 为可读文本,含「已推导的维度」「未推导的维度」两段,
  735. 每行格式为「维度:tree_node_name,匹配点:matched_point」
  736. (未推导行固定为「-」;matched_point 规则:is_fully_derived 为真取选题点否则取推导输出点)。
  737. """
  738. if ToolResult is None:
  739. return None
  740. logger.info(
  741. "round_pattern_dimension_analyze: account=%s post_id=%s log_id=%s round=%s",
  742. account_name,
  743. post_id,
  744. log_id,
  745. round,
  746. )
  747. try:
  748. r = int(round)
  749. if r < 1:
  750. return ToolResult(
  751. title="维度分析: 轮次无效",
  752. output="",
  753. error="round 须为 >= 1 的整数",
  754. )
  755. except (TypeError, ValueError):
  756. return ToolResult(
  757. title="维度分析: 轮次无效",
  758. output="",
  759. error="round 须为整数",
  760. )
  761. try:
  762. analyzed = round_pattern_dimension_analyze_core(
  763. account_name, post_id, log_id, r
  764. )
  765. if analyzed.get("error"):
  766. return ToolResult(
  767. title=f"维度分析 第{r}轮 失败",
  768. output="",
  769. error=str(analyzed["error"]),
  770. )
  771. # 保存到与 {轮次}_评估.json 同级目录
  772. out_dir = _round_eval_dir(account_name, post_id, log_id)
  773. out_dir.mkdir(parents=True, exist_ok=True)
  774. out_path = out_dir / f"{r}_维度分析.json"
  775. with open(out_path, "w", encoding="utf-8") as f:
  776. json.dump(analyzed, f, ensure_ascii=False, indent=2)
  777. derived = analyzed.get("derived_dims") or []
  778. underived = analyzed.get("underived_dims") or []
  779. text = _format_round_dimension_text(derived, underived)
  780. meta = (
  781. f"round={r}, derived={len(derived)}, underived={len(underived)}, "
  782. f"patterns={analyzed.get('patterns_count', 0)}"
  783. )
  784. return ToolResult(
  785. title=f"第 {r} 轮维度分析(已推导 {len(derived)} / 未推导 {len(underived)})",
  786. output=text,
  787. metadata={"round_pattern_dimension_analyze": meta},
  788. )
  789. except Exception as e:
  790. logger.exception("round_pattern_dimension_analyze failed: %s", e)
  791. return ToolResult(
  792. title="维度分析失败",
  793. output="",
  794. error=str(e),
  795. )
  796. def main(account_name, post_id, log_id) -> None:
  797. """本地简单测试:以家有大志账号的一次推导日志做分析,并将结果写入输出目录。"""
  798. result = pattern_dimension_analyze(
  799. account_name=account_name,
  800. post_id=post_id,
  801. log_id=log_id,
  802. )
  803. # 控制台打印前 4000 字符,便于快速查看
  804. # print(json.dumps(result, ensure_ascii=False, indent=2)[:4000] + "...")
  805. # 写入输出文件:../output/{account_name}/推导日志/{post_id}/{log_id}/pattern_dimension_analyze.json
  806. out_dir = _round_eval_dir(account_name, post_id, log_id)
  807. out_dir.mkdir(parents=True, exist_ok=True)
  808. output_file_name = f"{post_id}_pattern_dimension_analyze.json"
  809. out_path = out_dir / output_file_name
  810. with open(out_path, "w", encoding="utf-8") as f:
  811. json.dump(result, f, ensure_ascii=False, indent=2)
  812. print(f"\n分析结果已写入: {out_path}")
  813. def main_round_pattern_dimension_analyze(
  814. account_name: str,
  815. post_id: str,
  816. log_id: str,
  817. round_num: int,
  818. ) -> None:
  819. """本地测试:直接调用 round_pattern_dimension_analyze,打印 ToolResult。"""
  820. import asyncio
  821. async def _run() -> None:
  822. result = await round_pattern_dimension_analyze(
  823. account_name=account_name,
  824. post_id=post_id,
  825. log_id=log_id,
  826. round=round_num,
  827. context=None,
  828. )
  829. if result is None:
  830. print(
  831. "round_pattern_dimension_analyze 返回 None:请先将 Agent 项目根目录加入 PYTHONPATH,"
  832. "或在 __main__ 中保证能 import agent.tools.ToolResult"
  833. )
  834. return
  835. if result.error:
  836. print(f"错误: {result.error}")
  837. else:
  838. print(result.title)
  839. print(result.output)
  840. asyncio.run(_run())
  841. if __name__ == "__main__":
  842. import asyncio
  843. import importlib.util
  844. # 直接加载 ToolResult,避免 import agent 时拉全量依赖(如 langchain)
  845. _agent_root = Path(__file__).resolve().parents[3]
  846. _models_py = _agent_root / "agent" / "tools" / "models.py"
  847. if _models_py.is_file():
  848. _spec = importlib.util.spec_from_file_location(
  849. "_pattern_dim_tool_models", _models_py
  850. )
  851. if _spec and _spec.loader:
  852. _m = importlib.util.module_from_spec(_spec)
  853. _spec.loader.exec_module(_m)
  854. globals()["ToolResult"] = _m.ToolResult
  855. if getattr(_m, "ToolContext", None) is not None:
  856. globals()["ToolContext"] = _m.ToolContext
  857. # ---------- 开关与参数(改这里即可) ----------
  858. run_round_pattern_test = False
  859. run_full_pattern_analyze = True
  860. test_account_name = "家有大志"
  861. test_post_id = "68fb6a5c000000000302e5de"
  862. test_log_id = "20260317214307"
  863. test_round = 1
  864. items = [
  865. {"post_id": "68fb6a5c000000000302e5de", "log_id": "20260318172724"},
  866. # {"post_id": "69185d49000000000d00f94e", "log_id": "20260317214841"},
  867. # {"post_id": "6921937a000000001b0278d1", "log_id": "20260317215616"},
  868. ]
  869. if run_round_pattern_test:
  870. main_round_pattern_dimension_analyze(
  871. test_account_name,
  872. test_post_id,
  873. test_log_id,
  874. test_round,
  875. )
  876. elif run_full_pattern_analyze:
  877. for item in items:
  878. main(test_account_name, item["post_id"], item["log_id"])