码农知识堂 - 1000bd
  •   Python
  •   PHP
  •   JS/TS
  •   JAVA
  •   C/C++
  •   C#
  •   GO
  •   Kotlin
  •   Swift
  • 深度学习进阶(二十四)Swin 的二维 RPE


    合集 - 深度学习进阶(31)
    1.深度学习进阶(一)从注意力到自注意力03-312.深度学习进阶(二)多头自注意力机制(Multi-Head Attention)04-023.深度学习进阶(三)Transformer Block04-044.深度学习进阶(四)Transformer 整体结构04-065.深度学习进阶(五)Vision Transformer04-086.深度学习进阶(六)归纳偏置与蒸馏04-107.深度学习进阶(七)Data-efficient Image Transformer04-138.深度学习进阶(八)Swin Transformer04-159.深度学习进阶(九)池化技术的初步改进:RoI Pooling04-1910.深度学习进阶(十) RoI Align04-2111.深度学习进阶(十一)Position-Sensitive RoI Pooling04-2412.深度学习进阶(十二)可变形池化 deformable RS RoI Pooling04-2713.深度学习进阶(十三)可变形卷积 DCN04-2914.深度学习进阶(十四)ConvNeXt04-3015.深度学习进阶(十五)通道注意力 SE05-0316.深度学习进阶(十六) 混合注意力 CBAM05-0417.深度学习进阶(十七)高效通道注意力 ECA05-0518.深度学习进阶(十八)坐标注意力 CA05-0619.深度学习进阶(十九)相对位置编码 RPE05-0920.深度学习进阶(二十)Transformer-XL05-1121.深度学习进阶(二十一)跨窗口的 RPE05-1322.深度学习进阶(二十二)T5:NLP任务的首次大一统05-1523.深度学习进阶(二十三)偏置型 RPE05-18
    24.深度学习进阶(二十四)Swin 的二维 RPE05-20
    25.深度学习进阶(二十五)RoPE:现代 NLP 的位置编码范式05-2926.深度学习进阶(二十六)现代 LLM 的核心架构设计其一:RMSNorm06-0127.深度学习进阶(二十七)现代 LLM 的核心架构设计其二:SwiGLU06-0428.深度学习进阶(二十八)现代 LLM 的核心架构设计其三:Decoder-Only 下的 KV Cache06-0829.深度学习进阶(二十九)现代 LLM 的核心架构设计其四:GQA06-0930.深度学习进阶(三十)从 Transformer 到 LLaMA:现代 LLM 架构总览06-1631.深度学习进阶(三十一)FlashAttention:IO 感知的精确注意力06-19
    收起

    上一篇我们介绍了 T5 的偏置型 RPE,仅仅使用一个标量偏置,配合分桶策略,就用极低的复杂度实现了 NLP 的高效位置编码。

    而下一个问题就是:

    一维序列上的标量偏置,到了二维图像上要怎么做?

    这一篇我们来补上之前的 Swin Transformer 中一个当时没有展开的细节:二维 RPE。

    1. 为什么需要二维 RPE?#

    在 T5 中,相对位置是一个标量:i−j。因为文本是一维序列,两个 token 之间的关系只需要一个数字就能描述。

    但图像数据不同。一张图像被划分为 M×M 的 patch 网格后,两个 patch 之间的相对位置是二维的。

    一个 patch 到另一个 patch 的偏移,不仅有"水平方向的距离",还有"垂直方向的距离"。

    具体来说,对于图像中的位置 (x1,y1) 和 (x2,y2),相对位置是分成下面两部分:

    Δx=x1−x2,Δy=y1−y2

    这时如果再用一维标量来描述这个二维偏移,必然丢失方向信息。

    好在,Swin 本身的 Window Attention 设计其实已经为 RPE 减了负:注意力只在 7×7 的窗口内进行。
    这种设计让我们可以不再过多考虑 NLP 中的编码外推问题,但相应的,在这个局部范围内,精确的二维相对关系对建模视觉结构至关重要。

    因此,Swin 设计了一套二维的相对位置编码方案。

    2. 二维 RPE 如何构造?#

    我们直接来看 Swin 在窗口注意力中使用的公式:

    Attention(Q,K,V)=Softmax(QKTd+B)V

    公式本身在形式上和 T5 是完全相同的,关键在于偏置矩阵 B 的构造上。
    我们分点来展开:

    2.1 直接将 RPE 推广到二维#

    我们先来看看最直接的方法:
    对于一个 M×M 的窗口,直接设计 B∈RM2×M2,其中 Bij 表示窗口内第 i 个 patch 和第 j 个 patch 之间的偏置值。

    我们用一个简单的例子来演示为什么是 M2×M2 ,假设窗口大小:M=2 ,那么窗口就是:

    [t1t2t3t4]

    现在,每个 token 都要和另外所有 token 建立关系。那么 QKT 计算的注意力得分矩阵形状就是这样的:

    [t1→t1t1→t2t1→t3t1→t4t2→t1t2→t2t2→t3t2→t4t3→t1t3→t2t3→t3t3→t4t4→t1t4→t2t4→t3t4→t4]

    偏置矩阵必须和注意力矩阵一一对应。所以 B∈RM2×M2。
    这种方法当然是可以跑通的,但我们要考虑二维带来的参数问题:

    如果直接学习一个 M2×M2 的参数矩阵,那每个注意力头就得维护 M4 个参数。一个 Swin 有多个头和多个层,累计下来参数巨大。

    因此, Swin 自然有对应的改进。

    2.2 空间关系的平移不变性#

    在 NLP 中,我们只针对每种相对位置设计偏置,但是在上面方案里,你会发现直接推广会带来很多无意义的参数,核心是因为:

    在二维数据中相对逻辑更加凸显,窗口内大量位置对其实拥有相同的相对偏移。

    比如,patch (0,0) 和 (1,0) 之间的偏移是 (Δx=1,Δy=0),而 patch (2,0) 和 (3,0) 之间的偏移同样是 (Δx=1,Δy=0)。
    它们本质上描述的是同一种空间关系,理应共享同一个偏置值。

    于是 Swin 的做法是:推广相对逻辑,不直接学习 B,而是学习一个小得多的偏置表,再通过二维索引从中查值。

    3. 紧凑偏置表与查表逻辑#

    3.1 二维相对位置的计算#

    首先,对于一个 M×M 的窗口,给每个位置一个坐标 (x,y),显然:

    x,y∈[0,M−1]

    对于任意两个 patch ,二维相对偏移是:

    Δx=xi−xj,Δy=yi−yj

    那么,Δx 的取值范围就是 [−(M−1),M−1],一共 2M−1 种可能。
    Δy 同理,这部分的计算逻辑和 T5 是完全相同的。

    现在,我们知道了:所有可能的 (Δx,Δy) 组合一共有 (2M−1)2 种,也就是说:

    我们只需要一个 (2M−1)×(2M−1) 的偏置表,就能覆盖窗口内所有可能的位置关系。

    这就是 Swin 的紧凑偏置表 B^:

    B^∈R(2M−1)×(2M−1)

    建表本身的逻辑到此结束,但现在还有一个小问题:

    QKT 和 B^ 大小不一,对于每组注意力计算,我要如何查表注入相应偏置?

    3.2 查表过程#

    其实这步可以理解为:如何将 B^ 内的值映射到总公式里的 B 中?

    首先,前面我们已经知道了:QKT∈RM2×M2
    因此,真正参与 Attention 计算的偏置矩阵 B,也必须是 M2×M2。

    但我们刚刚学习的紧凑偏置表只有:

    B^∈R(2M−1)×(2M−1)

    不难理解,为了让二者适配,Swin 的设计是这样的:

    对于 Attention Matrix 中的每一个元素,都先计算两个 patch 的相对位移,再去 B^ 中查对应 bias。

    展开来说, QKT 中的每一个元素本质上都对应“一对 patch 的关系”,而每一对 patch 都有自己的 (Δx,Δy),因此,我们可以计算相对位移,实现查表取值:

    Bij=B^[Δx,Δy]

    这就实现了相同相对位移的 patch 对,共享同一个偏置。

    不过这在实现中还有一个问题:

    数组索引没有负数,负偏移并不能和其索引直接对应。

    而 Δx,Δy∈[−(M−1),M−1],因此 Swin 会先做一次平移去寻找正确索引:

    Δx′=Δx+(M−1)

    Δy′=Δy+(M−1)

    现在:

    Δx′,Δy′∈[0,2M−2]

    于是查表过程就变成:

    Bij=B^[Δx′,Δy′]

    image.png

    字母还是有些抽象,我们再举一个实例:设 M=3 ,那么 patch 网格可以就是:

    [(0,0)(0,1)(0,2)(1,0)(1,1)(1,2)(2,0)(2,1)(2,2)]

    此时 2M−1=5 ,因此 B^∈R5×5,如果当前 patch 为 (0,0),它去关注 (2,2) ,那么:

    ,,Δx=0−2=−2,Δy=0−2=−2

    现在,我们需要查:

    B^[−2,−2]

    显然,数组索引不能为负数。 所以进行平移:

    M−1=2

    于是:

    Δx′=Δx+2

    Δy′=Δy+2

    原本的 [−2,−2] ,就被平移成 [0,0]:

    B^[−2,−2]⇒B^[0,0]

    这里可能容易疑惑的一点是:

    B^∈R5×5 中存储的并不是“偏移坐标本身”,而是“对应相对位移的偏移参数”。

    展开来说:数学意义上的 (Δx,Δy)=(−2,−2) 会被映射到数组索引(0,0),因此,B^[0,0] 实际存储的就是相对位移为 (−2,−2) 时对应的偏置。

    这样,所有原本可能为负数的二维位移都被映射到了合法数组索引,可以稳定完成查表。
    最终所有 patch 两两之间都会完成一次查表从而动态构造出完整的偏置矩阵:

    B∈RM2×M2

    随后:

    QKTd+B

    即可完成二维相对位置信息的注入。
    值得一提的是在具体实现中,二维紧凑表会被展平成一维,以类似“编号”的逻辑取值,根本逻辑没变,明白即可。

    3.3 参数对比#

    来看看两种方式的参数对比:

    方式 M=7(Swin 默认) M=14
    暴力直接法 49×49=2401 196×196=38416
    Swin 紧凑法 13×13=169 27×27=729
    压缩比 约 14 倍 约 53 倍

    很明显,随着窗口增大,紧凑表的优势会更加明显。

    这便是 Swin 的二维 RPE 的完整逻辑,它十分符合 Swin 的整体设计逻辑,配合其实现了 Attention 在 CV 领域的推广,也在后续的很多混合架构中被使用。

    作者:哥布林学者

    出处:https://www.cnblogs.com/Goblinscholar/p/20095027

    版权:本作品采用「署名-非商业性使用-相同方式共享 4.0 国际」许可协议进行许可。

    给自己一些时间。

  • 相关阅读:
    ADSP-21489的开发详解:Norflash的编程和烧写
    面试题--SpringBoot
    IP地址规划设计
    什么是SHA384,SHA384和SHA512有什么区别
    【Bug排查】Uncaught (in promise) Error: Infinite redirect in navigation guard
    产品生命周期有哪些
    上海亚商投顾:沪指重返3100点 房地产板块掀涨停潮
    【C语言】通讯录
    湖北工业大学计算机考研资料汇总
    分库分表实战
  • 原文地址:https://www.cnblogs.com/Goblinscholar/p/20095027
  • 最新文章
  • 沪漂五周年了:我越来越迷茫了
    Agentic Skill Routing 实战:别再把所有 Skill 塞进 AI Agent 上下文
    MySQL-Seconds_behind_master的精度误差
    [MAF预定义ChatClient中间件-03]CachingChatClient——利用缓存省钱省时间
    AI的至暗历史:从万众期待到被政府撤资,AI的两次死亡徘徊
    Agent OS :五种驯服不确定性的范式
    PortSwigger SQL注入LAB11
    数据库即时编译JIT
    [Begin]AI Learn Data Day 0
    深度学习进阶(二十七)现代 LLM 的核心架构设计其二:SwiGLU
  • 热门文章
  • 十款代码表白小特效 一个比一个浪漫 赶紧收藏起来吧!!!
    奉劝各位学弟学妹们,该打造你的技术影响力了!
    五年了,我在 CSDN 的两个一百万。
    Java俄罗斯方块,老程序员花了一个周末,连接中学年代!
    面试官都震惊,你这网络基础可以啊!
    你真的会用百度吗?我不信 — 那些不为人知的搜索引擎语法
    心情不好的时候,用 Python 画棵樱花树送给自己吧
    通宵一晚做出来的一款类似CS的第一人称射击游戏Demo!原来做游戏也不是很难,连憨憨学妹都学会了!
    13 万字 C 语言从入门到精通保姆级教程2021 年版
    10行代码集2000张美女图,Python爬虫120例,再上征途
小工具 小游戏
Copyright © 2022 侵权请联系2656653265@qq.com    京ICP备2022015340号-1

京公网安备 11010502049817号