feat: add yolox algo

This commit is contained in:
2026-08-05 21:19:01 +08:00
parent b399d94de3
commit 79fde03972
13 changed files with 368 additions and 2 deletions
+68
View File
@@ -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)