YoLo World代码块解读

MaxSigmoidAttnBlock

分别处理图像与文本特征,计算这两者的相关性,得到整个句子所有word中最大的相关性数值作为attention作用于图像特征中。

python 复制代码
class MaxSigmoidAttnBlock(nn.Module):
    """Max Sigmoid attention block."""

    def __init__(self, c1, c2, nh=1, ec=128, gc=512, scale=False):
        """Initializes MaxSigmoidAttnBlock with specified arguments."""
        super().__init__()
        self.nh = nh
        self.hc = c2 // nh
        self.ec = Conv(c1, ec, k=1, act=False) if c1 != ec else None
        self.gl = nn.Linear(gc, ec)
        self.bias = nn.Parameter(torch.zeros(nh))
        self.proj_conv = Conv(c1, c2, k=3, s=1, act=False)
        self.scale = nn.Parameter(torch.ones(1, nh, 1, 1)) if scale else 1.0

    def forward(self, x, guide):
        """Forward process."""
        bs, _, h, w = x.shape

        guide = self.gl(guide)
        guide = guide.view(bs, -1, self.nh, self.hc)
        embed = self.ec(x) if self.ec is not None else x
        embed = embed.view(bs, self.nh, self.hc, h, w)

        aw = torch.einsum("bmchw,bnmc->bmhwn", embed, guide) #文本与图像的交叉注意力
        aw = aw.max(dim=-1)[0] #整个句子最大相关性数值作为attention作用于图像特征中。
        aw = aw / (self.hc**0.5)
        aw = aw + self.bias[None, :, None, None]
        aw = aw.sigmoid() * self.scale

        x = self.proj_conv(x)
        x = x.view(bs, self.nh, -1, h, w)
        x = x * aw.unsqueeze(2)
        return x.view(bs, -1, h, w)

WorldDetect

基本和Detect一样,主要是分类分支需要通过BNContrastiveHead将图像特征于文字特征计算相关性得到每个grid中所有类别的置信度。

python 复制代码
class WorldDetect(Detect):
    def __init__(self, nc=80, embed=512, with_bn=False, ch=()):
        """Initialize YOLOv8 detection layer with nc classes and layer channels ch."""
        super().__init__(nc, ch)
        c3 = max(ch[0], min(self.nc, 100))
        self.cv3 = nn.ModuleList(nn.Sequential(Conv(x, c3, 3), Conv(c3, c3, 3), nn.Conv2d(c3, embed, 1)) for x in ch)
        self.cv4 = nn.ModuleList(BNContrastiveHead(embed) if with_bn else ContrastiveHead() for _ in ch)

    def forward(self, x, text):
        """Concatenates and returns predicted bounding boxes and class probabilities."""
        for i in range(self.nl):
            x[i] = torch.cat((self.cv2[i](x[i]), self.cv4[i](self.cv3[i](x[i]), text)), 1)
            #cv4可以获得类别编码,[512,80,80]*[4,512]=[4,80,80]
相关推荐
YOLO数据集集合15 小时前
大模型融合YOLO铁路要素缺陷分析系统 | 铁路缺陷检测 YOLO DeepSeek 大语言模型 智能巡检 9141期
人工智能·yolo·目标检测·语言模型·铁路缺陷·轨道缺陷
Tingmanyi15 小时前
YOLOv13改进策略【Neck篇】| AFPN 渐进式特征金字塔,整体替换 Head 参数反降 82 万
yolo
YOLO_DATA2 天前
YOLO27防震锤缺陷检测数据集 防震锤数据集 1000张 防震锤 带标注 voc yolo 2 类 目标检测
人工智能·深度学习·yolo·目标检测·计算机视觉·数据集·无人机
YOLO数据集集合2 天前
红外多类型无人机检测数据集 | 红外检测 无人机识别 多类型分类 低空安防 反无人机 9139期
yolo·目标检测·分类·数据挖掘·无人机·无人机检测·红外无人机
YOLO数据集集合3 天前
电力设备目标检测数据集 | 电力设备 变电站巡检 部件识别 目标检测 YOLO格式 9127期
人工智能·yolo·目标检测·计算机视觉·目标跟踪·电力设备
FPGA小徐3 天前
FPGA 部署 YOLO 完整指南:从模型量化到 Zynq PS+PL 硬件加速
yolo·fpga开发
FPGA小徐3 天前
Ubuntu 22.04 安装 Miniconda、PyTorch 与 YOLOv8 并完成 CPU 推理
pytorch·yolo·ubuntu
YOLO数据集集合3 天前
桥梁病害检测数据集 | 桥梁病害 结构健康监测 裂缝检测 钢筋外露9134期
yolo·目标检测·外墙裂缝·建筑病害·桥梁病害·桥梁缺陷
zy_destiny3 天前
深度学习实战-基于YOLOv8的玉米雄穗目标检测:从数据集下载到训练部署全流程实录
yolo·目标检测·机器学习
YOLO数据集集合4 天前
外墙裂缝目标检测YOLO数据集:6,296张图像助力建筑病害智能识别| 外墙裂缝检测 YOLO数据集 建筑健康监测 无人机巡检 目标检测8030期
yolo·目标检测·无人机·建筑·外墙裂缝·建筑病害