如何可视化Python字典或NumPy二维数组格式的SARSA Q表
Q表可视化解决方案
前置说明
以下方案基于matplotlib实现,可直接在Jupyter Notebook中运行,支持定期输出训练中间结果,提供两种常用呈现样式,你可以按需选择:
- 样式1:每个网格拆分为4个三角形,分别对应四个方向的Q值(匹配你提到的参考样式)
- 样式2:网格热图叠加最优动作箭头,更直观呈现当前训练得到的策略
步骤1:字典格式Q表转NumPy数组
首先把你的字典格式Q表转换为形状为(网格行数, 网格列数, 4)的三维数组,最后一维依次对应U、R、D、L四个动作:
import numpy as np def q_dict_to_array(q_dict, n_rows, n_cols): # 动作到索引的映射 action_map = {'U': 0, 'R': 1, 'D': 2, 'L': 3} # 初始化Q数组,默认填充nan处理未探索的状态动作对 q_array = np.full((n_rows, n_cols, 4), np.nan) for (state, action), value in q_dict.items(): # 假设状态编号按行优先排列,即第一行从左到右是0,1,..n_cols-1,第二行是n_cols,...2n_cols-1 row = state // n_cols col = state % n_cols q_array[row, col, action_map[action]] = value return q_array
如果你的状态编号规则不是行优先,自行修改row和col的计算逻辑即可。
步骤2:可视化实现
样式1:四三角形拆分显示所有方向Q值
import matplotlib.pyplot as plt from matplotlib.patches import Polygon def plot_q_table_split(q_array, title="Q表可视化"): n_rows, n_cols, _ = q_array.shape fig, ax = plt.subplots(figsize=(n_cols*1.2, n_rows*1.2)) # 配色方案,红蓝渐变,可自行更换cmap cmap = plt.cm.coolwarm vmin = np.nanmin(q_array) vmax = np.nanmax(q_array) # 每个格子四个三角形的顶点坐标 triangles = [ # U: 上三角 [[0,1], [1,1], [0.5, 0.5]], # R: 右三角 [[1,1], [1,0], [0.5, 0.5]], # D: 下三角 [[1,0], [0,0], [0.5, 0.5]], # L: 左三角 [[0,0], [0,1], [0.5, 0.5]] ] text_offsets = [(0.5, 0.85), (0.85, 0.5), (0.5, 0.15), (0.15, 0.5)] for row in range(n_rows): for col in range(n_cols): for a in range(4): val = q_array[row, col, a] # 计算三角形的实际坐标 poly_coords = np.array(triangles[a]) + [col, n_rows - 1 - row] poly = Polygon(poly_coords, facecolor=cmap((val - vmin)/(vmax - vmin)) if not np.isnan(val) else 'white', edgecolor='black') ax.add_patch(poly) # 标注Q值,保留2位小数 if not np.isnan(val): ax.text(col + text_offsets[a][0], n_rows - 1 - row + text_offsets[a][1], f"{val:.2f}", ha='center', va='center', fontsize=8) # 设置坐标轴 ax.set_xlim(0, n_cols) ax.set_ylim(0, n_rows) ax.set_xticks(np.arange(0.5, n_cols, 1)) ax.set_xticklabels(np.arange(n_cols)) ax.set_yticks(np.arange(0.5, n_rows, 1)) ax.set_yticklabels(np.arange(n_rows)[::-1]) ax.grid(color='black') ax.set_title(title) plt.tight_layout() plt.show()
样式2:热图叠加最优动作箭头(推荐用于训练过程策略观察)
这个样式更简洁,能快速看到每个状态的最优选择:
def plot_q_table_policy(q_array, title="Q表+最优策略可视化"): n_rows, n_cols, _ = q_array.shape # 计算每个状态的最大Q值和最优动作 max_q = np.nanmax(q_array, axis=2) best_action = np.nanargmax(q_array, axis=2) # 箭头映射 arrow_map = {0: '↑', 1: '→', 2: '↓', 3: '←'} fig, ax = plt.subplots(figsize=(n_cols*1.2, n_rows*1.2)) # 绘制热图 im = ax.imshow(max_q, cmap='coolwarm', vmin=np.nanmin(max_q), vmax=np.nanmax(max_q)) # 标注数值和箭头 for i in range(n_rows): for j in range(n_cols): if not np.isnan(max_q[i,j]): text = ax.text(j, i, f"{max_q[i,j]:.2f}\n{arrow_map[best_action[i,j]]}", ha='center', va='center', color='black', fontsize=10) # 设置坐标轴 ax.set_xticks(np.arange(n_cols)) ax.set_yticks(np.arange(n_rows)) ax.set_title(title) plt.colorbar(im, ax=ax, shrink=0.8) plt.tight_layout() plt.show()
训练过程调用示例
你在SARSA的训练循环中每1万轮调用一次即可,假设你的网格是4行4列:
total_episodes = 100000 for episode in range(total_episodes): # 你的SARSA训练逻辑 # ... # 每1万轮可视化一次 if (episode + 1) % 10000 == 0: q_arr = q_dict_to_array(Q, n_rows=4, n_cols=4) # 两种样式二选一即可 plot_q_table_split(q_arr, title=f"Q表(训练轮次:{episode+1})") # plot_q_table_policy(q_arr, title=f"Q表+策略(训练轮次:{episode+1})")
可调参数说明
- 可以通过修改
figsize调整图表大小 - 更换
cmap参数修改配色,常用可选RdBu、viridis等 - 调整
fontsize修改文字大小 - 数值保留位数可以修改格式化字符串里的
.2f为你需要的精度
内容的提问来源于stack exchange,提问作者waffledood
相关产品推荐
相关产品推荐

