初始化 observer 项目:纯代码,不含权重与训练数据
This commit is contained in:
@@ -0,0 +1,154 @@
|
||||
#!/usr/bin/env python3
|
||||
"""
|
||||
使用 Hugging Face transformers 进行零样本目标检测
|
||||
使用 OWL-ViT 模型检测野鸡
|
||||
"""
|
||||
|
||||
import cv2
|
||||
import torch
|
||||
from pathlib import Path
|
||||
from transformers import OwlViTProcessor, OwlViTForObjectDetection
|
||||
from PIL import Image
|
||||
|
||||
# 配置
|
||||
IMAGES_DIR = Path(__file__).parent / "datasets" / "images" / "pheasant"
|
||||
OUTPUT_DIR = Path(__file__).parent / "datasets" / "yolo_format"
|
||||
|
||||
# 检测配置
|
||||
CONFIDENCE_THRESHOLD = 0.1 # 置信度阈值
|
||||
|
||||
# 文本描述
|
||||
TEXT_PROMPTS = ["pheasant", "wild bird", "bird in grass"]
|
||||
|
||||
class ZeroShotDetector:
|
||||
def __init__(self):
|
||||
print("加载 OWL-ViT 模型...")
|
||||
self.processor = OwlViTProcessor.from_pretrained("google/owlvit-base-patch32")
|
||||
self.model = OwlViTForObjectDetection.from_pretrained("google/owlvit-base-patch32")
|
||||
self.model.eval()
|
||||
print("OWL-ViT 模型加载完成")
|
||||
|
||||
def detect(self, image_path):
|
||||
"""检测图片中的野鸡"""
|
||||
# 读取图片
|
||||
image = Image.open(image_path).convert("RGB")
|
||||
|
||||
# 准备输入
|
||||
inputs = self.processor(text=TEXT_PROMPTS, images=image, return_tensors="pt")
|
||||
|
||||
# 推理
|
||||
with torch.no_grad():
|
||||
outputs = self.model(**inputs)
|
||||
|
||||
# 获取结果
|
||||
target_sizes = torch.tensor([image.size[::-1]]) # [height, width]
|
||||
results = self.processor.post_process_grounded_object_detection(
|
||||
outputs, threshold=CONFIDENCE_THRESHOLD, target_sizes=target_sizes
|
||||
)[0]
|
||||
|
||||
# 解析结果
|
||||
detections = []
|
||||
for score, label, box in zip(results["scores"], results["labels"], results["boxes"]):
|
||||
box = box.tolist()
|
||||
detections.append({
|
||||
'bbox': box, # [x1, y1, x2, y2]
|
||||
'score': score.item(),
|
||||
'label': TEXT_PROMPTS[label],
|
||||
})
|
||||
|
||||
return detections
|
||||
|
||||
def convert_to_yolo_format(bbox, img_width, img_height):
|
||||
"""将边界框转换为 YOLO 格式 (cx, cy, w, h)"""
|
||||
x1, y1, x2, y2 = bbox
|
||||
cx = (x1 + x2) / 2 / img_width
|
||||
cy = (y1 + y2) / 2 / img_height
|
||||
w = (x2 - x1) / img_width
|
||||
h = (y2 - y1) / img_height
|
||||
return cx, cy, w, h
|
||||
|
||||
def main():
|
||||
# 初始化检测器
|
||||
detector = ZeroShotDetector()
|
||||
|
||||
# 创建输出目录
|
||||
(OUTPUT_DIR / "images" / "train").mkdir(parents=True, exist_ok=True)
|
||||
(OUTPUT_DIR / "labels" / "train").mkdir(parents=True, exist_ok=True)
|
||||
|
||||
# 获取所有图片
|
||||
image_extensions = {".jpg", ".jpeg", ".png", ".bmp"}
|
||||
image_files = [
|
||||
f for f in IMAGES_DIR.iterdir()
|
||||
if f.suffix.lower() in image_extensions
|
||||
]
|
||||
|
||||
print(f"找到 {len(image_files)} 张图片")
|
||||
|
||||
# 处理每张图片
|
||||
labeled_count = 0
|
||||
for img_path in image_files:
|
||||
print(f"\n处理: {img_path.name}")
|
||||
|
||||
# 读取图片
|
||||
image = cv2.imread(str(img_path))
|
||||
if image is None:
|
||||
print(f" 跳过: 无法读取 {img_path.name}")
|
||||
continue
|
||||
|
||||
img_height, img_width = image.shape[:2]
|
||||
|
||||
# 检测野鸡
|
||||
detections = detector.detect(img_path)
|
||||
|
||||
# 生成标注
|
||||
labels = []
|
||||
for det in detections:
|
||||
bbox = det['bbox']
|
||||
score = det['score']
|
||||
|
||||
# 转换为 YOLO 格式
|
||||
cx, cy, w, h = convert_to_yolo_format(bbox, img_width, img_height)
|
||||
labels.append(f"0 {cx:.6f} {cy:.6f} {w:.6f} {h:.6f}\n")
|
||||
|
||||
# 在图片上绘制检测框(用于可视化)
|
||||
x1, y1, x2, y2 = bbox
|
||||
cv2.rectangle(image, (int(x1), int(y1)), (int(x2), int(y2)), (0, 255, 0), 2)
|
||||
cv2.putText(image, f"pheasant: {score:.2f}", (int(x1), int(y1) - 10),
|
||||
cv2.FONT_HERSHEY_SIMPLEX, 0.6, (0, 255, 0), 2)
|
||||
|
||||
# 保存标注文件
|
||||
if len(labels) > 0:
|
||||
labeled_count += 1
|
||||
# 保存带标注的图片(用于可视化)
|
||||
cv2.imwrite(str(OUTPUT_DIR / "images" / "train" / img_path.name), image)
|
||||
|
||||
# 保存标注
|
||||
label_path = OUTPUT_DIR / "labels" / "train" / (img_path.stem + ".txt")
|
||||
with open(label_path, "w") as f:
|
||||
f.writelines(labels)
|
||||
|
||||
print(f" ✓ 标注了 {len(labels)} 个目标")
|
||||
else:
|
||||
print(f" - 未检测到野鸡")
|
||||
|
||||
print(f"\n完成!")
|
||||
print(f" 总图片数: {len(image_files)}")
|
||||
print(f" 有效标注: {labeled_count}")
|
||||
print(f" 输出目录: {OUTPUT_DIR}")
|
||||
|
||||
# 创建 data.yaml
|
||||
data_yaml = OUTPUT_DIR / "data.yaml"
|
||||
with open(data_yaml, "w") as f:
|
||||
f.write(f"""# Observer 数据集配置 - 只识别野鸡
|
||||
path: {OUTPUT_DIR}
|
||||
train: images/train
|
||||
val: images/train
|
||||
|
||||
# 类别
|
||||
nc: 1
|
||||
names: ['pheasant']
|
||||
""")
|
||||
print(f" 数据配置: {data_yaml}")
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
Reference in New Issue
Block a user