如何在Keras绘制学习曲线时自定义图例?
修改TensorFlow/Keras训练曲线图例的正确方法
问题描述
使用TensorFlow 2/Keras训练CNN模型后,通过history属性绘制准确率和损失曲线,当前图例显示为accuracy/val_accuracy,希望修改为train_accuracy/validation_accuracy。尝试仅修改plt.legend的第二个参数无效,使用代理艺术家的方法也失败并出现警告。
解决方案
方法一:直接指定完整图例文本
你之前的问题在于只修改了图例的第二个标签,第一个标签仍沿用原string变量值(比如"accuracy")。只需同时修改两个标签的文本即可:
import matplotlib.pyplot as plt def plot_graphs(cnn_trained_model, string): plt.plot(cnn_trained_model.history[string]) plt.plot(cnn_trained_model.history['val_'+string]) plt.xlabel("Epochs") plt.ylabel(string.capitalize()) # 同时指定训练和验证的图例文本 plt.legend([f'train_{string}', f'validation_{string}']) plt.show() # 绘制准确率和损失曲线 plot_graphs(cnn_trained_model, "accuracy") plot_graphs(cnn_trained_model, "loss")
方法二:为曲线添加label参数(更规范)
给每条绘制的曲线添加label属性,再直接调用plt.legend()自动识别标签,这种方式更易维护:
import matplotlib.pyplot as plt def plot_graphs(cnn_trained_model, string): # 为每条曲线设置label plt.plot(cnn_trained_model.history[string], label=f'train_{string}') plt.plot(cnn_trained_model.history['val_'+string], label=f'validation_{string}') plt.xlabel("Epochs") plt.ylabel(string.capitalize()) # 自动读取label生成图例 plt.legend() plt.show() # 绘制准确率和损失曲线 plot_graphs(cnn_trained_model, "accuracy") plot_graphs(cnn_trained_model, "loss")
关于代理艺术家失败的说明
你尝试的代理艺术家方法完全没必要,因为matplotlib本身支持直接通过label参数或legend的文本列表设置图例,无需额外创建代理对象。之前的警告是因为传入的Line2D对象使用方式错误,直接用上述两种方法即可解决问题。
内容的提问来源于stack exchange,提问作者Bluetail
相关产品推荐
相关产品推荐

