cartesian.py 13 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316
  1. """Cartesian creation-query generator and legacy preview projection."""
  2. from __future__ import annotations
  3. import json
  4. import random
  5. from dataclasses import dataclass
  6. from itertools import product
  7. from pathlib import Path
  8. from typing import Any
  9. from acquisition.queries.axes import ACTIONS, STAGES, _nonleaf_d4
  10. from core.config import Settings
  11. from query_planning import (
  12. GenerationRequest,
  13. GeneratorKind,
  14. UnifiedQueryGenerationService,
  15. )
  16. ROOT = Path(__file__).resolve().parents[1]
  17. TREES = ROOT / "scope_trees" / "trees_index.json"
  18. KTYPE = ["怎么做", "有哪些", "为什么"]
  19. MODALITY = ["视频", "图文"]
  20. INTENT = ["灵感", "选题", "脚本"]
  21. DEFAULT_ACTIVE_FAMILY_KEYS = ("f1", "f2")
  22. def _planning_service() -> UnifiedQueryGenerationService:
  23. return UnifiedQueryGenerationService()
  24. @dataclass(frozen=True)
  25. class QueryBuildOptions:
  26. per: int = 0
  27. batch_n: int = 0
  28. seed: int = 7
  29. active_family_keys: tuple[str, ...] = DEFAULT_ACTIVE_FAMILY_KEYS
  30. def _segs(path: str | None) -> list[str]:
  31. return [x for x in (path or "").split("/") if x]
  32. def _leaves(idx: list[dict[str, Any]], source_type: str, under: str | None = None) -> list[str]:
  33. paths = [(_segs(n.get("path")), n.get("name")) for n in idx if n.get("source_type") == source_type]
  34. if under:
  35. paths = [(s, nm) for s, nm in paths if under in s]
  36. all_paths = {"/".join(s) for s, _ in paths}
  37. out: list[str] = []
  38. seen: set[str] = set()
  39. for segs, name in paths:
  40. if len(segs) < 2:
  41. continue
  42. full = "/".join(segs)
  43. is_leaf = not any(other != full and other.startswith(full + "/") for other in all_paths)
  44. value = name or segs[-1]
  45. if is_leaf and value and value not in seen:
  46. seen.add(value)
  47. out.append(value)
  48. return out
  49. def _nodes_at_depth(
  50. idx: list[dict[str, Any]],
  51. source_type: str,
  52. depths: tuple[int, ...] = (3, 4),
  53. under: str | None = None,
  54. ) -> list[str]:
  55. paths = [(_segs(n.get("path")), n.get("name")) for n in idx if n.get("source_type") == source_type]
  56. if under:
  57. paths = [(s, nm) for s, nm in paths if under in s]
  58. out: list[str] = []
  59. seen: set[str] = set()
  60. for segs, name in paths:
  61. if len(segs) not in depths:
  62. continue
  63. value = name or (segs[-1] if segs else "")
  64. if value and value not in seen:
  65. seen.add(value)
  66. out.append(value)
  67. return out
  68. def _axis_tree(idx: list[dict[str, Any]], source_type: str) -> list[dict[str, Any]]:
  69. nodes: dict[str, dict[str, Any]] = {}
  70. for row in idx:
  71. if row.get("source_type") != source_type:
  72. continue
  73. segs = _segs(row.get("path"))
  74. if len(segs) not in (3, 4):
  75. continue
  76. path = "/" + "/".join(segs)
  77. nodes[path] = {
  78. "name": row.get("name") or segs[-1],
  79. "path": path,
  80. "level": len(segs),
  81. "children": [],
  82. }
  83. roots: list[dict[str, Any]] = []
  84. for path, node in nodes.items():
  85. if node["level"] == 3:
  86. roots.append(node)
  87. continue
  88. parent_path = "/" + "/".join(_segs(path)[:-1])
  89. parent = nodes.get(parent_path)
  90. if parent:
  91. parent["children"].append(node)
  92. roots.sort(key=lambda node: node["path"])
  93. for node in roots:
  94. node["children"].sort(key=lambda child: child["path"])
  95. return roots
  96. def _sample_axis(values: list[str], limit: int, rng: random.Random) -> list[str]:
  97. if limit <= 0 or limit >= len(values):
  98. return values
  99. return rng.sample(values, limit)
  100. def build_creation_query_batch(
  101. settings: Settings,
  102. *,
  103. tree_path: Path = TREES,
  104. options: QueryBuildOptions | None = None,
  105. ) -> dict[str, Any]:
  106. """Build a creation-query batch without writing files or database rows."""
  107. opts = options or QueryBuildOptions()
  108. rng = random.Random(opts.seed)
  109. idx = json.loads(tree_path.read_text("utf-8"))
  110. shi = _nodes_at_depth(idx, "实质", depths=(3, 4))
  111. xing = _nodes_at_depth(idx, "形式", depths=(3, 4))
  112. purpose_pool = _leaves(idx, "作用") + _leaves(idx, "感受") + _leaves(idx, "意图")
  113. shi_batch = _sample_axis(shi, opts.batch_n, rng)
  114. xing_batch = _sample_axis(xing, opts.batch_n, rng)
  115. purpose_batch = _sample_axis(purpose_pool, opts.batch_n, rng)
  116. f6_zy = _nonleaf_d4("作用", 10, tree_path=tree_path)
  117. f6_stage_act = [(s, a) for s in STAGES for a in ACTIONS] + [("", "")]
  118. if not shi_batch or not xing_batch or not purpose_batch or not f6_zy:
  119. raise RuntimeError("scope tree does not contain enough creation query axes")
  120. product_pools = {
  121. "实质": shi_batch,
  122. "形式": xing_batch,
  123. "目的": purpose_batch,
  124. "模态": MODALITY,
  125. "业务阶段": INTENT,
  126. "知识类型": KTYPE,
  127. }
  128. def gen_product(keys: list[str]):
  129. for values in product(*(product_pools[key] for key in keys)):
  130. yield {"parts": dict(zip(keys, values, strict=True))}
  131. legacy_per = opts.per if opts.per > 0 else 30
  132. master: list[dict[str, str]] = []
  133. for i in range(legacy_per):
  134. stage, action = f6_stage_act[i % len(f6_stage_act)]
  135. master.append(
  136. {
  137. "实质": shi_batch[i % len(shi_batch)],
  138. "形式": xing_batch[i % len(xing_batch)],
  139. "目的": purpose_batch[i % len(purpose_batch)],
  140. "模态": MODALITY[i % len(MODALITY)],
  141. "业务阶段": INTENT[i % len(INTENT)],
  142. "知识类型": KTYPE[(i // 3) % len(KTYPE)],
  143. "阶段": stage or "/",
  144. "动作": action or "/",
  145. "_st": stage,
  146. "_ac": action,
  147. "作用": f6_zy[i % len(f6_zy)],
  148. }
  149. )
  150. def gen_old(
  151. i: int,
  152. *,
  153. shi_axis: bool = False,
  154. xing_axis: bool = False,
  155. purpose_axis: bool = False,
  156. effect_axis: bool = True,
  157. ) -> dict[str, Any]:
  158. row = master[i]
  159. segment = (row["_st"] + row["_ac"]) if row["_ac"] else ""
  160. head = (
  161. ([row["实质"]] if shi_axis else [])
  162. + ([row["形式"]] if xing_axis else [])
  163. + ([row["目的"]] if purpose_axis else [])
  164. )
  165. tail = ([row["作用"]] if effect_axis else []) + [row["知识类型"]]
  166. query = " ".join(head + ([segment] if segment else []) + tail)
  167. parts: dict[str, str] = {}
  168. if shi_axis:
  169. parts["实质"] = row["实质"]
  170. if xing_axis:
  171. parts["形式"] = row["形式"]
  172. if purpose_axis:
  173. parts["目的"] = row["目的"]
  174. parts["阶段"], parts["动作"] = row["阶段"], row["动作"]
  175. if effect_axis:
  176. parts["作用"] = row["作用"]
  177. parts["知识类型"] = row["知识类型"]
  178. return {"parts": parts, "query": query}
  179. families = [
  180. {"key": "f1", "axes": ["实质", "模态", "业务阶段", "知识类型"], "items": lambda: gen_product(["实质", "模态", "业务阶段", "知识类型"])},
  181. {"key": "f2", "axes": ["形式", "模态", "业务阶段", "知识类型"], "items": lambda: gen_product(["形式", "模态", "业务阶段", "知识类型"])},
  182. {"key": "f4", "axes": ["作用/感受/意图", "模态", "业务阶段", "知识类型"], "items": lambda: gen_product(["目的", "模态", "业务阶段", "知识类型"])},
  183. {"key": "f3", "axes": ["实质", "形式", "模态", "业务阶段", "知识类型"], "items": lambda: gen_product(["实质", "形式", "模态", "业务阶段", "知识类型"])},
  184. {"key": "f5", "axes": ["模态", "业务阶段", "知识类型"], "items": lambda: gen_product(["模态", "业务阶段", "知识类型"])},
  185. {"key": "a_shi", "axes": ["实质", "阶段", "动作", "作用", "知识类型"], "gen": lambda i: gen_old(i, shi_axis=True)},
  186. {"key": "a_xing", "axes": ["形式", "阶段", "动作", "作用", "知识类型"], "gen": lambda i: gen_old(i, xing_axis=True)},
  187. {"key": "a_both", "axes": ["实质", "形式", "阶段", "动作", "作用", "知识类型"], "gen": lambda i: gen_old(i, shi_axis=True, xing_axis=True)},
  188. {"key": "a_purpose", "axes": ["作用/感受/意图", "阶段", "动作", "作用", "知识类型"], "gen": lambda i: gen_old(i, purpose_axis=True)},
  189. {"key": "a_tail", "axes": ["阶段", "动作", "作用", "知识类型"], "gen": lambda i: gen_old(i)},
  190. {"key": "b_shi", "axes": ["实质", "阶段", "动作", "知识类型"], "gen": lambda i: gen_old(i, shi_axis=True, effect_axis=False)},
  191. {"key": "b_xing", "axes": ["形式", "阶段", "动作", "知识类型"], "gen": lambda i: gen_old(i, xing_axis=True, effect_axis=False)},
  192. {"key": "b_both", "axes": ["实质", "形式", "阶段", "动作", "知识类型"], "gen": lambda i: gen_old(i, shi_axis=True, xing_axis=True, effect_axis=False)},
  193. {"key": "b_purpose", "axes": ["作用/感受/意图", "阶段", "动作", "知识类型"], "gen": lambda i: gen_old(i, purpose_axis=True, effect_axis=False)},
  194. {"key": "b_tail", "axes": ["阶段", "动作", "知识类型"], "gen": lambda i: gen_old(i, effect_axis=False)},
  195. ]
  196. order = ["实质", "形式", "目的", "模态", "业务阶段", "知识类型"]
  197. out: dict[str, Any] = {
  198. "axis_values": {
  199. "实质": shi,
  200. "形式": xing,
  201. "目的池": purpose_pool,
  202. "业务阶段": INTENT,
  203. "模态": MODALITY,
  204. "知识类型": KTYPE,
  205. "阶段": STAGES,
  206. "动作": ACTIONS,
  207. "作用": f6_zy,
  208. },
  209. "axis_trees": {
  210. "实质": _axis_tree(idx, "实质"),
  211. "形式": _axis_tree(idx, "形式"),
  212. },
  213. "metadata": {
  214. "seed": opts.seed,
  215. "per": opts.per,
  216. "batch_n": opts.batch_n,
  217. "active_family_keys": list(opts.active_family_keys),
  218. },
  219. "families": [],
  220. }
  221. family_by_key = {family["key"]: family for family in families}
  222. unknown = [key for key in opts.active_family_keys if key not in family_by_key]
  223. if unknown:
  224. raise ValueError(f"unknown query family key(s): {', '.join(unknown)}")
  225. for family_key in opts.active_family_keys:
  226. family = family_by_key[family_key]
  227. name = " × ".join(family["axes"])
  228. seen: set[str] = set()
  229. items: list[dict[str, Any]] = []
  230. generated_items = family["items"]() if "items" in family else (family["gen"](i) for i in range(legacy_per))
  231. for generated in generated_items:
  232. parts = generated["parts"]
  233. query = generated.get("query") or " ".join(parts[k] for k in order if k in parts)
  234. if query in seen:
  235. continue
  236. seen.add(query)
  237. items.append({"query": query, "parts": parts})
  238. if opts.per > 0 and len(items) >= opts.per:
  239. break
  240. for item in items:
  241. item.update({"keep": True, "reason": ""})
  242. planned = _planning_service().generate(
  243. GenerationRequest(
  244. generator_kind=GeneratorKind.CARTESIAN,
  245. name=f"cartesian-preview:{family['key']}",
  246. payload={
  247. "candidates": [
  248. {
  249. "query_text": item["query"],
  250. "axes": item["parts"],
  251. "filter_reason": item["reason"],
  252. "metadata": {
  253. "family_key": family["key"],
  254. "family_name": name,
  255. "family_axes": family["axes"],
  256. },
  257. "source_refs": [
  258. {
  259. "generator_kind": GeneratorKind.CARTESIAN.value,
  260. "family_key": family["key"],
  261. "axes": item["parts"],
  262. }
  263. ],
  264. }
  265. for item in items
  266. ],
  267. "input_snapshot": {
  268. "family_key": family["key"],
  269. "family_axes": family["axes"],
  270. "options": out["metadata"],
  271. },
  272. },
  273. )
  274. )
  275. items = [
  276. {
  277. "query": candidate.query_text,
  278. "parts": candidate.axes,
  279. "keep": True,
  280. "reason": candidate.filter_reason or "",
  281. }
  282. for candidate in planned.selected_queries
  283. ]
  284. out["families"].append(
  285. {"key": family["key"], "name": name, "axes": family["axes"], "items": items}
  286. )
  287. return out