Zig-RiR简介

        Zig-RiR(Zigzag RWKV-in-RWKV)是一种专为高效医学图像分割设计的新型神经网络架构,由 Chen 等人在 2025 年提出。它基于 RWKV(Receptance Weighted Key Value) 模型——一种具有线性计算复杂度、能高效建模长序列的类 Transformer 架构——并针对医学图像的特性进行了关键性改进。

注:关于RWKV的讲解请看博主的这篇文章:

https://blog.csdn.net/qq_73038863/article/details/153137485?fromshare=blogdetail&sharetype=blogdetail&sharerId=153137485&sharerefer=PC&sharesource=qq_73038863&sharefrom=from_link

        传统 Vision Transformer(ViT)及其变体虽然能捕捉全局上下文,但其自注意力机制带来二次方计算开销,在高分辨率医学图像(如 1024×1024 的病理切片或 3D CT 体积)上效率低下。而直接将 RWKV 应用于视觉任务时,又因忽略局部细节和破坏空间连续性(例如按行扫描图像块导致相邻像素在序列中不连续),导致分割精度下降。

        为解决这些问题,Zig-RiR 基于TNT(Transformer in Transformer)的思想创新性地提出了一种嵌套式“句子-词语”结构:

(1)将输入图像划分为若干“视觉句子”(patches);

(2)每个“句子”再细分为“视觉词语”(sub-patches);

(3)使用 Outer Zig-RWKV 块建模全局语义关系;

(4)通过 Inner Zig-RWKV 块在每个局部区域内精细挖掘细节特征。

注:TNT(Transformer in Transformer)是一种改进的视觉Transformer架构,通过嵌套的Transformer结构增强模型对局部和全局特征的建模能力。其核心思想是在全局Transformer块中嵌入局部Transformer块,形成层次化特征提取机制,提升图像分类、目标检测等任务的性能。

        关于TNT的讲解请看博主的这篇文章:

https://blog.csdn.net/qq_73038863/article/details/153183386?fromshare=blogdetail&sharetype=blogdetail&sharerId=153183386&sharerefer=PC&sharesource=qq_73038863&sharefrom=from_link

        更重要的是,Zig-RiR 引入了Zigzag 扫描策略(Zigzag Scan),在将 2D/3D 图像展平为序列时,沿“之”字形路径遍历相邻块,最大程度保留空间连续性,从而更好地利用图像的局部归纳偏置。

        实验表明,Zig-RiR 在四个 2D/3D 医学分割数据集上均达到当前最优(SOTA)性能,同时在 1024×1024 高分辨率图像上实现:14.4 倍推理加速、89.5% GPU 显存降低。这使得 Zig-RiR 成为兼顾高精度、高效率、低资源消耗的理想选择,尤其适用于临床部署中的实时医学图像分析场景。

模型讲解

        模型的整体架构如下图所示:

第一步:输入图像 → 划分“视觉句子”和“视觉词语”

        通过一个卷积主干(Convolutional stem),将图像划分为多个区域:

        视觉句子(Visual Sentence):较大的图像块(如 8×8 像素),代表一个语义单元。

        视觉词语(Visual Word):每个句子内部再细分为更小的块(如 4×4),相当于“词”。

注:视觉句子和视觉单词不是指简单的图像分割后的块,是分割后的图像经过Convolutional stem处理后的结果。注意这点以防对原文中的尺寸描述有误解。

        每个“句子”被编码为一个向量(visual sentence embedding);每个“词语”也被编码为一个向量(visual word embedding)。这些向量按 Zigzag 顺序排列成序列,确保相邻像素在序列中也邻近,保留空间结构。

第二步:进入第一个 Zig-RiR Block 层(L₁)

        现在我们有了一个由“视觉句子”和“视觉词语”组成的序列。核心思想是RWKV-in-RWKV 嵌套结构。

        外层(Outer):使用 Zig-RWKV 对所有“句子”建模全局关系。

        内层(Inner):对每个“句子”内部的“词语”使用另一个 RWKV 模块进行精细建模。

