You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

自定义数据集微调YOLOX模型无法使用SAHI推理的问题

使用SAHI推理自定义微调YOLOX模型的解决方案

SAHI没有内置YOLOX的model_type,但可以通过自定义DetectionModel子类实现适配,核心是把YOLOX的推理输出转换成SAHI兼容的格式。以下是具体步骤:

1. 加载自定义YOLOX模型

按YOLOX官方流程加载你微调后的模型:

import torch
from yolox.exp import get_exp
from yolox.utils import fuse_model

# 加载微调时用的exp配置文件
exp = get_exp(exp_file="path/to/your/yolox_exp_config.py", name=None)
# 初始化模型
model = exp.get_model()
# 加载微调权重
ckpt = torch.load("path/to/your/fine_tuned_weights.pth", map_location="cpu")
model.load_state_dict(ckpt["model"])
# 可选:融合模型层加速推理
model = fuse_model(model)
# 切换到评估模式并指定设备
device = "cuda:0" if torch.cuda.is_available() else "cpu"
model.to(device).eval()

2. 实现SAHI兼容的DetectionModel子类

继承SAHI的BaseDetectionModel,实现predict方法完成YOLOX推理到SAHI格式的转换:

from sahi.models.base import BaseDetectionModel
from sahi.prediction import ObjectPrediction
from yolox.utils import postprocess

class YOLOXSAHIDetectionModel(BaseDetectionModel):
    def __init__(
        self,
        model,
        confidence_threshold: float = 0.3,
        device: str = "cpu",
        image_size: int = 640,
        num_classes: int = 80,  # 替换为你的数据集类别数
        nms_threshold: float = 0.3,  # 替换为你微调时用的NMS阈值
    ):
        super().__init__(
            model=model,
            confidence_threshold=confidence_threshold,
            device=device,
            image_size=image_size,
        )
        self.num_classes = num_classes
        self.nms_threshold = nms_threshold
        self.model.to(device).eval()

    def predict(self, image: torch.Tensor):
        # 执行YOLOX推理
        with torch.no_grad():
            outputs = self.model(image)
            # 后处理得到检测框
            outputs = postprocess(
                outputs,
                self.num_classes,
                self.confidence_threshold,
                self.nms_threshold,
                class_agnostic=False
            )
        
        # 转换为SAHI的ObjectPrediction格式
        object_predictions = []
        if outputs[0] is not None:
            for output in outputs[0]:
                bbox_xyxy = output[:4].cpu().numpy()
                score = output[4].cpu().numpy()
                category_id = int(output[5].cpu().numpy())
                
                object_pred = ObjectPrediction(
                    bbox=bbox_xyxy,
                    score=score,
                    category_id=category_id,
                    category_name=self.category_mapping.get(category_id, str(category_id))
                )
                object_predictions.append(object_pred)
        
        return object_predictions

# 初始化适配后的模型
# 可选:添加类别名称映射(替换为你的数据集类别)
category_mapping = {0: "cat", 1: "dog", 2: "bird"}
sahi_model = YOLOXSAHIDetectionModel(
    model=model,
    confidence_threshold=0.4,
    device=device,
    image_size=640,
    num_classes=exp.num_classes,  # 从exp配置中读取类别数
    nms_threshold=exp.nmsthre  # 从exp配置中读取NMS阈值
)
sahi_model.category_mapping = category_mapping

3. 执行SAHI切片推理

现在可以用SAHI的标准接口进行切片推理:

from sahi.predict import get_sliced_prediction
import cv2

# 加载目标图像
image = cv2.imread("path/to/your/test_image.jpg")
# 执行切片推理
sliced_result = get_sliced_prediction(
    image,
    sahi_model,
    slice_height=256,  # 根据你的需求调整切片尺寸
    slice_width=256,
    overlap_height_ratio=0.2,  # 切片重叠比例
    overlap_width_ratio=0.2
)

# 可视化结果
sliced_result.visualize()
# 打印检测结果
for pred in sliced_result.object_prediction_list:
    print(f"类别: {pred.category_name}, 置信度: {pred.score:.2f}, 坐标: {pred.bbox}")

关键注意事项

  • 确保SAHI的图像预处理(resize、归一化)和你微调YOLOX时的逻辑完全一致,否则会影响检测精度。
  • 如果你的YOLOX用了自定义预处理/后处理逻辑,需要对应修改YOLOXSAHIDetectionModel中的推理流程。

内容的提问来源于stack exchange,提问作者sama acm

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.15 06:53:20