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

PyTorch二分类:如何基于最高概率类别绘制图像Bounding Box

猫狗分类模型能否绘制Bounding Box?如何获取坐标?

你目前用PyTorch实现了猫狗二分类,现有预训练模型仅能输出图像的类别概率,相关代码及模型输出如下:

分类模型代码

device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
class_names = ['dogs', 'cats']
    
def predict_label(image_path, model):
    img = Image.open(image_path)
    img_transformed = transformer(img)
    
    with torch.no_grad():
        model.eval()
        output = model(img_transformed)
        print(output)
        index = output.data.cpu().numpy().argmax()
        return class_names[index]

model = torch.load('dogs_cats')
predict_label("./husky.jpg", model)

模型输出

tensor([[ 0.4717, -0.1059]], device='cuda:0')

能否绘制Bounding Box?

当前的分类模型无法直接输出目标的Bounding Box坐标,因为普通图像分类模型的核心任务是预测整张图像的类别,它学习的是全局特征,没有针对目标的位置信息进行训练,输出仅包含类别概率,不涉及任何位置数据。

但可以通过两种方式实现类似需求:


方法1:改用目标检测模型

目标检测模型(如YOLO、Faster R-CNN、SSD)的设计目标就是同时预测图像中目标的类别和Bounding Box坐标,输出结果会包含每个目标的边界框信息(通常是左上角(x1,y1)和右下角(x2,y2)坐标)。

你可以基于PyTorch加载预训练的目标检测模型,示例代码大致如下:

import torch

# 加载预训练YOLOv5模型
model = torch.hub.load('ultralytics/yolov5', 'yolov5s', pretrained=True)
# 预测图像
results = model("./husky.jpg")
# 获取检测结果(包含类别、置信度、Bounding Box)
results.print()
# 保存带Bounding Box的图像
results.save()

这类模型会直接输出目标的精确Bounding Box,适合需要准确边界的场景。


方法2:用Grad-CAM生成近似关注区域(无需重新训练)

如果不想更换模型,可以使用**Grad-CAM(梯度加权类激活映射)**可视化模型关注的区域,再从生成的热力图中提取近似的Bounding Box。这种方法通过计算模型最后一层特征图的梯度,得到对当前类别贡献最大的区域,再将热力图转化为边界框。

示例代码如下:

import torch
import torch.nn.functional as F
from PIL import Image
import numpy as np
import cv2

# 定义Grad-CAM类
class GradCAM:
    def __init__(self, model, target_layer):
        self.model = model
        self.target_layer = target_layer
        self.gradients = None
        self.activations = None
        
        # 注册钩子捕获梯度和特征图
        def backward_hook(module, grad_in, grad_out):
            self.gradients = grad_out[0]
        def forward_hook(module, input, output):
            self.activations = output
        
        self.target_layer.register_backward_hook(backward_hook)
        self.target_layer.register_forward_hook(forward_hook)
    
    def generate_cam(self, input_tensor, target_class):
        # 前向传播
        output = self.model(input_tensor)
        # 清零梯度
        self.model.zero_grad()
        # 计算目标类别的损失并反向传播
        class_loss = output[0, target_class]
        class_loss.backward()
        
        # 计算梯度权重
        gradients = self.gradients.data.cpu().numpy()[0]
        activations = self.activations.data.cpu().numpy()[0]
        weights = np.mean(gradients, axis=(1,2))
        
        # 生成热力图
        cam = np.zeros(activations.shape[1:], dtype=np.float32)
        for i, w in enumerate(weights):
            cam += w * activations[i]
        
        # 归一化热力图
        cam = np.maximum(cam, 0)
        cam = cv2.resize(cam, (input_tensor.shape[3], input_tensor.shape[2]))
        cam = cam / cam.max()
        return cam

# 加载模型和图像
model = torch.load('dogs_cats')
model.eval()
img_path = "./husky.jpg"
img = Image.open(img_path).convert('RGB')
img_tensor = transformer(img).unsqueeze(0).to(device)

# 选择目标层(需根据你的模型结构调整,比如ResNet的layer4[-1])
target_layer = model.layer4[-1]

# 生成Grad-CAM热力图
grad_cam = GradCAM(model, target_layer)
pred_class = predict_label(img_path, model)
target_class = class_names.index(pred_class)
cam = grad_cam.generate_cam(img_tensor, target_class)

# 从热力图提取Bounding Box
thresh = 0.5
binary_cam = (cam > thresh).astype(np.uint8)
# 寻找轮廓
contours, _ = cv2.findContours(binary_cam, cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_SIMPLE)
# 获取最大轮廓的Bounding Box
if contours:
    largest_contour = max(contours, key=cv2.contourArea)
    x, y, w, h = cv2.boundingRect(largest_contour)
    # 转换为原始图像的坐标(需对应transformer的缩放/裁剪逻辑调整)
    img_width, img_height = img.size
    x = int(x * img_width / cam.shape[1])
    y = int(y * img_height / cam.shape[0])
    w = int(w * img_width / cam.shape[1])
    h = int(h * img_height / cam.shape[0])
    # 在原始图像上绘制Bounding Box
    img_np = np.array(img)
    cv2.rectangle(img_np, (x, y), (x+w, y+h), (0, 255, 0), 2)
    # 保存或显示图像
    Image.fromarray(img_np).save("husky_with_box.jpg")

注意:这种方法得到的是模型关注的近似区域,不是精确的目标边界,适合可视化模型关注点,而非精准检测。


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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.12 05:50:37