Inner Zig-RWKV 处理

        输入句子按 Zigzag 顺序排列(如:左上 → 右上 → 左下 → 右下)送入轻量 RWKV 模块(Inner Block)。对每个“视觉句子”内部的 4 个“视觉词语”进行精细建模,更新这些词语的表示。这一步确保了局部特征得到了充分的捕捉和增强。

        输入:每个句子内的 4 个词

        输出:更新后的 4 个词

聚合为句子表示

        在将这些更新后的词语送入 Outer Zig-RWKV 之前,通常会对每个句子内部的词语进行某种形式的聚合操作(如平均池化或线性组合),以生成该句子的一个代表向量。这样做的目的是减少后续处理的复杂度,并且使得每个句子能够作为一个整体来考虑其与其他句子的关系。

Outer Zig-RWKV 处理

        所有这些句子级的表示被组织成一个序列,按照 Zigzag 扫描的方式排列,然后输入到 Outer Zig-RWKV 模块中。这里的 Zigzag 扫描是为了保持空间邻近性的同时降低维度,从而有效地捕捉全局依赖关系。

        输入:重组后的句子序列

        输出:更新后的句子表示

        下面这张图展示的是 Inner Zig-RWKV 模块,它通过多路径 Zigzag 扫描和双向 RWKV (Bi-WKV)对每个“视觉句子”内部的局部结构进行精细建模。

        对于Spatial Mix块:输入的 2D 特征图首先通过 Q-Shift 操作增强跨区域交互,随后其 Key(Ks)和 Value(Vs)序列按照预定义的 Zigzag 扫描路径(如 Δ_zig^(1,1) 和 Δ_zig^(1,2))被重排为一维序列,并依次送入两个双向 RWKV(Bi-WKV)模块。

第三步:Patch Merging —— 合并信息,减少分辨率

        为了逐步提升抽象层次,模型引入 Patch Merging 模块:

        将相邻的 2×2 个“句子”合并成一个更大的“超级句子”,特征维度翻倍(通道数 ↑),空间尺寸 ↓(H/W ↓)。这类似于 CNN 中的池化层,但这里是基于 Transformer 式聚合。

第四步:重复结构 —— 多级 Zig-RiR Block(L₂, L₃, L₄)与Patch Merging

        每一级都更关注更大尺度的语义结构,减少空间分辨率,增加通道深度,提升对全局形状、边界等的理解能力。

        相比纯 ViT,避免了在高分辨率下计算爆炸;相比 CNN,保留了更强的全局感知能力。

第五步:上采样与重建 —— 从粗到精恢复细节

        到了最深的一层(如 L₄),我们已经拥有了高度抽象的全局特征。但我们需要的是像素级预测,所以要“往上走”——解码器部分开始工作。

        Up-sampling Module(上采样):将低分辨率特征放大回原尺寸;

        Patch Expanding:将合并后的“句子”拆分成更细粒度的“词语”

        卷积层:进一步融合多尺度信息;

        重复以上步骤,逐级还原细节。

        与传统 U-Net 不同的是:这里没有跳跃连接,而是依靠 Zig-RiR 的长期记忆机制自动关联不同层级的信息。

第六步:预测头(Prediction Head)→ 输出分割掩码

        最后,经过多次上采样和卷积后,进入 Prediction Head:使用 1x1 卷积生成最终的预测图。

整个流程的例子

假设输入为:

        图像尺寸:3 × 512 × 512(RGB 皮肤镜图像)

        任务:分割出病灶区域(输出 2 类掩码)

步骤 1:卷积 Stem → 生成“视觉词语”

        使用 4×4 卷积,stride=4,将图像划分为 4×4 的小块(patch),每个 patch 映射为 64 维向量(64 维是一个在表达能力、计算效率和架构一致性之间取得最佳平衡的经验值,既能有效编码局部图像信息,又不会在高分辨率阶段造成计算或显存瓶颈。)。

        空间分辨率:512 ÷ 4 = 128

        总 patch 数:128 × 128 = 16,384 个

        输出形状:64 × 128 × 128

