feat: add yolox algo
This commit is contained in:
@@ -0,0 +1,56 @@
|
||||
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()
|
||||
Reference in New Issue
Block a user