当前位置:首页>python>Ccy:别再只画折线图了,用 Python 做一张“模型鲁棒性相变图”

Ccy:别再只画折线图了,用 Python 做一张“模型鲁棒性相变图”

  • 2026-10-11 06:21:33
Ccy:别再只画折线图了,用 Python 做一张“模型鲁棒性相变图”

在做模型鲁棒性、测试时自适应、退化恢复、可靠性分析时,我们经常会遇到一个问题:

模型性能不是简单地随某个变量下降,而是在多个测试时压力共同作用下,呈现出“稳定—退化—崩溃—恢复”的复杂状态变化。

如果只用普通折线图,很难表达这种二维测试条件下的性能分布。

我们这里构建一个实用的 Python 可视化示例:绘制一张 Robustness Phase Diagram,即鲁棒性相变图。它可以直观展示模型在不同测试时压力下的性能状态,以及自适应策略是否能把模型从退化或崩溃区域“拉回来”。

先放结果,中间是介绍,文末是代码:

1. 这张图想表达什么?

我们考虑两个测试时压力轴:

  • 横轴:测试时压力 1,例如 domain mismatch、sensor drift;
  • 纵轴:测试时压力 2,例如 observation sparsity、resource constraint。

在这两个压力共同作用下,模型可能处于四种状态:

区域
含义
Robust region
模型性能稳定,处于高性能状态
Degraded region
模型性能下降,但仍可使用
Collapse region
模型性能低于可接受阈值,进入崩溃状态
Recovery region
原始性能较差,但经过自适应策略后恢复到较高性能

相比只画 accuracy 曲线,这类图更适合表达:

  • 多因素联合退化;
  • 模型性能边界;
  • 崩溃阈值;
  • 自适应策略的恢复能力;
  • 从高退化输入到恢复状态的过程轨迹。

2.图中各区域如何解读?

最终图可以这样理解:

Robust region

模型在低测试时压力下保持稳定高性能。这通常对应接近训练分布、传感器稳定、观测充分的情况。

Degraded region

模型性能下降,但尚未完全失效。这一区域通常是鲁棒性方法最需要关注的部分,因为模型仍有恢复空间。

Collapse region

模型性能低于可接受阈值。在实际任务中,这可能意味着分类错误率急剧升高、检测漏检严重、分割结果不可用,或者回归误差超过任务容忍范围。

Recovery region

这是最值得关注的区域。它表示原始模型本来已经退化,但经过自适应或恢复策略后,性能重新回到高水平。

3.这张图适合放在哪里?

这类图适合用于:

  • 论文 method motivation;
  • robustness analysis;
  • ablation study;
  • test-time adaptation 可视化;
  • 复杂退化实验结果总结;
  • PPT 中解释方法为什么有效;
  • 展示模型从崩溃区域恢复到稳定区域的过程。

尤其是当你的方法不是单纯提高平均性能,而是强调“在复杂测试条件下安全恢复模型表现”时,这张图会比普通柱状图更有解释力。


4.小结

代码本质上完成了三件事:

  1. 构造二维测试时压力空间;
  2. 根据模型性能划分鲁棒、退化、崩溃和恢复区域;
  3. 用边界线、等高线和轨迹展示模型状态变化。

相比普通折线图,鲁棒性相变图能更清楚地回答几个问题:

  • 模型在哪些测试条件下仍然稳定?
  • 模型什么时候开始明显退化?
  • 模型在哪些区域会崩溃?
  • 自适应策略到底救回了哪些样本?
  • 恢复过程是否具有清晰的方向性?

对于鲁棒性、测试时自适应、模型恢复、序列决策相关研究,这种图是一种很实用的可视化方式。

如果你正在做相关实验,可以尝试把自己的真实性能矩阵替换进去,也许会比单纯的 accuracy 表格更有说服力。

