57 lines
1.5 KiB
Python
57 lines
1.5 KiB
Python
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()
|