1
0

main.py 8.5 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232
  1. """
  2. 配载智能体 · 命令行入口
  3. 把 src/ 里那七个模块串起来,变成一个能用的东西。
  4. 常用命令::
  5. python main.py tools # 看看有哪些工具
  6. python main.py ask "40尺箱写哪个贝号?" # 问一个问题
  7. python main.py ask "..." --show-steps # 顺便看它思考了几个来回
  8. python main.py ask "..." --no-llm # 只用检索,不调大模型(免费、看原理)
  9. python main.py chat # 连续对话
  10. python main.py eval --offline # 工具自检(免费、秒出)
  11. python main.py eval --online # 端到端评估(要调大模型)
  12. python main.py eval --all # 两层都跑
  13. """
  14. from __future__ import annotations
  15. import argparse
  16. import json
  17. import sys
  18. from pathlib import Path
  19. # 让 `python main.py` 在任意目录下执行都能找到 src 包
  20. ROOT = Path(__file__).resolve().parent
  21. sys.path.insert(0, str(ROOT))
  22. from src import ( # noqa: E402
  23. BayTool,
  24. CoordTool,
  25. KnowledgeTool,
  26. LLM,
  27. LLMError,
  28. Memory,
  29. ReActAgent,
  30. Retriever,
  31. StowageCheckTool,
  32. ToolRegistry,
  33. check_tools,
  34. evaluate_agent,
  35. )
  36. from src.evaluate import print_agent_report, print_tool_report # noqa: E402
  37. KB_FILE = ROOT / "data" / "配载知识库.txt"
  38. CASE_FILE = ROOT / "data" / "测试用例.json"
  39. NOTE_FILE = ROOT / "outputs" / "长期笔记.md"
  40. # ======================================================================
  41. def build_registry(mode: str = "tfidf") -> ToolRegistry:
  42. """装好 4 个工具(工具系统)。"""
  43. print(f"⏳ 正在加载知识库并建索引(模式:{mode})…")
  44. retriever = Retriever(KB_FILE, mode=mode)
  45. retriever.index()
  46. print(f"✅ 知识库已就绪,共 {len(retriever.chunks)} 段资料")
  47. registry = ToolRegistry()
  48. registry.register(BayTool())
  49. registry.register(CoordTool())
  50. registry.register(StowageCheckTool())
  51. registry.register(KnowledgeTool(retriever, top_k=3))
  52. return registry
  53. def build_agent(args) -> ReActAgent:
  54. llm = LLM()
  55. print(f"✅ 已接入大模型:{llm.model}")
  56. registry = build_registry(args.mode)
  57. # chat / eval 子命令不一定带了 --show-steps 参数,用 getattr 兜一下默认值
  58. agent = ReActAgent(
  59. llm=llm,
  60. registry=registry,
  61. memory=Memory(note_path=NOTE_FILE),
  62. max_steps=getattr(args, "max_steps", 5),
  63. verbose=getattr(args, "show_steps", False),
  64. )
  65. return agent
  66. # ======================================================================
  67. # 命令一:列出工具
  68. # ======================================================================
  69. def cmd_tools(args) -> None:
  70. registry = build_registry(args.mode)
  71. print("\n可用工具:\n")
  72. for name in registry.names:
  73. tool = registry.get(name)
  74. print(f" 🔧 {name}")
  75. print(f" {tool.description}\n")
  76. print("想试试?python main.py ask \"贝位 01 和 03 合成几号?\"")
  77. # ======================================================================
  78. # 命令二:问一句(--no-llm 时只做检索,看 RAG 的原理)
  79. # ======================================================================
  80. def cmd_ask(args) -> None:
  81. if args.no_llm:
  82. registry = build_registry(args.mode)
  83. hits = registry.get("知识库检索").run(args.question)
  84. print("\n" + "=" * 56)
  85. print("【检索结果】(这些资料会被贴到问题前面,一起发给大模型)")
  86. print("=" * 56)
  87. print(hits)
  88. print("\n" + "=" * 56)
  89. print("【实际会发给大模型的提示词长这样】")
  90. print("=" * 56)
  91. print("只准根据下面资料回答,资料里没有就说「没有提到」。\n")
  92. print(f"资料:\n{hits}\n")
  93. print(f"问题:{args.question}")
  94. return
  95. agent = build_agent(args)
  96. result = agent.run(args.question)
  97. print("\n" + "=" * 56)
  98. print("最终答案")
  99. print("=" * 56)
  100. print(result["answer"])
  101. print("\n" + "-" * 56)
  102. print(f"思考-行动轮数:{len(result['steps'])} 是否正常结束:{result['ok']}")
  103. # ======================================================================
  104. # 命令三:连续对话(演示"短期记忆")
  105. # ======================================================================
  106. def cmd_chat(args) -> None:
  107. agent = build_agent(args)
  108. print("\n进入连续对话模式。它能记住你前面说过的话(短期记忆)。")
  109. print("输入 exit / quit 退出,输入 :clear 清空记忆。\n")
  110. while True:
  111. try:
  112. question = input("你 > ").strip()
  113. except (EOFError, KeyboardInterrupt):
  114. print()
  115. break
  116. if not question:
  117. continue
  118. if question.lower() in {"exit", "quit", ":q"}:
  119. break
  120. if question == ":clear":
  121. agent.memory.clear()
  122. print("(已清空短期记忆)")
  123. continue
  124. result = agent.run(question)
  125. print(f"\n助手 > {result['answer']}\n")
  126. print("-" * 56)
  127. # ======================================================================
  128. # 命令四:评估
  129. # ======================================================================
  130. def cmd_eval(args) -> None:
  131. report = {}
  132. if args.offline or args.all:
  133. registry = build_registry("tfidf") # 离线自检固定用 tfidf,结果才稳定
  134. report["工具自检"] = check_tools(registry, CASE_FILE)
  135. print_tool_report(report["工具自检"])
  136. if args.online or args.all:
  137. agent = build_agent(args)
  138. print("\n开始端到端评估(每一题都要调一次大模型,慢是正常的)…\n")
  139. report["端到端"] = evaluate_agent(agent.run, CASE_FILE, verbose=True)
  140. print_agent_report(report["端到端"])
  141. # 报告做增量合并:只跑了一层时,不要把另一层的历史结果抹掉
  142. out = ROOT / "outputs" / "评估报告.json"
  143. out.parent.mkdir(parents=True, exist_ok=True)
  144. merged = {}
  145. if out.exists():
  146. try:
  147. merged = json.loads(out.read_text(encoding="utf-8"))
  148. except json.JSONDecodeError:
  149. merged = {}
  150. merged.update(report)
  151. out.write_text(json.dumps(merged, ensure_ascii=False, indent=2), encoding="utf-8")
  152. print(f"\n📄 报告已保存:{out}(本次覆盖的层:{'、'.join(report)})")
  153. # ======================================================================
  154. def main() -> int:
  155. parser = argparse.ArgumentParser(
  156. description="配载智能体 —— 集装箱船配载规则问答与坐标计算",
  157. formatter_class=argparse.RawDescriptionHelpFormatter,
  158. epilog=__doc__,
  159. )
  160. parser.add_argument("--mode", default="tfidf", choices=["tfidf", "vector"],
  161. help="检索模式:tfidf(默认,秒开)或 vector(语义向量,首次要下模型)")
  162. parser.add_argument("--max-steps", type=int, default=5, help="ReAct 最多循环几轮")
  163. sub = parser.add_subparsers(dest="command", required=True)
  164. sub.add_parser("tools", help="列出所有工具")
  165. p_ask = sub.add_parser("ask", help="问一个问题")
  166. p_ask.add_argument("question", help="你要问的话")
  167. p_ask.add_argument("--show-steps", action="store_true", help="打印每一轮的思考过程")
  168. p_ask.add_argument("--no-llm", action="store_true", help="只用检索,不调用大模型")
  169. sub.add_parser("chat", help="连续对话")
  170. p_eval = sub.add_parser("eval", help="跑评估")
  171. p_eval.add_argument("--offline", action="store_true", help="只跑工具自检(不花钱)")
  172. p_eval.add_argument("--online", action="store_true", help="只跑端到端评估")
  173. p_eval.add_argument("--all", action="store_true", help="两层都跑")
  174. p_eval.add_argument("--show-steps", action="store_true", help="端到端评估时打印思考过程")
  175. args = parser.parse_args()
  176. if args.command == "eval" and not (args.offline or args.online or args.all):
  177. args.offline = True # 默认跑最便宜的那层
  178. try:
  179. {
  180. "tools": cmd_tools,
  181. "ask": cmd_ask,
  182. "chat": cmd_chat,
  183. "eval": cmd_eval,
  184. }[args.command](args)
  185. except LLMError as exc:
  186. print(f"\n❌ 大模型调用出问题:{exc}")
  187. print(" 检查一下 .env 里的 LLM_API_KEY / LLM_BASE_URL / LLM_MODEL_ID。")
  188. return 1
  189. except FileNotFoundError as exc:
  190. print(f"\n❌ 找不到文件:{exc}")
  191. return 1
  192. return 0
  193. if __name__ == "__main__":
  194. raise SystemExit(main())