import numpy as npimport matplotlib.pyplot as pltfrom matplotlib.colors import ListedColormap, BoundaryNormfrom matplotlib.patches import Patchnp.random.seed(7)x = np.linspace(0, 1, 260)y = np.linspace(0, 1, 260)X, Y = np.meshgrid(x, y)raw_task_perf = (    92    - 18 * X**1.6    - 22 * Y**1.8    - 35 * (X * Y)**1.2    + 2.5 * np.sin(3 * np.pi * X) * np.cos(2 * np.pi * Y))adaptation_gain = (    18 * np.exp(-((X - 0.62)**2 / 0.055 + (Y - 0.52)**2 / 0.075))    + 10 * np.exp(-((X - 0.78)**2 / 0.035 + (Y - 0.30)**2 / 0.045)))adapted_task_perf = raw_task_perf + adaptation_gainadapted_task_perf = np.clip(adapted_task_perf, 35, 94)phase = np.zeros_like(adapted_task_perf, dtype=int)phase[adapted_task_perf < 68] = 2phase[(adapted_task_perf >= 68) & (adapted_task_perf < 82)] = 1phase[adapted_task_perf >= 82] = 0recovery_mask = (    (raw_task_perf < 78)    & (adapted_task_perf >= 82)    & (adaptation_gain > 6))phase[recovery_mask] = 3plt.rcParams["font.family"] = "DejaVu Sans"plt.rcParams["axes.unicode_minus"] = Falsefig, ax = plt.subplots(figsize=(10.8, 7.2), dpi=150)phase_colors = [    "#D7F2E3",  # Robust    "#FFF1B8",  # Degraded    "#F6B6B6",  # Collapse    "#BFD7FF",  # Recovery]cmap = ListedColormap(phase_colors)norm = BoundaryNorm([-0.5, 0.5, 1.5, 2.5, 3.5], cmap.N)ax.pcolormesh(X, Y, phase, cmap=cmap, norm=norm, shading="auto", alpha=0.96)contours = ax.contour(    X, Y, adapted_task_perf,    levels=[60, 68, 75, 82, 88],    colors="k",    linewidths=[1.0, 1.8, 1.0, 1.8, 1.0],    alpha=0.62)ax.clabel(contours, inline=True, fontsize=9, fmt="%d")ax.contour(    X, Y, adapted_task_perf,    levels=[68],    colors="#7A1F1F",    linewidths=2.8,    linestyles="--")ax.contour(    X, Y, adapted_task_perf,    levels=[82],    colors="#1C6B43",    linewidths=2.8,    linestyles="--")ax.contour(    X, Y, adaptation_gain,    levels=[6],    colors="#2356A4",    linewidths=2.4,    linestyles="-.")trajectory = np.array([    [0.82, 0.76],    [0.72, 0.66],    [0.61, 0.56],    [0.49, 0.43],    [0.34, 0.28],])ax.plot(    trajectory[:, 0],    trajectory[:, 1],    color="#222222",    linewidth=2.4,    marker="o",    markersize=5.5,    zorder=5)for i in range(len(trajectory) - 1):    ax.annotate(        "",        xy=trajectory[i + 1],        xytext=trajectory[i],        arrowprops=dict(            arrowstyle="->",            color="#222222",            lw=2.0,            shrinkA=5,            shrinkB=5        ),        zorder=6    )ax.text(    trajectory[0, 0] + 0.02,    trajectory[0, 1] + 0.025,    "high-shift input",    fontsize=10,    weight="bold")ax.text(    trajectory[-1, 0] - 0.055,    trajectory[-1, 1] - 0.055,    "adapted state",    fontsize=10,    weight="bold")ax.scatter(    [0.08], [0.08],    s=170,    marker="*",    color="#1B7F4C",    edgecolor="white",    linewidth=1.4,    zorder=7)ax.text(    0.105, 0.085,    "nominal state",    fontsize=10,    weight="bold",    va="center")ax.text(    0.16, 0.18,    "Robust\nregion",    fontsize=16,    weight="bold",    color="#145A32",    ha="center",    va="center")ax.text(    0.42, 0.55,    "Degraded\nregion",    fontsize=15,    weight="bold",    color="#8A6D00",    ha="center",    va="center")ax.text(    0.82, 0.88,    "Collapse\nregion",    fontsize=15,    weight="bold",    color="#7A1F1F",    ha="center",    va="center")ax.text(    0.66, 0.42,    "Recovery\nregion",    fontsize=15,    weight="bold",    color="#1F4E8C",    ha="center",    va="center")legend_elements = [    Patch(facecolor=phase_colors[0], edgecolor="none", label="Robust region"),    Patch(facecolor=phase_colors[1], edgecolor="none", label="Degraded region"),    Patch(facecolor=phase_colors[2], edgecolor="none", label="Collapse region"),    Patch(facecolor=phase_colors[3], edgecolor="none", label="Recovery region"),]line_legend = [    plt.Line2D([0], [0], color="#1C6B43", lw=2.6, ls="--", label="Robust boundary"),    plt.Line2D([0], [0], color="#7A1F1F", lw=2.6, ls="--", label="Collapse boundary"),    plt.Line2D([0], [0], color="#2356A4", lw=2.4, ls="-.", label="Recovery frontier"),    plt.Line2D([0], [0], color="#222222", lw=2.4, marker="o", label="Adaptation trajectory"),]ax.legend(    handles=legend_elements + line_legend,    loc="upper left",    bbox_to_anchor=(1.08, 1.0),    frameon=True,    framealpha=0.94,    fontsize=9.5,    borderpad=0.9)ax.set_title(    "Robustness Phase Diagram: Performance Phase Transition under Compound Test-Time Shifts",    fontsize=15,    weight="bold",    pad=14)ax.set_xlabel("Test-time stress axis 1  |  domain mismatch / sensor drift", fontsize=11)ax.set_ylabel("Test-time stress axis 2  |  observation sparsity / resource constraint", fontsize=11)ax.set_xlim(0, 1)ax.set_ylim(0, 1)ax.set_xticks(np.linspace(0, 1, 6))ax.set_yticks(np.linspace(0, 1, 6))ax.grid(    color="white",    linestyle="-",    linewidth=0.8,    alpha=0.65)ax.set_aspect("equal", adjustable="box")ax.text(    0.02, -0.12,    "Contours indicate adapted downstream task performance. "    "Dashed lines mark phase-transition boundaries.",    transform=ax.transAxes,    fontsize=9.5,    color="#444444")plt.tight_layout(rect=[0, 0.03, 0.78, 1])plt.show()

最新文章

随机文章