feat: add yolox algo
This commit is contained in:
@@ -0,0 +1,68 @@
|
||||
import argparse
|
||||
import os
|
||||
import shutil
|
||||
import csv
|
||||
from ultralytics import YOLO
|
||||
|
||||
parser = argparse.ArgumentParser(description="YOLO26 岩性分类评估")
|
||||
parser.add_argument("--model", default="runs/classify/train/weights/best.pt", help="模型权重路径")
|
||||
parser.add_argument("--val-root", default="custom_dataset/val", help="验证集路径")
|
||||
parser.add_argument("--save-root", default="岩性分类结果", help="分类结果保存目录")
|
||||
parser.add_argument("--csv", default="分类预测结果.csv", help="预测结果 CSV 输出路径")
|
||||
parser.add_argument("--device", default="cuda", help="推理设备")
|
||||
args = parser.parse_args()
|
||||
|
||||
model = YOLO(args.model)
|
||||
|
||||
all_images = []
|
||||
for cls_folder in os.listdir(args.val_root):
|
||||
cls_path = os.path.join(args.val_root, cls_folder)
|
||||
if os.path.isdir(cls_path):
|
||||
for img_name in os.listdir(cls_path):
|
||||
all_images.append(os.path.join(cls_path, img_name))
|
||||
|
||||
total_samples = 0
|
||||
correct_samples = 0
|
||||
|
||||
with open(args.csv, "w", newline="", encoding="utf-8-sig") as f:
|
||||
writer = csv.writer(f)
|
||||
writer.writerow(["图片名", "真实标签", "预测类别", "置信度"])
|
||||
|
||||
for img_path in all_images:
|
||||
res = model(img_path, verbose=False)[0]
|
||||
pred_class = res.names[res.probs.top1]
|
||||
conf = f"{res.probs.top1conf:.2%}"
|
||||
|
||||
true_class = os.path.basename(os.path.dirname(img_path))
|
||||
img_name = os.path.basename(img_path)
|
||||
|
||||
writer.writerow([img_name, true_class, pred_class, conf])
|
||||
|
||||
target_dir = os.path.join(args.save_root, true_class)
|
||||
os.makedirs(target_dir, exist_ok=True)
|
||||
shutil.copy(img_path, os.path.join(target_dir, img_name))
|
||||
|
||||
total_samples += 1
|
||||
|
||||
is_correct = False
|
||||
if true_class == "泥页岩" and ("泥" in pred_class or "页" in pred_class):
|
||||
is_correct = True
|
||||
elif true_class == "砂砾岩" and ("砂" in pred_class or "粉砂" in pred_class):
|
||||
is_correct = True
|
||||
elif true_class == "灰岩" and ("灰" in pred_class or "石灰岩" in pred_class):
|
||||
is_correct = True
|
||||
elif true_class == "云岩" and ("云" in pred_class or "白云" in pred_class):
|
||||
is_correct = True
|
||||
elif true_class == "特殊岩" and ("石膏" in pred_class or "煤" in pred_class or "膏盐" in pred_class):
|
||||
is_correct = True
|
||||
|
||||
if is_correct:
|
||||
correct_samples += 1
|
||||
|
||||
accuracy = (correct_samples / total_samples) * 100
|
||||
print("\n" + "=" * 60)
|
||||
print("分类完成!准确率统计结果:")
|
||||
print(f"总测试样本数:{total_samples} 张")
|
||||
print(f"预测正确样本数:{correct_samples} 张")
|
||||
print(f"整体分类准确率:{accuracy:.2f}%")
|
||||
print("=" * 60)
|
||||
Reference in New Issue
Block a user