import argparse import csv import matplotlib.pyplot as plt import numpy as np parser = argparse.ArgumentParser(description="绘制训练曲线") parser.add_argument("--csv", default="runs/classify/train/results.csv", help="训练结果 CSV 路径") parser.add_argument("--output", default="训练曲线.png", help="输出图片路径") args = parser.parse_args() epochs = [] train_loss = [] val_loss = [] val_acc = [] with open(args.csv, "r", encoding="utf-8") as f: reader = csv.DictReader(f) for row in reader: epochs.append(int(row["epoch"])) train_loss.append(float(row["train/loss"])) val_loss.append(float(row["val/loss"])) val_acc.append(float(row["metrics/accuracy_top1"])) avg_acc = np.mean(val_acc) print("=" * 60) print(f"最终预测准确率:{val_acc[-1]:.2%}") print(f"最高准确率:{max(val_acc):.2%}") print(f"平均准确率:{avg_acc:.2%}") print("=" * 60) plt.rcParams["font.sans-serif"] = ["SimHei"] plt.rcParams["axes.unicode_minus"] = False plt.figure(figsize=(10, 4)) plt.subplot(121) plt.plot(epochs, train_loss, label="训练损失", linewidth=2) plt.plot(epochs, val_loss, label="验证损失", linewidth=2) plt.title("损失函数曲线") plt.xlabel("Epoch") plt.ylabel("Loss") plt.legend() plt.grid(True) plt.subplot(122) plt.plot(epochs, val_acc, label="准确率", color="#2ecc71", linewidth=2) plt.title("准确率曲线") plt.xlabel("Epoch") plt.ylabel("Accuracy") plt.legend() plt.grid(True) plt.tight_layout() plt.savefig(args.output, dpi=300) plt.show()