chart.py 18 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465466467468469470471472473474475476
  1. import json
  2. import os
  3. import matplotlib.pyplot as plt
  4. import matplotlib.patches as mpatches
  5. import numpy as np
  6. from matplotlib.ticker import FuncFormatter
  7. path_agent_history = '../../outputs/agent/'
  8. filename_log_pssr = '../../outputs/log/log_pssr_agent.json'
  9. def load_json(filename):
  10. """打开json文件并解析成dict"""
  11. with open(filename, "r", encoding="utf-8") as f:
  12. return json.load(f)
  13. def save_json(data, filename):
  14. """将dict存储为json文件"""
  15. filename_bk = filename + '.bk'
  16. with open(filename_bk, "w", encoding="utf-8") as f:
  17. json.dump(data, f, ensure_ascii=False, indent=2)
  18. if os.path.exists(filename):
  19. os.remove(filename)
  20. os.rename(filename_bk, filename)
  21. def get_problem_level():
  22. """返回问题problem_id到问题难度的映射"""
  23. map_pid_to_level = {}
  24. for problem_id in load_json(filename_log_pssr)['total']:
  25. theorem_length = len(load_json(f'../../datasets/problems/{problem_id}.json')['theorem_seqs'])
  26. if theorem_length > 12:
  27. map_pid_to_level[problem_id] = 6
  28. else:
  29. map_pid_to_level[problem_id] = int(theorem_length / 2) + theorem_length % 2
  30. return map_pid_to_level
  31. def get_avg_context_len():
  32. """
  33. 按照问题难度,统计平均上下文长度。每个问题的上下文长度是 solving_history_pid.json 文件中,所有content长度的和; len(content)
  34. 此外,还要分为 已求解的问题 和 其他问题
  35. """
  36. log = load_json(filename_log_pssr)
  37. map_pid_to_level = get_problem_level()
  38. solved_pids = {int(k) for k in log['solved']}
  39. # {level: [total_len, count]}
  40. solved_accum = {l: [0, 0] for l in range(1, 7)}
  41. others_accum = {l: [0, 0] for l in range(1, 7)}
  42. for pid in log['total']:
  43. level = map_pid_to_level.get(pid)
  44. if level is None:
  45. continue
  46. hist = load_json(f'{path_agent_history}solving_history_{pid}.json')
  47. ctx_len = sum(
  48. len(str(msg.get('content', '')))
  49. for round_msgs in hist.get('history', [])
  50. if isinstance(round_msgs, list)
  51. for msg in round_msgs
  52. )
  53. accum = solved_accum if pid in solved_pids else others_accum
  54. accum[level][0] += ctx_len
  55. accum[level][1] += 1
  56. # dict: key 为 problem_level; value 为 avg_context_length
  57. avg_context_len_solved = {
  58. l: (solved_accum[l][0] / solved_accum[l][1] if solved_accum[l][1] > 0 else None)
  59. for l in range(1, 7)
  60. }
  61. # unsolved + timeout + error; 如果当前等级的问题没有,则key 为 None
  62. avg_context_length_others = {
  63. l: (others_accum[l][0] / others_accum[l][1] if others_accum[l][1] > 0 else None)
  64. for l in range(1, 7)
  65. }
  66. return avg_context_len_solved, avg_context_length_others
  67. def get_avg_epoch():
  68. """
  69. 按照问题难度,统计平均交互次数。每个问题的交互次数存储在 solving_history_pid.json 文件中。
  70. 此外,还要分为 已求解的问题 和 其他问题
  71. """
  72. log = load_json(filename_log_pssr)
  73. map_pid_to_level = get_problem_level()
  74. # {level: [total_epoch, count]}
  75. solved_accum = {l: [0, 0] for l in range(1, 7)}
  76. others_accum = {l: [0, 0] for l in range(1, 7)}
  77. for cat in ('solved', 'unsolved', 'timeout', 'error'):
  78. accum = solved_accum if cat == 'solved' else others_accum
  79. for pid_str, info in log[cat].items():
  80. level = map_pid_to_level.get(int(pid_str))
  81. if level is None:
  82. continue
  83. accum[level][0] += info['epoch']
  84. accum[level][1] += 1
  85. # dict: key 为 problem_level; value 为 avg_epoch
  86. avg_epoch_solved = {
  87. l: (solved_accum[l][0] / solved_accum[l][1] if solved_accum[l][1] > 0 else None)
  88. for l in range(1, 7)
  89. }
  90. # unsolved + timeout + error; 如果当前等级的问题没有,则key 为 None
  91. avg_epoch_others = {
  92. l: (others_accum[l][0] / others_accum[l][1] if others_accum[l][1] > 0 else None)
  93. for l in range(1, 7)
  94. }
  95. return avg_epoch_solved, avg_epoch_others
  96. def get_tool_call():
  97. """
  98. 按照问题难度,统计所有工具的平均调用次数。需要解析每个问题solving_history_pid.json 文件中 role 为 assistance 的消息
  99. 当json解析出错时,记为error
  100. """
  101. tool_keys = ['apply', 'decompose', 'find', 'check', 'error']
  102. log = load_json(filename_log_pssr)
  103. map_pid_to_level = get_problem_level()
  104. solved_pids = {int(k) for k in log['solved']}
  105. # {level: {tool: total_count}}, {level: problem_count}
  106. solved_count = {l: {t: 0 for t in tool_keys} for l in range(1, 7)}
  107. others_count = {l: {t: 0 for t in tool_keys} for l in range(1, 7)}
  108. solved_n = {l: 0 for l in range(1, 7)}
  109. others_n = {l: 0 for l in range(1, 7)}
  110. for pid in log['total']:
  111. level = map_pid_to_level.get(pid)
  112. if level is None:
  113. continue
  114. hist = load_json(f'{path_agent_history}solving_history_{pid}.json')
  115. is_solved = pid in solved_pids
  116. count = solved_count[level] if is_solved else others_count[level]
  117. for round_msgs in hist.get('history', []):
  118. if not isinstance(round_msgs, list):
  119. continue
  120. for msg in round_msgs:
  121. if msg.get('role') != 'assistant':
  122. continue
  123. try:
  124. parsed = json.loads(str(msg.get('content', '')))
  125. tool = parsed.get('action', '').split('(')[0].strip()
  126. if tool in ['find_fact', 'find_goal']:
  127. tool = 'find'
  128. if tool in tool_keys:
  129. count[tool] += 1
  130. except Exception:
  131. count['error'] += 1
  132. if is_solved:
  133. solved_n[level] += 1
  134. else:
  135. others_n[level] += 1
  136. # dict: key 为 problem_level; value 为 平均tool_call次数
  137. avg_tool_call_solved = {
  138. l: ({t: solved_count[l][t] / solved_n[l] for t in tool_keys} if solved_n[l] > 0 else None)
  139. for l in range(1, 7)
  140. }
  141. # unsolved + timeout + error; 如果当前等级的问题没有,则key 为 None
  142. avg_tool_call_others = {
  143. l: ({t: others_count[l][t] / others_n[l] for t in tool_keys} if others_n[l] > 0 else None)
  144. for l in range(1, 7)
  145. }
  146. return avg_tool_call_solved, avg_tool_call_others
  147. def draw_figure():
  148. """
  149. 结合上述三个数据画图
  150. """
  151. avg_context_len_solved, avg_context_length_others = get_avg_context_len()
  152. avg_epoch_solved, avg_epoch_others = get_avg_epoch()
  153. avg_tool_call_solved, avg_tool_call_others = get_tool_call()
  154. levels = [1, 2, 3, 4, 5, 6]
  155. tool_keys = ['apply', 'decompose', 'find', 'check', 'error']
  156. n_tools = len(tool_keys)
  157. # 全局设置
  158. plt.rcParams.update({
  159. 'font.family': 'serif',
  160. 'font.size': 10,
  161. 'axes.linewidth': 1.0,
  162. 'xtick.direction': 'out',
  163. 'ytick.direction': 'out',
  164. 'xtick.major.size': 4,
  165. 'ytick.major.size': 4,
  166. 'figure.dpi': 150,
  167. })
  168. # 柱状图色板
  169. tool_colors = [
  170. '#55A868', '#5DA5DA', '#9970AB', '#E6AB02', '#E7298A',
  171. ]
  172. bar_width = 0.2
  173. level_spacing = n_tools * bar_width + 0.2
  174. x_centers = np.arange(len(levels)) * level_spacing
  175. fig, ax_bar = plt.subplots(figsize=(10, 4.2))
  176. ax_ctx = ax_bar.twinx()
  177. ax_epoch = ax_bar.twinx()
  178. ax_ctx.yaxis.set_label_position('left')
  179. ax_ctx.yaxis.tick_left()
  180. ax_bar.yaxis.set_visible(False)
  181. # 顶部封边
  182. ax_bar.spines['top'].set_visible(True)
  183. ax_ctx.spines['top'].set_visible(False)
  184. ax_epoch.spines['top'].set_visible(False)
  185. # --- 发散柱状图 ---
  186. for i, tool in enumerate(tool_keys):
  187. sv = [avg_tool_call_solved[l][tool] if avg_tool_call_solved[l] is not None else 0 for l in levels]
  188. ov = [avg_tool_call_others[l][tool] if avg_tool_call_others[l] is not None else 0 for l in levels]
  189. x_pos = x_centers + (i - n_tools / 2 + 0.5) * bar_width
  190. bars_s = ax_bar.bar(x_pos, [-v for v in sv], width=bar_width,
  191. color=tool_colors[i], edgecolor='white', linewidth=0.3, zorder=2)
  192. bars_o = ax_bar.bar(x_pos, ov, width=bar_width,
  193. color=tool_colors[i], edgecolor='white', linewidth=0.3,
  194. hatch='////', alpha=0.75, zorder=2)
  195. # 柱子向下 (Solved) 的文本
  196. for bar, val in zip(bars_s, sv):
  197. ax_epoch.text(bar.get_x() + bar.get_width() / 2, -val - 0.05,
  198. f'{val:.1f}', ha='center', va='top', fontsize=6,
  199. color='#333333', fontfamily='sans-serif',
  200. fontweight='bold',
  201. zorder=10,
  202. transform=ax_bar.transData)
  203. # 柱子向上 (Failed / Others) 的文本
  204. for bar, val in zip(bars_o, ov):
  205. ax_epoch.text(bar.get_x() + bar.get_width() / 2, bar.get_height() + 0.05,
  206. f'{val:.1f}', ha='center', va='bottom', fontsize=6,
  207. color='#333333', fontfamily='sans-serif',
  208. fontweight='bold',
  209. zorder=10,
  210. transform=ax_bar.transData)
  211. ax_bar.axhline(0, color='#333333', linewidth=0.8, zorder=3)
  212. ax_bar.set_xticks(x_centers)
  213. ax_bar.set_xticklabels([f'Level {l}' for l in levels], fontsize=10, fontweight='bold')
  214. # --- 折线 ---
  215. def plot_line(ax, solved_dict, others_dict, color, label_s, label_o,
  216. marker_s='o', marker_o='s'):
  217. s_pts = [(x_centers[j], v)
  218. for j, (l, v) in enumerate(zip(levels, [solved_dict.get(l) for l in levels]))
  219. if v is not None]
  220. o_pts = [(x_centers[j], v)
  221. for j, (l, v) in enumerate(zip(levels, [others_dict.get(l) for l in levels]))
  222. if v is not None]
  223. h1 = h2 = None
  224. if s_pts:
  225. xs, vs = zip(*s_pts)
  226. h1, = ax.plot(xs, vs, color=color, linestyle='-', linewidth=1.8,
  227. marker=marker_s, markersize=6, markeredgecolor='white',
  228. markeredgewidth=0.8, label=label_s, zorder=5)
  229. if o_pts:
  230. xo, vo = zip(*o_pts)
  231. h2, = ax.plot(xo, vo, color=color, linestyle='--', linewidth=1.8,
  232. marker=marker_o, markersize=6, markeredgecolor='white',
  233. markeredgewidth=0.8, label=label_o, zorder=5)
  234. return h1, h2
  235. h_ctx_s, h_ctx_o = plot_line(ax_ctx, avg_context_len_solved, avg_context_length_others,
  236. '#1A6FAF', 'Context Length (Solved)', 'Context Length (Failed)')
  237. ax_ctx.set_ylabel('Avg. Context Length', fontsize=11, color='black', labelpad=6, fontweight='bold')
  238. ax_ctx.tick_params(axis='y', labelcolor='black', labelsize=9)
  239. ax_ctx.spines['left'].set_edgecolor('black')
  240. def format_k(x, pos):
  241. return f'{x / 1000:g}k' if x >= 1000 else f'{x:g}'
  242. ax_ctx.yaxis.set_major_formatter(FuncFormatter(format_k))
  243. h_ep_s, h_ep_o = plot_line(ax_epoch, avg_epoch_solved, avg_epoch_others,
  244. '#C0392B', 'Avg. Epoch (Solved)', 'Avg. Epoch (Failed)',
  245. marker_s='^', marker_o='v')
  246. ax_epoch.set_ylabel('Avg. Epoch', fontsize=11, color='black', labelpad=6, fontweight='bold')
  247. ax_epoch.tick_params(axis='y', colors='black', labelsize=9)
  248. ax_epoch.spines['right'].set_edgecolor('black')
  249. # --- 图例 ---
  250. tool_patches = [mpatches.Patch(facecolor=tool_colors[i], edgecolor='#555555',
  251. linewidth=0.5, label=tool_keys[i])
  252. for i in range(n_tools)]
  253. line_handles = [h for h in [h_ctx_s, h_ctx_o, h_ep_s, h_ep_o] if h is not None]
  254. all_handles = tool_patches + line_handles
  255. n_cols = 5
  256. ordered_handles = [h for i in range(n_cols) for h in all_handles[i::n_cols]]
  257. legend = ax_bar.legend(
  258. handles=ordered_handles,
  259. fontsize=8,
  260. loc='lower left',
  261. bbox_to_anchor=(0, 1.05, 1, 0.1),
  262. mode="expand",
  263. ncol=n_cols,
  264. framealpha=0.9,
  265. edgecolor='#CCCCCC',
  266. borderpad=0.6,
  267. borderaxespad=0.
  268. )
  269. for text in legend.get_texts():
  270. text.set_fontweight('bold')
  271. # 如果图例有标题,也加粗
  272. if legend.get_title():
  273. legend.get_title().set_fontweight('bold')
  274. plt.tight_layout()
  275. plt.savefig('../../outputs/fig-statistics.pdf', bbox_inches='tight')
  276. plt.show()
  277. def draw_table(level=6, span=2, latex=True, show_complete=False):
  278. filenames = {
  279. 'Backward-DFS': 'log_pssr_formalgeo7k-bw-dfs.json', # symbolic solver
  280. 'Backward-RS': 'log_pssr_formalgeo7k-bw-rs.json',
  281. 'Backward-BFS': 'log_pssr_formalgeo7k-bw-bfs.json',
  282. 'Forward-DFS': 'log_pssr_formalgeo7k-fw-dfs.json',
  283. 'Forward-BFS': 'log_pssr_formalgeo7k-fw-bfs.json',
  284. 'Forward-RS': 'log_pssr_formalgeo7k-fw-rs.json',
  285. 'Kimi-K2': 'log_pssr_kimi-k2.json', # neural solver
  286. 'DeepSeek v3': 'log_pssr_deepseek-v3.json',
  287. 'GPT-5 mini': [64.79, 74.11, 63.30, 64.66, 53.50, 53.23, 41.46],
  288. 'Qwen3-VL': [65.93, 74.53, 65.43, 72.18, 50.96, 41.94, 36.67],
  289. 'Doubao seed 1.8': [69.14, 74.11, 69.15, 71.43, 64.33, 50.00, 51.67],
  290. 'GPT-5.2': [73.14, 80.38, 73.40, 74.81, 63.06, 59.68, 46.67],
  291. 'Claude4.5 Sonnet': [75.79, 84.55, 73.94, 76.32, 67.52, 64.52, 48.33],
  292. 'T5-small': 'log_pssr_t5-small_bs20_timeout600.json', # neural-symbolic solver (training-based)
  293. 'BART-base': 'log_pssr_bart-base_bs20_timeout600.json',
  294. 'Inter-GPS': 'log_pssr_intergps.json',
  295. 'DualGeoSolver': 'log_pssr_dualgeosolver_bs10_timeout600.json',
  296. 'NGS': 'log_pssr_ngs_bs10_timeout600.json',
  297. 'FGeo-DRL': 'log_pssr_fgeodrl.json',
  298. 'FGeo-TP': [80.86, 96.43, 85.44, 76.12, 62.26, 48.88, 29.55],
  299. 'FGeo-ISRL': 'log_pssr_res_bdrl.json',
  300. 'HyperGNet': 'log_pssr_hypergnet_TTT_bs5_gb_tm600.json',
  301. 'NSS': 'log_pssr_nss_FFFF_bs5_tm600.json',
  302. 'Pri-TPG': [89.29, 99.16, 96.28, 87.92, 77.07, 66.13, 30.00], # neural-symbolic solver (training-free)
  303. 'Ours': 'log_pssr_agent.json'
  304. }
  305. last_methods = ["Forward-RS", "Claude4.5 Sonnet", "NSS", 'Ours']
  306. problem_level = {} # map problem_id to level
  307. level_map = {} # map t_length to level (start from 0)
  308. for i in range(level):
  309. for j in range(span):
  310. level_map[i * span + j + 1] = i + 1
  311. save_json({'info': 'map theorem_length to problem level.', 'map': level_map},
  312. '../../outputs/log/log_level_map.json')
  313. for pid in range(7000):
  314. pid += 1
  315. t_length = len(load_json(f'../../datasets/problems/{pid}.json')['theorem_seqs'])
  316. problem_level[pid] = level_map[t_length] if t_length <= level * span else level
  317. method_name_max_len = max([len(m) for m in filenames.keys()] + [6]) + 1
  318. outputs = []
  319. if not show_complete:
  320. head = ['Method' + "".join([" "] * (method_name_max_len - 6)),
  321. 'Total', 'L1 ', 'L2 ', 'L3 ', 'L4 ', 'L5 ', 'L6 ']
  322. line = ''.join(['-'] * (7 * 8 + method_name_max_len))
  323. else:
  324. head = ['Method' + "".join([" "] * (method_name_max_len - 6)),
  325. ' A ', ' T ', 'Total', 'L1 ', 'L2 ', 'L3 ', 'L4 ', 'L5 ', 'L6 ']
  326. line = ''.join(['-'] * (9 * 8 + method_name_max_len))
  327. if latex:
  328. print(' & '.join(head))
  329. outputs.append(' & '.join(head))
  330. else:
  331. print(' | '.join(head))
  332. outputs.append(' | '.join(head))
  333. print(line)
  334. outputs.append(line)
  335. for method in filenames.keys(): # pssr_log
  336. lines = [method + "".join([" "] * (method_name_max_len - len(method)))]
  337. if isinstance(filenames[method], list):
  338. lines.extend([' - ', ' - '])
  339. for r in filenames[method]:
  340. lines.append(str(r))
  341. lines[-1] = lines[-1] + ' ' * (5 - len(lines[-1]))
  342. else:
  343. pssr_log = load_json(f"../../outputs/log/{filenames[method]}")
  344. GT = (len(pssr_log["solved"]) + len(pssr_log["unsolved"]) + # 事实求解成功率,分母为已求解的题目
  345. len(pssr_log["timeout"]) + len(pssr_log["error"]))
  346. lines.append(str(round(GT / len(pssr_log["total"]) * 100, 2)))
  347. lines[-1] = lines[-1] + ' ' * (5 - len(lines[-1]))
  348. lines.append(str(round(len(pssr_log["solved"]) / GT * 100, 2)))
  349. lines[-1] = lines[-1] + ' ' * (5 - len(lines[-1]))
  350. total_level_count = [0 for _ in range(level + 1)] # [total, l1, l2, ...]
  351. solved_level_count = [0 for _ in range(level + 1)]
  352. for pid in pssr_log["total"]:
  353. total_level_count[0] += 1
  354. total_level_count[problem_level[pid]] += 1
  355. if str(pid) in pssr_log["solved"]:
  356. solved_level_count[0] += 1
  357. solved_level_count[problem_level[pid]] += 1
  358. # print()
  359. # print(total_level_count)
  360. # print(solved_level_count)
  361. for i in range(level + 1):
  362. if total_level_count[i] == 0:
  363. lines.append('Nan')
  364. else:
  365. lines.append(str(round(solved_level_count[i] / total_level_count[i] * 100, 2)))
  366. lines[-1] = lines[-1] + ' ' * (5 - len(lines[-1]))
  367. if not show_complete:
  368. lines = [lines[0]] + lines[3:]
  369. if latex:
  370. print(' & '.join(lines))
  371. outputs.append(' & '.join(lines))
  372. else:
  373. print(' | '.join(lines))
  374. outputs.append(' | '.join(lines))
  375. if method in last_methods:
  376. print(line)
  377. outputs.append(line)
  378. with open('../../outputs/tab-main_results.txt', 'w', encoding='utf-8') as f:
  379. f.write('\n'.join(outputs))
  380. def lmm_call_statistic():
  381. data = {'solved': [], 'unsolved': []}
  382. log = load_json('../../outputs/log/log_pssr_agent.json')
  383. for filename in os.listdir('../../outputs/agent'):
  384. count = 0
  385. for history in load_json(f'../../outputs/agent/{filename}')['history']:
  386. for msg in history:
  387. if msg['role'] == 'assistant':
  388. count += 1
  389. pid = filename.split('.')[0].split('_')[-1]
  390. if pid in log['solved']:
  391. data['solved'].append(count)
  392. else:
  393. data['unsolved'].append(count)
  394. print('solved', sum(data['solved']) / len(data['solved']))
  395. print('unsolved', sum(data['unsolved']) / len(data['unsolved']))
  396. if __name__ == '__main__':
  397. draw_figure()
  398. draw_table()
  399. lmm_call_statistic()