goat / Scripts /refine_annotations.py
LightChuan's picture
Upload folder using huggingface_hub
6a5bb7e verified
Raw
History Blame Contribute Delete
7.44 kB
"""标注精修工具 — 用模型预测辅助修正GT标注
原理:
1. 用最佳模型预测训练集
2. 对每个GT框,找最匹配的预测框
3. 如果IoU在0.5~0.85之间(有偏差但能匹配), 用加权平均修正GT
4. IoU>0.85的保持不变(标注已经很好)
5. IoU<0.5的标记为需人工审核
为什么有效:
模型在大量数据上训练后, 对目标边界的回归比单个标注员更稳定。
特别是小目标, 标注时2px偏差就会导致IoU从0.9降到0.5,
而模型预测虽然不完美, 但平均来看比有噪声的标注更一致。
用法:
python Scripts/refine_annotations.py
python Scripts/refine_annotations.py --dry-run # 只统计, 不修改
python Scripts/refine_annotations.py --ratio 0.7 # GT权重0.7, 预测权重0.3
"""
import sys, os, glob, argparse, shutil
import numpy as np
from PIL import Image
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
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 refine_box(gt_xyxy, pred_xyxy, gt_weight=0.6):
"""加权平均修正: 偏向GT但吸收模型的定位优势"""
pw = 1 - gt_weight
return [g * gt_weight + p * pw for g, p in zip(gt_xyxy, pred_xyxy)]
def xyxy_to_xywhn(box_xyxy, img_w, img_h):
"""xyxy像素坐标 → 归一化xywh"""
cx = (box_xyxy[0] + box_xyxy[2]) / 2 / img_w
cy = (box_xyxy[1] + box_xyxy[3]) / 2 / img_h
w = (box_xyxy[2] - box_xyxy[0]) / img_w
h = (box_xyxy[3] - box_xyxy[1]) / img_h
return [max(0, min(1, v)) for v in [cx, cy, w, h]]
def main():
parser = argparse.ArgumentParser()
parser.add_argument("--model", default="runs/detect/Detection_experiments/v5_5_s_highres/weights/best.pt")
parser.add_argument("--data-dir", default="Data/Detection_dataset")
parser.add_argument("--imgsz", type=int, default=1536)
parser.add_argument("--ratio", type=float, default=0.6,
help="GT权重 (0.5=各一半, 0.7=偏GT, 0.3=偏模型)")
parser.add_argument("--iou-low", type=float, default=0.4,
help="IoU低于此值不修正(可能是漏检/误检)")
parser.add_argument("--iou-high", type=float, default=0.85,
help="IoU高于此值不修正(标注已精确)")
parser.add_argument("--dry-run", action="store_true",
help="只统计不修改")
parser.add_argument("--split", default="train",
help="处理哪个split")
args = parser.parse_args()
from ultralytics import YOLO
model = YOLO(args.model)
img_dir = os.path.join(args.data_dir, f"images/{args.split}")
lbl_dir = os.path.join(args.data_dir, f"labels/{args.split}")
# 备份标注
if not args.dry_run:
backup_dir = lbl_dir + "_backup_before_refine"
if not os.path.exists(backup_dir):
shutil.copytree(lbl_dir, backup_dir)
print(f"[Backup] 原始标注已备份到: {backup_dir}")
else:
print(f"[Backup] 备份已存在: {backup_dir}")
imgs = sorted(glob.glob(os.path.join(img_dir, "*")))
print(f"\n处理 {len(imgs)} 张图片...")
stats = {"total": 0, "refined": 0, "kept": 0, "unmatched": 0, "need_review": []}
for idx, img_path in enumerate(imgs):
basename = os.path.splitext(os.path.basename(img_path))[0]
lbl_path = os.path.join(lbl_dir, basename + ".txt")
if not os.path.exists(lbl_path):
continue
with open(lbl_path) as f:
gt_lines = [l.strip() for l in f if l.strip()]
if not gt_lines:
continue
# 模型预测
results = model.predict(img_path, imgsz=args.imgsz, conf=0.25, iou=0.6, verbose=False)
pred_boxes = results[0].boxes.xyxy.cpu().numpy()
pred_confs = results[0].boxes.conf.cpu().numpy() if len(results[0].boxes) > 0 else np.array([])
img = Image.open(img_path)
W, H = img.size
new_lines = []
review_count = 0
for line in gt_lines:
parts = line.split()
cls_id = parts[0]
cx, cy, w, h = float(parts[1]), float(parts[2]), float(parts[3]), float(parts[4])
gt_xyxy = [(cx-w/2)*W, (cy-h/2)*H, (cx+w/2)*W, (cy+h/2)*H]
area = w * h
stats["total"] += 1
if len(pred_boxes) == 0:
new_lines.append(line)
stats["unmatched"] += 1
continue
# 找最匹配的预测框
best_iou, best_idx = 0, -1
for j, pb in enumerate(pred_boxes):
iou = compute_iou(gt_xyxy, pb)
if iou > best_iou:
best_iou = iou
best_idx = j
if best_iou >= args.iou_high:
# 标注已经很好, 不修改
new_lines.append(line)
stats["kept"] += 1
elif best_iou >= args.iou_low:
# 有偏差, 加权修正
# 小目标更信任模型(标注噪声大), 大目标更信任GT
if area < 0.003:
local_ratio = max(args.ratio - 0.1, 0.3) # 小目标: 更信模型
else:
local_ratio = min(args.ratio + 0.1, 0.8) # 大目标: 更信GT
refined = refine_box(gt_xyxy, pred_boxes[best_idx], local_ratio)
new_xywhn = xyxy_to_xywhn(refined, W, H)
new_lines.append(f"{cls_id} {new_xywhn[0]:.6f} {new_xywhn[1]:.6f} {new_xywhn[2]:.6f} {new_xywhn[3]:.6f}")
stats["refined"] += 1
else:
# IoU太低, 可能是漏检或误标, 保持原样但标记
new_lines.append(line)
stats["unmatched"] += 1
review_count += 1
if review_count > 3:
stats["need_review"].append(basename)
# 写入修正后的标注
if not args.dry_run:
with open(lbl_path, "w") as f:
f.write("\n".join(new_lines) + "\n")
if (idx + 1) % 100 == 0:
print(f" [{idx+1}/{len(imgs)}] refined={stats['refined']} kept={stats['kept']} unmatched={stats['unmatched']}")
print(f"\n{'='*50}")
print(f"标注精修完成{'(dry-run)' if args.dry_run else ''}!")
print(f" 总框数: {stats['total']}")
print(f" 已修正: {stats['refined']} ({stats['refined']/max(stats['total'],1)*100:.1f}%)")
print(f" 已保留: {stats['kept']} ({stats['kept']/max(stats['total'],1)*100:.1f}%)")
print(f" 未匹配: {stats['unmatched']} ({stats['unmatched']/max(stats['total'],1)*100:.1f}%)")
print(f" 需审核图: {len(stats['need_review'])} 张")
if stats["need_review"]:
review_file = "need_review_images.txt"
with open(review_file, "w") as f:
f.write("\n".join(stats["need_review"]))
print(f" 需审核列表: {review_file}")
if not args.dry_run:
print(f"\n下一步: 用精修后的标注重新训练")
print(f" 原始标注备份: {lbl_dir}_backup_before_refine/")
print(f"{'='*50}")
if __name__ == "__main__":
main()