goat / Scripts /eval_sahi.py
LightChuan's picture
Upload folder using huggingface_hub
6a5bb7e verified
Raw
History Blame Contribute Delete
6.82 kB
"""SAHI 切片推理评估 — 对比不同切片参数的mAP50-95
原理:
大图 (3200x1800) 切成重叠小块 → 每块独立检测 → WBF/NMS融合
小目标在切片中变大 → 定位精度显著提升
测试配置:
1. 基线: 全图推理 (当前最佳方式)
2. SAHI 640切片
3. SAHI 960切片
4. SAHI 1280切片
"""
import sys, os, time, json
import numpy as np
from pathlib import Path
os.chdir("/home/user/goat")
sys.path.insert(0, "/home/user/goat")
from ultralytics import YOLO
from sahi import AutoDetectionModel
from sahi.predict import get_sliced_prediction, get_prediction
from sahi.utils.coco import Coco, CocoAnnotation, CocoImage, CocoPrediction
import glob
from PIL import Image
def load_gt_boxes(lbl_path, img_w, img_h):
"""读取YOLO格式GT, 返回xyxy列表"""
boxes = []
if not os.path.exists(lbl_path):
return boxes
with open(lbl_path) as f:
for line in f:
parts = line.strip().split()
if len(parts) < 5:
continue
cls, cx, cy, w, h = int(parts[0]), float(parts[1]), float(parts[2]), float(parts[3]), float(parts[4])
x1 = (cx - w/2) * img_w
y1 = (cy - h/2) * img_h
x2 = (cx + w/2) * img_w
y2 = (cy + h/2) * img_h
boxes.append([x1, y1, x2, y2, cls])
return boxes
def compute_iou(box1, box2):
x1 = max(box1[0], box2[0]); y1 = max(box1[1], box2[1])
x2 = min(box1[2], box2[2]); y2 = min(box1[3], box2[3])
inter = max(0, x2-x1) * max(0, y2-y1)
a1 = (box1[2]-box1[0]) * (box1[3]-box1[1])
a2 = (box2[2]-box2[0]) * (box2[3]-box2[1])
return inter / (a1 + a2 - inter + 1e-7)
def compute_ap_at_iou(gt_boxes_all, pred_boxes_all, iou_threshold):
"""计算单个IoU阈值下的AP"""
all_preds = []
n_gt_total = 0
for img_idx, (gts, preds) in enumerate(zip(gt_boxes_all, pred_boxes_all)):
n_gt_total += len(gts)
gt_matched = [False] * len(gts)
for pred in preds:
best_iou, best_gt = 0, -1
for gi, gt in enumerate(gts):
iou = compute_iou(pred[:4], gt[:4])
if iou > best_iou:
best_iou = iou
best_gt = gi
if best_iou >= iou_threshold and best_gt >= 0 and not gt_matched[best_gt]:
all_preds.append((pred[4], 1, img_idx)) # (conf, tp, img)
gt_matched[best_gt] = True
else:
all_preds.append((pred[4], 0, img_idx)) # (conf, fp, img)
if n_gt_total == 0:
return 0.0
# 按置信度排序
all_preds.sort(key=lambda x: -x[0])
tp_cum, fp_cum = 0, 0
precisions, recalls = [], []
for conf, is_tp, _ in all_preds:
if is_tp:
tp_cum += 1
else:
fp_cum += 1
precisions.append(tp_cum / (tp_cum + fp_cum))
recalls.append(tp_cum / n_gt_total)
# COCO 101-point interpolation
ap = 0
for t in np.linspace(0, 1, 101):
p_at_r = 0
for p, r in zip(precisions, recalls):
if r >= t:
p_at_r = max(p_at_r, p)
ap += p_at_r / 101
return ap
def eval_config(model_path, val_img_dir, val_lbl_dir, slice_size=None, overlap_ratio=0.2, imgsz=1536, conf=0.25):
"""评估一种推理配置"""
detection_model = AutoDetectionModel.from_pretrained(
model_type="ultralytics",
model_path=model_path,
confidence_threshold=conf,
device="cuda:0",
)
imgs = sorted(glob.glob(os.path.join(val_img_dir, "*")))
gt_boxes_all = []
pred_boxes_all = []
t0 = time.time()
for idx, img_path in enumerate(imgs):
basename = os.path.splitext(os.path.basename(img_path))[0]
lbl_path = os.path.join(val_lbl_dir, basename + ".txt")
img = Image.open(img_path)
W, H = img.size
gts = load_gt_boxes(lbl_path, W, H)
gt_boxes_all.append(gts)
if slice_size:
result = get_sliced_prediction(
img_path,
detection_model,
slice_height=slice_size,
slice_width=slice_size,
overlap_height_ratio=overlap_ratio,
overlap_width_ratio=overlap_ratio,
perform_standard_pred=True, # 也做全图预测
postprocess_type="NMS",
postprocess_match_threshold=0.5,
verbose=0,
)
else:
result = get_prediction(
img_path,
detection_model,
verbose=0,
)
preds = []
for pred in result.object_prediction_list:
bb = pred.bbox
preds.append([bb.minx, bb.miny, bb.maxx, bb.maxy, pred.score.value])
pred_boxes_all.append(preds)
if (idx+1) % 20 == 0:
print(f" [{idx+1}/{len(imgs)}]...")
elapsed = time.time() - t0
# 计算mAP50-95
iou_thresholds = np.arange(0.5, 1.0, 0.05)
aps = []
ap_details = {}
for iou_t in iou_thresholds:
ap = compute_ap_at_iou(gt_boxes_all, pred_boxes_all, iou_t)
aps.append(ap)
ap_details[f"AP@{int(iou_t*100)}"] = ap
map50 = aps[0]
map5095 = np.mean(aps)
return map50, map5095, ap_details, elapsed
def main():
model_path = "runs/detect/Detection_experiments/v6_1_s_refined/weights/best.pt"
val_img_dir = "Data/Detection_dataset/images/val"
val_lbl_dir = "Data/Detection_dataset/labels/val"
configs = [
("全图推理 (baseline)", None, 0.2),
("SAHI 切片 640", 640, 0.25),
("SAHI 切片 960", 960, 0.25),
("SAHI 切片 1280", 1280, 0.2),
]
results = []
for name, slice_sz, overlap in configs:
print(f"\n{'='*50}")
print(f" 测试: {name}")
print(f"{'='*50}")
map50, map5095, details, elapsed = eval_config(
model_path, val_img_dir, val_lbl_dir,
slice_size=slice_sz, overlap_ratio=overlap,
)
results.append((name, map50, map5095, details, elapsed))
print(f" mAP50={map50:.4f} mAP50-95={map5095:.4f} 用时={elapsed:.1f}s")
for k in ["AP@50","AP@75","AP@80","AP@90","AP@95"]:
print(f" {k} = {details[k]:.4f}")
print(f"\n\n{'='*70}")
print(f"{'配置':<25} {'mAP50':>8} {'mAP50-95':>10} {'AP@75':>8} {'AP@90':>8} {'耗时':>8}")
print(f"{'='*70}")
for name, m50, m5095, det, t in results:
print(f"{name:<25} {m50:>8.4f} {m5095:>10.4f} {det['AP@75']:>8.4f} {det['AP@90']:>8.4f} {t:>7.1f}s")
print(f"{'='*70}")
if __name__ == "__main__":
main()