Zig-RiR(Zigzag RWKV-in-RWKV)模型讲解
Zig-RiR简介
Zig-RiR(Zigzag RWKV-in-RWKV)是一种专为高效医学图像分割设计的新型神经网络架构,由 Chen 等人在 2025 年提出。它基于 RWKV(Receptance Weighted Key Value) 模型——一种具有线性计算复杂度、能高效建模长序列的类 Transformer 架构——并针对医学图像的特性进行了关键性改进。
注:关于RWKV的讲解请看博主的这篇文章:
传统 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的讲解请看博主的这篇文章:
更重要的是,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]
| B | Batch size(批大小) | 一次输入多少张图像,比如 B=8 表示同时处理 8 张图 |
| 4 | Sequence Length(序列长度) | 每个“视觉句子”由 4 个“视觉词语” 组成,这 4 个词被排成一个短序列 |
| 64 | Feature Dimension(特征维度) | 每个“词”用一个 64 维向量 表示其语义和外观信息 |
总输出仍为: 64 × 128 × 128
作用:增强局部细节(如边缘、纹理)。
步骤 4:聚合 → 得到句子表示
将每个句子的 4 个词平均池化成 1 个向量,得到句子级序列:[B, 4096, 64] , 即 64 × 64 × 64。
| B | Batch size(批大小) | 一次处理多少张图像 |
| 4096 | Token 数量(句子数量) | 每张图被划分为 64 × 64 = 4096 个“视觉句子” |
| 64 | 特征维度 | 每个“句子”用一个 64 维向量 表示其整体语义 |
[B, 4096, 64] | 序列形式:适合送入 RWKV(处理一维序列) |
64 × 64 × 64 | 2D 特征图形式: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 × 32 | Zig-RiR Block + Merge | 256 × 16 × 16 |
| L₃ | 256 × 16 × 16 | Zig-RiR Block + Merge | 512 × 8 × 8 |
| L₄ | 512 × 8 × 8 | Zig-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 |
| Stem | 64 × 128 × 128 |
| L₁ 后 | 64 × 128 × 128 |
| Merge 1 | 128 × 64 × 64 |
| Merge 2 | 256 × 32 × 32 |
| Merge 3 | 512 × 16 × 16 |
| L₄ 输出 | 512 × 8 × 8 |
| 解码后 | 32 × 512 × 512 |
| 输出 | 2 × 512 × 512 |
更多推荐

所有评论(0)