You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何可视化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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.10.04 01:06:03