作用:把图像变成“视觉词语”序列。

步骤 2:分组 → 构建“视觉句子”

        每 2×2 = 4 个相邻词语组成一个“句子”。

        句子网格:128 ÷ 2 = 64 → 共 64×64 = 4,096 个句子。

作用:为嵌套建模准备局部语义单元。

步骤 3:Inner Zig-RWKV(局部建模)

        对每个句子内部的 4 个词,用 RWKV 建模局部关系。

输入:每个句子 [B, 4, 64]

输出:更新后的 [B, 4, 64]

BBatch size(批大小)一次输入多少张图像,比如 B=8 表示同时处理 8 张图
4Sequence Length(序列长度)每个“视觉句子”由 4 个“视觉词语” 组成,这 4 个词被排成一个短序列
64Feature Dimension(特征维度)每个“词”用一个 64 维向量 表示其语义和外观信息

        总输出仍为: 64 × 128 × 128

作用:增强局部细节(如边缘、纹理)。

步骤 4:聚合 → 得到句子表示

        将每个句子的 4 个词平均池化成 1 个向量,得到句子级序列:[B, 4096, 64] , 即 64 × 64 × 64。

BBatch size(批大小)一次处理多少张图像
4096Token 数量(句子数量)每张图被划分为 64 × 64 = 4096 个“视觉句子”
64特征维度每个“句子”用一个 64 维向量 表示其整体语义
[B, 4096, 64]序列形式:适合送入 RWKV(处理一维序列)
64 × 64 × 642D 特征图形式:64 通道,空间尺寸 64×64(适合可视化或后续卷积操作)

作用:降 token 数,生成“句子”级别的特征。

步骤 5:Outer Zig-RWKV(全局建模)

        将 4,096 个句子按 Zigzag 路径(蛇形)排成一维序列,输入到 RWKV 模块,建模长距离依赖(如病灶中心 ↔ 边界)。

输出:更新后的句子表示 [B, 4096, 64]

作用:捕捉全局语义结构。

步骤 6:Patch Merging(第一次降采样)

        将 64 × 64 × 64 特征图,每 2×2 区域合并,通道数翻倍(64 → 128),空间减半(64 → 32)。

输出: 128 × 32 × 32

作用:压缩信息,进入更高抽象层级。

步骤 7:重复 Block + Merge(共 4 层)

L₂128 × 32 × 32Zig-RiR Block + Merge256 × 16 × 16
L₃256 × 16 × 16Zig-RiR Block + Merge512 × 8 × 8
L₄512 × 8 × 8Zig-RiR Block(不再 Merge)512 × 8 × 8

步骤 8:上采样(解码器)

从 512 × 8 × 8 开始,逐步恢复分辨率:

上采样 → 256 × 16 × 16

上采样 → 128 × 32 × 32

上采样 → 64 × 64 × 64

上采样 → 32 × 128 × 128

最后插值 → 32 × 512 × 512

作用:重建高分辨率分割图。

步骤 9:预测头(Prediction Head)

        用 1×1 卷积 将 32 通道映射为 2 类(背景 / 病变)。

输出:2 × 512 × 512。

        应用 Sigmoid → 得到概率掩码。

阶段特征图尺寸 (C × H × W)
输入3 × 512 × 512
Stem64 × 128 × 128
L₁ 后64 × 128 × 128
Merge 1128 × 64 × 64
Merge 2256 × 32 × 32
Merge 3512 × 16 × 16
L₄ 输出512 × 8 × 8
解码后32 × 512 × 512
输出2 × 512 × 512
Logo

分享最新、最前沿的AI大模型技术,吸纳国内前几批AI大模型开发者

更多推荐