Files
2026-08-05 21:19:01 +08:00

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()