27 lines
938 B
Python
27 lines
938 B
Python
import argparse
|
|
from ultralytics import YOLO
|
|
import matplotlib.pyplot as plt
|
|
|
|
parser = argparse.ArgumentParser(description="YOLO26 单张图片岩性分类预测")
|
|
parser.add_argument("--image", required=True, help="待预测的图片路径")
|
|
parser.add_argument("--model", default="runs/classify/train/weights/best.pt", help="模型权重路径")
|
|
parser.add_argument("--output", default="rock_result.png", help="结果图片保存路径")
|
|
parser.add_argument("--device", default="cuda", help="推理设备")
|
|
args = parser.parse_args()
|
|
|
|
model = YOLO(args.model)
|
|
|
|
results = model(args.image, device=args.device)
|
|
|
|
print("=" * 50)
|
|
print("岩石分类预测结果")
|
|
print(f"图片:{args.image}")
|
|
print(f"类别:{results[0].names[results[0].probs.top1]}")
|
|
print(f"置信度:{results[0].probs.top1conf:.2%}")
|
|
print("=" * 50)
|
|
|
|
plt.imshow(results[0].plot())
|
|
plt.axis("off")
|
|
plt.savefig(args.output, dpi=300, bbox_inches="tight")
|
|
plt.show()
|