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

使用VisionClassifierTrainer获取验证准确率及绘制准确率-epoch曲线的问题

如何用VisionClassifierTrainer获取验证准确率并绘制曲线

1. 开启每轮验证跟踪

VisionClassifierTrainer默认不会在每轮训练后自动评估验证集,需要初始化时添加几个关键参数,指定评估策略和指标:

修改后的初始化代码:

from hugsvision.nnet.VisionClassifierTrainer import VisionClassifierTrainer
from transformers import ViTFeatureExtractor, ViTForImageClassification

trainer = VisionClassifierTrainer(
    model_name   = "MyKvasirV2Model11",
    train        = train,
    test         = test,
    output_dir   = "/content/drive/MyDrive/Untitled Folder",
    max_epochs   = 2,
    batch_size   = 50,
    lr           = 2e-5,
    # 添加验证相关配置
    evaluation_strategy="epoch",  # 每轮epoch结束后运行验证
    metric_for_best_model="accuracy",  # 以准确率作为最优模型评估指标
    load_best_model_at_end=True,  # 可选:训练结束后自动加载最优模型
    model = ViTForImageClassification.from_pretrained(
        huggingface_model,
        num_labels = len(label2id),
        label2id   = label2id,
        id2label   = id2label
    ),
    feature_extractor = ViTFeatureExtractor.from_pretrained(
        huggingface_model,
    ),
)

2. 提取每轮验证准确率

启动训练后,通过trainer.state.log_history可以获取所有训练日志,从中筛选出每轮的验证准确率:

# 开始训练
trainer.train()

# 提取验证准确率数据
val_accuracies = []
epochs = []
for log_entry in trainer.state.log_history:
    if "eval_accuracy" in log_entry:
        val_accuracies.append(log_entry["eval_accuracy"])
        epochs.append(log_entry["epoch"])

# 打印结果
print("每轮验证准确率:", val_accuracies)

3. 绘制验证准确率Vs epoch曲线

用matplotlib基于提取的数据绘制曲线:

import matplotlib.pyplot as plt

plt.figure(figsize=(8, 5))
plt.plot(epochs, val_accuracies, marker='o', color='#2ecc71', label='验证准确率')
plt.xlabel('训练轮次(Epoch)')
plt.ylabel('准确率(Accuracy)')
plt.title('验证准确率随训练轮次变化曲线')
plt.legend()
plt.grid(alpha=0.3)
plt.show()

额外提示

  • 确保test变量对应的是验证集,不要用测试集替代验证集。如果没有单独的验证集,可以用trainer.split_train_test(train, test_size=0.2)拆分训练集得到验证子集。
  • log_history中还包含训练损失、验证损失等其他指标,可按相同逻辑提取使用。

内容的提问来源于stack exchange,提问作者Anu Raj

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.22 19:13:33