DeepQTMT: 深度学习加速VVC帧内编码CU划分 —— 论文深度解读与源码拆解

论文 :IEEE TIP 2021 | 作者 :Tianyi Li, Mai Xu 等 | 机构 :北京航空航天大学

链接arXiv:2006.13125 | 代码raulkviana/MSE-CNN-Implementations | 数据库tianyili2017/CPIV


目录

  1. 背景与问题动机
  2. [VVC QTMT划分机制深度解析](#VVC QTMT划分机制深度解析)
  3. 大规模数据库与数据洞察
  4. MSE-CNN网络架构详解
  5. 自适应损失函数设计
  6. 多阈值决策机制
  7. 开源代码拆解
  8. 实验结果与消融分析
  9. 总结与思考

1. 背景与问题动机

1.1 VVC编码:效率与复杂度的矛盾

2020年,联合视频专家组(JVET)正式发布了新一代视频编码标准------多功能视频编码 (Versatile Video Coding, VVC/H.266)。相比其前身HEVC/H.265,VVC在编码效率上取得了质的飞跃:在同等主观质量下,码率可降低约50%

然而,这巨大的编码增益是以指数级增长的编码复杂度 为代价的------VTM参考软件在帧内(All-Intra)模式下的编码复杂度是HEVC的平均18倍 。其中最核心的瓶颈在于:QTMT(四叉树+多类型树)CU划分 占据了编码总时间的超过97% Tissier et al., MMSP 2019

💡 核心矛盾: VVC通过引入极其灵活的QTMT划分结构获得了编码增益,但这种灵活性需要编码器通过穷举式RDO搜索来探索所有可能的划分组合,导致编码时间呈数量级增长,严重阻碍了VVC的实际工业落地。

1.2 已有方案与局限

在DeepQTMT提出之前,加速VVC划分的方案主要分为两类:

类别 方法 代表工作 局限性
启发式方法 利用纹理同质性、空间相关性等中间编码特征建立统计模型 Fu et al. (ICME 2019), Yang et al. (TCSVT 2020) 严重依赖手工特征提取,泛化性差,在不同视频内容上表现不稳定
数据驱动方法 使用CNN等模型自动学习CU划分规律 Jin et al. (VCIP 2017), Wang et al. (ICIP 2018) 仅支持简单的QTBT结构,无法处理VVC的QTMT多类型划分,缺乏对RD代价的联合优化

1.3 DeepQTMT的四大贡献

  1. 大规模数据库 :建立了CPIH-Intra数据库(后序发展为CPIV),包含充足且多样化的QTMT CU划分模式,为数据驱动的VVC加速研究奠定基础。
  2. 多阶段出口CNN(MSE-CNN) :提出了一种带有提前退出(early-exit)机制的模型架构,与QTMT多层次结构天然契合------深度越深的CU,特征越丰富,决策越精确。
  3. 自适应损失函数 :在交叉熵损失中融入RD代价(Rate-Distortion cost),使得模型不仅关注划分模式的分类准确率,同时也敏感于错误分类对编码效率的实际损害。
  4. 多阈值决策方案 :通过在不同划分阶段设置不同的置信度阈值,实现了编码复杂度与RD性能之间的精细调优,用户可根据需求选择"faster"或"moderate"模式。

实验结果:在VTM-7.0上,编码时间减少44.65%∼66.88% ,BD-BR仅增加1.322%∼3.188%,显著超越所有现有方法。


2. VVC QTMT划分机制深度解析

2.1 从HEVC到VVC:划分粒度的质变

理解DeepQTMT,首先要理解VVC的QTMT划分机制为什么如此复杂。我们先来看HEVC和VVC的根本区别:
#mermaid-svg-ji3n5SNLjsyZt131{font-family:"trebuchet ms",verdana,arial,sans-serif;font-size:16px;fill:#333;}@keyframes edge-animation-frame{from{stroke-dashoffset:0;}}@keyframes dash{to{stroke-dashoffset:0;}}#mermaid-svg-ji3n5SNLjsyZt131 .edge-animation-slow{stroke-dasharray:9,5!important;stroke-dashoffset:900;animation:dash 50s linear infinite;stroke-linecap:round;}#mermaid-svg-ji3n5SNLjsyZt131 .edge-animation-fast{stroke-dasharray:9,5!important;stroke-dashoffset:900;animation:dash 20s linear infinite;stroke-linecap:round;}#mermaid-svg-ji3n5SNLjsyZt131 .error-icon{fill:#552222;}#mermaid-svg-ji3n5SNLjsyZt131 .error-text{fill:#552222;stroke:#552222;}#mermaid-svg-ji3n5SNLjsyZt131 .edge-thickness-normal{stroke-width:1px;}#mermaid-svg-ji3n5SNLjsyZt131 .edge-thickness-thick{stroke-width:3.5px;}#mermaid-svg-ji3n5SNLjsyZt131 .edge-pattern-solid{stroke-dasharray:0;}#mermaid-svg-ji3n5SNLjsyZt131 .edge-thickness-invisible{stroke-width:0;fill:none;}#mermaid-svg-ji3n5SNLjsyZt131 .edge-pattern-dashed{stroke-dasharray:3;}#mermaid-svg-ji3n5SNLjsyZt131 .edge-pattern-dotted{stroke-dasharray:2;}#mermaid-svg-ji3n5SNLjsyZt131 .marker{fill:#333333;stroke:#333333;}#mermaid-svg-ji3n5SNLjsyZt131 .marker.cross{stroke:#333333;}#mermaid-svg-ji3n5SNLjsyZt131 svg{font-family:"trebuchet ms",verdana,arial,sans-serif;font-size:16px;}#mermaid-svg-ji3n5SNLjsyZt131 p{margin:0;}#mermaid-svg-ji3n5SNLjsyZt131 .label{font-family:"trebuchet ms",verdana,arial,sans-serif;color:#333;}#mermaid-svg-ji3n5SNLjsyZt131 .cluster-label text{fill:#333;}#mermaid-svg-ji3n5SNLjsyZt131 .cluster-label span{color:#333;}#mermaid-svg-ji3n5SNLjsyZt131 .cluster-label span p{background-color:transparent;}#mermaid-svg-ji3n5SNLjsyZt131 .label text,#mermaid-svg-ji3n5SNLjsyZt131 span{fill:#333;color:#333;}#mermaid-svg-ji3n5SNLjsyZt131 .node rect,#mermaid-svg-ji3n5SNLjsyZt131 .node circle,#mermaid-svg-ji3n5SNLjsyZt131 .node ellipse,#mermaid-svg-ji3n5SNLjsyZt131 .node polygon,#mermaid-svg-ji3n5SNLjsyZt131 .node path{fill:#ECECFF;stroke:#9370DB;stroke-width:1px;}#mermaid-svg-ji3n5SNLjsyZt131 .rough-node .label text,#mermaid-svg-ji3n5SNLjsyZt131 .node .label text,#mermaid-svg-ji3n5SNLjsyZt131 .image-shape .label,#mermaid-svg-ji3n5SNLjsyZt131 .icon-shape .label{text-anchor:middle;}#mermaid-svg-ji3n5SNLjsyZt131 .node .katex path{fill:#000;stroke:#000;stroke-width:1px;}#mermaid-svg-ji3n5SNLjsyZt131 .rough-node .label,#mermaid-svg-ji3n5SNLjsyZt131 .node .label,#mermaid-svg-ji3n5SNLjsyZt131 .image-shape .label,#mermaid-svg-ji3n5SNLjsyZt131 .icon-shape .label{text-align:center;}#mermaid-svg-ji3n5SNLjsyZt131 .node.clickable{cursor:pointer;}#mermaid-svg-ji3n5SNLjsyZt131 .root .anchor path{fill:#333333!important;stroke-width:0;stroke:#333333;}#mermaid-svg-ji3n5SNLjsyZt131 .arrowheadPath{fill:#333333;}#mermaid-svg-ji3n5SNLjsyZt131 .edgePath .path{stroke:#333333;stroke-width:2.0px;}#mermaid-svg-ji3n5SNLjsyZt131 .flowchart-link{stroke:#333333;fill:none;}#mermaid-svg-ji3n5SNLjsyZt131 .edgeLabel{background-color:rgba(232,232,232, 0.8);text-align:center;}#mermaid-svg-ji3n5SNLjsyZt131 .edgeLabel p{background-color:rgba(232,232,232, 0.8);}#mermaid-svg-ji3n5SNLjsyZt131 .edgeLabel rect{opacity:0.5;background-color:rgba(232,232,232, 0.8);fill:rgba(232,232,232, 0.8);}#mermaid-svg-ji3n5SNLjsyZt131 .labelBkg{background-color:rgba(232, 232, 232, 0.5);}#mermaid-svg-ji3n5SNLjsyZt131 .cluster rect{fill:#ffffde;stroke:#aaaa33;stroke-width:1px;}#mermaid-svg-ji3n5SNLjsyZt131 .cluster text{fill:#333;}#mermaid-svg-ji3n5SNLjsyZt131 .cluster span{color:#333;}#mermaid-svg-ji3n5SNLjsyZt131 div.mermaidTooltip{position:absolute;text-align:center;max-width:200px;padding:2px;font-family:"trebuchet ms",verdana,arial,sans-serif;font-size:12px;background:hsl(80, 100%, 96.2745098039%);border:1px solid #aaaa33;border-radius:2px;pointer-events:none;z-index:100;}#mermaid-svg-ji3n5SNLjsyZt131 .flowchartTitleText{text-anchor:middle;font-size:18px;fill:#333;}#mermaid-svg-ji3n5SNLjsyZt131 rect.text{fill:none;stroke-width:0;}#mermaid-svg-ji3n5SNLjsyZt131 .icon-shape,#mermaid-svg-ji3n5SNLjsyZt131 .image-shape{background-color:rgba(232,232,232, 0.8);text-align:center;}#mermaid-svg-ji3n5SNLjsyZt131 .icon-shape p,#mermaid-svg-ji3n5SNLjsyZt131 .image-shape p{background-color:rgba(232,232,232, 0.8);padding:2px;}#mermaid-svg-ji3n5SNLjsyZt131 .icon-shape .label rect,#mermaid-svg-ji3n5SNLjsyZt131 .image-shape .label rect{opacity:0.5;background-color:rgba(232,232,232, 0.8);fill:rgba(232,232,232, 0.8);}#mermaid-svg-ji3n5SNLjsyZt131 .label-icon{display:inline-block;height:1em;overflow:visible;vertical-align:-0.125em;}#mermaid-svg-ji3n5SNLjsyZt131 .node .label-icon path{fill:currentColor;stroke:revert;stroke-width:revert;}#mermaid-svg-ji3n5SNLjsyZt131 :root{--mermaid-font-family:"trebuchet ms",verdana,arial,sans-serif;} VVC: QTMT
QT
BT_H
BT_V
TT_H
TT_V
NS
128x128 CTU
强制QT分割
64x64 CU
4个32x32
2个64x32
2个32x64
1:2:1 水平三分
1:2:1 垂直三分
不分割
HEVC: 仅四叉树
128x128 CTU
64x64 CU
32x32 CU
16x16 CU
8x8 CU

图1: HEVC与VVC的CU划分方式对比 --- VVC引入了五种非四叉树划分类型

具体来说,VVC的QTMT结构包含 6种划分模式

索引 模式 符号 描述 示意图
0 不分割 (Non-Split) NS 当前CU不再继续划分
1 四叉树 (Quad-Tree) QT 等分成4个相同大小的正方形子块
2 水平二叉树 (BT_H) HBT 水平方向等分成2个矩形子块 ⬜⬜
3 垂直二叉树 (BT_V) VBT 垂直方向等分成2个矩形子块 ⬛⬛
4 水平三叉树 (TT_H) HTT 水平方向1:2:1分成3个子块 ⬜⬜⬜
5 垂直三叉树 (TT_V) VTT 垂直方向1:2:1分成3个子块 ⬛⬛⬛

QTMT 六阶段递归划分流程:

复制代码
Stage 1: 128×128 CTU → [强制QT→4个64×64]
Stage 2: 64×64 CU   → [6模式分类]
Stage 3: 32×32 CU   → [6模式分类]
Stage 4: 16×16 CU   → [6模式分类]
Stage 5: 8×8 CU     → [6模式分类]
Stage 6: 4×4 CU     → [6模式分类]

Stage 1 为确定性QT(128→64×64),Stage 2~6 为模型预测阶段

2.2 为什么QTMT的复杂度如此之高?

VTM编码器在每帧图像的每个CTU(Coding Tree Unit,128×128)上,从Stage 1到Stage 6对每个CU递归地进行RDO搜索,尝试所有可行的划分方式。对于一个CTU来说,可能的划分路径是组合爆炸级别的------每个阶段都有最多6种选择,且各子块的划分又独立进行。

关键观察:论文通过统计分析发现,不同大小CU的划分模式分布差异巨大。

  1. 64×64 CU:非分割(NS)占多数(~36%),说明大块CU倾向于不细分
  2. 32×32 CU:四叉树(QT)占比最高(~35%),这是最常见的大块划分方式
  3. 16×16 CU:N种划分模式分布相对均匀,决策难度最大
  4. 8×8 CU:非分割(NS)和二叉树(HBT/VBT)占主导,因为小CU难以支持TT划分

这直接引出一个设计洞察:不同尺寸的CU需要不同的处理策略------不能用一个统一模型来解决所有划分问题。


3. 大规模数据库与数据洞察

3.1 CPIH-Intra / CPIV 数据库

DeepQTMT的第一个贡献是建立了大规模VVC CU划分数据库。数据库基本信息:

属性 详情
原始视频来源 RAISE, Xiph.org, CDVL, 自拍数据等
视频数量 2,000个原始序列(训练集)
分辨率范围 256×256 到 2048×1536
编码器 VTM-7.0, All-Intra配置
QP值 22, 27, 32, 37(覆盖全部质量范围)
标注信息 最优划分模式 + 6种模式的RD代价
测试集 JCT-VC标准测试序列22条(Class A1~E)

🔑 数据集的独特价值: 不仅记录了每个CU的最优划分模式 (VTM RDO结果),还记录了所有6种划分模式的完整RD代价。这让模型训练时可以知道"选错A比选错B代价更大"------这是后续自适应损失函数设计的数据基础。

3.2 数据集获取

数据集在GitHub开源:tianyili2017/CPIV,包含完整的CU划分标注和RD代价信息。每条数据记录包含:CTU的Y分量像素值、CU位置坐标、CU尺寸、最优划分模式标签、以及6种模式的RD代价值。


4. MSE-CNN网络架构详解

4.1 整体思路:多阶段 + 提前退出

MSE-CNN(Multi-Stage Exit CNN)的核心思想可以概括为一句话:网络架构的设计应当反映问题的自然层次结构
#mermaid-svg-se731DZMmrc8gGSy{font-family:"trebuchet ms",verdana,arial,sans-serif;font-size:16px;fill:#333;}@keyframes edge-animation-frame{from{stroke-dashoffset:0;}}@keyframes dash{to{stroke-dashoffset:0;}}#mermaid-svg-se731DZMmrc8gGSy .edge-animation-slow{stroke-dasharray:9,5!important;stroke-dashoffset:900;animation:dash 50s linear infinite;stroke-linecap:round;}#mermaid-svg-se731DZMmrc8gGSy .edge-animation-fast{stroke-dasharray:9,5!important;stroke-dashoffset:900;animation:dash 20s linear infinite;stroke-linecap:round;}#mermaid-svg-se731DZMmrc8gGSy .error-icon{fill:#552222;}#mermaid-svg-se731DZMmrc8gGSy .error-text{fill:#552222;stroke:#552222;}#mermaid-svg-se731DZMmrc8gGSy .edge-thickness-normal{stroke-width:1px;}#mermaid-svg-se731DZMmrc8gGSy .edge-thickness-thick{stroke-width:3.5px;}#mermaid-svg-se731DZMmrc8gGSy .edge-pattern-solid{stroke-dasharray:0;}#mermaid-svg-se731DZMmrc8gGSy .edge-thickness-invisible{stroke-width:0;fill:none;}#mermaid-svg-se731DZMmrc8gGSy .edge-pattern-dashed{stroke-dasharray:3;}#mermaid-svg-se731DZMmrc8gGSy .edge-pattern-dotted{stroke-dasharray:2;}#mermaid-svg-se731DZMmrc8gGSy .marker{fill:#333333;stroke:#333333;}#mermaid-svg-se731DZMmrc8gGSy .marker.cross{stroke:#333333;}#mermaid-svg-se731DZMmrc8gGSy svg{font-family:"trebuchet ms",verdana,arial,sans-serif;font-size:16px;}#mermaid-svg-se731DZMmrc8gGSy p{margin:0;}#mermaid-svg-se731DZMmrc8gGSy .label{font-family:"trebuchet ms",verdana,arial,sans-serif;color:#333;}#mermaid-svg-se731DZMmrc8gGSy .cluster-label text{fill:#333;}#mermaid-svg-se731DZMmrc8gGSy .cluster-label span{color:#333;}#mermaid-svg-se731DZMmrc8gGSy .cluster-label span p{background-color:transparent;}#mermaid-svg-se731DZMmrc8gGSy .label text,#mermaid-svg-se731DZMmrc8gGSy span{fill:#333;color:#333;}#mermaid-svg-se731DZMmrc8gGSy .node rect,#mermaid-svg-se731DZMmrc8gGSy .node circle,#mermaid-svg-se731DZMmrc8gGSy .node ellipse,#mermaid-svg-se731DZMmrc8gGSy .node polygon,#mermaid-svg-se731DZMmrc8gGSy .node path{fill:#ECECFF;stroke:#9370DB;stroke-width:1px;}#mermaid-svg-se731DZMmrc8gGSy .rough-node .label text,#mermaid-svg-se731DZMmrc8gGSy .node .label text,#mermaid-svg-se731DZMmrc8gGSy .image-shape .label,#mermaid-svg-se731DZMmrc8gGSy .icon-shape .label{text-anchor:middle;}#mermaid-svg-se731DZMmrc8gGSy .node .katex path{fill:#000;stroke:#000;stroke-width:1px;}#mermaid-svg-se731DZMmrc8gGSy .rough-node .label,#mermaid-svg-se731DZMmrc8gGSy .node .label,#mermaid-svg-se731DZMmrc8gGSy .image-shape .label,#mermaid-svg-se731DZMmrc8gGSy .icon-shape .label{text-align:center;}#mermaid-svg-se731DZMmrc8gGSy .node.clickable{cursor:pointer;}#mermaid-svg-se731DZMmrc8gGSy .root .anchor path{fill:#333333!important;stroke-width:0;stroke:#333333;}#mermaid-svg-se731DZMmrc8gGSy .arrowheadPath{fill:#333333;}#mermaid-svg-se731DZMmrc8gGSy .edgePath .path{stroke:#333333;stroke-width:2.0px;}#mermaid-svg-se731DZMmrc8gGSy .flowchart-link{stroke:#333333;fill:none;}#mermaid-svg-se731DZMmrc8gGSy .edgeLabel{background-color:rgba(232,232,232, 0.8);text-align:center;}#mermaid-svg-se731DZMmrc8gGSy .edgeLabel p{background-color:rgba(232,232,232, 0.8);}#mermaid-svg-se731DZMmrc8gGSy .edgeLabel rect{opacity:0.5;background-color:rgba(232,232,232, 0.8);fill:rgba(232,232,232, 0.8);}#mermaid-svg-se731DZMmrc8gGSy .labelBkg{background-color:rgba(232, 232, 232, 0.5);}#mermaid-svg-se731DZMmrc8gGSy .cluster rect{fill:#ffffde;stroke:#aaaa33;stroke-width:1px;}#mermaid-svg-se731DZMmrc8gGSy .cluster text{fill:#333;}#mermaid-svg-se731DZMmrc8gGSy .cluster span{color:#333;}#mermaid-svg-se731DZMmrc8gGSy div.mermaidTooltip{position:absolute;text-align:center;max-width:200px;padding:2px;font-family:"trebuchet ms",verdana,arial,sans-serif;font-size:12px;background:hsl(80, 100%, 96.2745098039%);border:1px solid #aaaa33;border-radius:2px;pointer-events:none;z-index:100;}#mermaid-svg-se731DZMmrc8gGSy .flowchartTitleText{text-anchor:middle;font-size:18px;fill:#333;}#mermaid-svg-se731DZMmrc8gGSy rect.text{fill:none;stroke-width:0;}#mermaid-svg-se731DZMmrc8gGSy .icon-shape,#mermaid-svg-se731DZMmrc8gGSy .image-shape{background-color:rgba(232,232,232, 0.8);text-align:center;}#mermaid-svg-se731DZMmrc8gGSy .icon-shape p,#mermaid-svg-se731DZMmrc8gGSy .image-shape p{background-color:rgba(232,232,232, 0.8);padding:2px;}#mermaid-svg-se731DZMmrc8gGSy .icon-shape .label rect,#mermaid-svg-se731DZMmrc8gGSy .image-shape .label rect{opacity:0.5;background-color:rgba(232,232,232, 0.8);fill:rgba(232,232,232, 0.8);}#mermaid-svg-se731DZMmrc8gGSy .label-icon{display:inline-block;height:1em;overflow:visible;vertical-align:-0.125em;}#mermaid-svg-se731DZMmrc8gGSy .node .label-icon path{fill:currentColor;stroke:revert;stroke-width:revert;}#mermaid-svg-se731DZMmrc8gGSy :root{--mermaid-font-family:"trebuchet ms",verdana,arial,sans-serif;} Stage 4: MseCnnStgX
Stage 3: MseCnnStgX
Stage 2: MseCnnStg1
Stage 1 (确定性)
NS
继续
NS
继续
...
128x128 CTU (Y分量)
强制QT分割
重叠卷积层

1→16通道, 3x3
条件卷积 Stage2

Residual Units x2
子网络

Conv+FC+Softmax
早退判断
条件卷积 Stage3

Residual Units
子网络 min=32
早退判断
条件卷积 Stage4

Residual Units
子网络 min=16
早退判断
停止划分
继续Stage 5, 6...

图2: MSE-CNN整体推理流程 --- 每个阶段都可以提前退出

4.2 核心组件拆解

(1) 重叠卷积层(Overlapping Convolution Layer)

CTU的亮度分量(128×128单通道)首先通过一个3×3重叠卷积层,将通道数从1扩展到16:

F 0 = PReLU ( Conv 3 × 3 ( X CTU , 1 → 16 ) ) \mathbf{F}0 = \text{PReLU}(\text{Conv}{3\times3}(\mathbf{X}_{\text{CTU}}, 1 \to 16)) F0=PReLU(Conv3×3(XCTU,1→16))

这里使用PReLU(Parametric ReLU)而非标准ReLU,因为PReLU的负半轴斜率是可学习的参数,对视频编码中的纹理特征有更好的适配性。

(2) 条件卷积(Conditional Convolution)

这是MSE-CNN最核心的设计。条件卷积的本质是根据CU的尺寸动态调整残差单元的数量

n r = ⌊ log ⁡ 2 ( a p / a c ) ⌋ n_r = \lfloor \log_2(a_p / a_c) \rfloor nr=⌊log2(ap/ac)⌋

其中 a p a_p ap 是父CU的最小轴长度, a c a_c ac 是当前CU的最小轴长度。例如:

  • 128×128 → 64×64 子CU: a p = 128 , a c = 64 a_p=128, a_c=64 ap=128,ac=64 → n r = log ⁡ 2 ( 128 / 64 ) = 1 n_r = \log_2(128/64) = 1 nr=log2(128/64)=1 个残差单元
  • 128×128 → 32×32 子CU: a p = 128 , a c = 32 a_p=128, a_c=32 ap=128,ac=32 → n r = log ⁡ 2 ( 128 / 32 ) = 2 n_r = \log_2(128/32) = 2 nr=log2(128/32)=2 个残差单元

💡 设计原理: 小的CU需要更多的卷积层来提取足够的上下文信息(因为像素更少),而大的CU纹理信息已经足够,不需要过度处理。这种自适应深度的设计避免了"一刀切"的固定网络深度带来的冗余计算。

每个残差单元的结构是经典的ResNet风格跳跃连接:

F out = PReLU ( F in + Conv 3 × 3 ( PReLU ( Conv 3 × 3 ( F in ) ) ) ) \mathbf{F}{\text{out}} = \text{PReLU}(\mathbf{F}{\text{in}} + \text{Conv}{3\times3}(\text{PReLU}(\text{Conv}{3\times3}(\mathbf{F}_{\text{in}})))) Fout=PReLU(Fin+Conv3×3(PReLU(Conv3×3(Fin))))

复制代码
Fin → Conv3x3+PReLU → Conv3x3 → [+] → PReLU → Fout
          ↑                            |
          └──────── Shortcut ──────────┘

图3: 残差单元详细结构 --- 跳跃连接确保梯度传播

(3) QP半掩码(QP Half-Mask)

量化参数(QP)直接影响编码后的图像质量和CU划分决策。为了让网络感知QP,作者设计了一种精巧的QP半掩码机制

q ~ = Q P 51 + 0.5 \tilde{q} = \frac{QP}{51} + 0.5 q~=51QP+0.5

在前向传播中,将特征图沿通道维度均分为两半,后半部分乘以归一化后的 q ~ \tilde{q} q~ 值:

F ′ = Concat ( F 0 : C / 2 , q ~ ⋅ F C / 2 : C ) \mathbf{F}' = \text{Concat}(\mathbf{F}{0:C/2},\ \tilde{q} \cdot \mathbf{F}{C/2:C}) F′=Concat(F0:C/2, q~⋅FC/2:C)

💡 为什么是"一半"? 如果所有通道都乘以QP,网络可能会过度依赖QP信号而忽略纹理特征。只修改一半通道,可以在"QP感知"和"纹理感知"之间取得平衡。这种设计类似于特征调制(Feature Modulation)中的FiLM层思想。

(4) 多尺度子网络(Multi-Scale Sub-Networks)

Stage 2的子网络处理固定的64×64 CU,而Stage 3+的MseCnnStgX需要处理多种形状的CU。为此,作者设计了形状自适应池化 + 分类器结构:

CU最小轴长度 预处理卷积 分类器结构
64 直接送入子网络 Conv(4×4)→Conv(4×4)→FC(128→8)→FC(8→6)+Softmax
32 自适应kernel/stride的卷积(如32×64→4×8) PReLU→Conv(4×4)→Conv(2×2)→FC(128→64)→FC(64→6)+Softmax
16 自适应kernel/stride的卷积 PReLU→Conv(2×2)→Conv(2×2)→FC(64→64)→FC(64→6)+Softmax
8 自适应kernel/stride的卷积(如8×8→4×4) PReLU→Conv(2×2)→FC(32→16)→FC(16→6)+Softmax
4 自适应kernel/stride的卷积 PReLU→Conv(2×2)→FC(32→16)→FC(16→6)+Softmax

核心思路:使用步长等于kernel size的卷积一次性将所有CU统一到相同空间尺寸(1×1或2×2),然后送入全连接分类器。

4.3 迁移学习训练策略

MSE-CNN的各阶段不是独立训练的,而是采用级联迁移学习策略:

  1. 训练Stage 1&2 :在64×64 CU数据上训练MseCnnStg1
  2. 迁移到Stage 3 :将Stage 2的条件卷积权重复制到Stage 3的MseCnnStgX作为初始化,然后在32×32 CU数据上微调
  3. 迁移到Stage 4:将Stage 3的权重作为Stage 4的初始化,在16×16级别的CU数据上微调
  4. 依次类推到Stage 5和Stage 6

这种策略利用了低层纹理特征在不同阶段之间的可迁移性,加速了深层Stage的训练收敛。


5. 自适应损失函数设计

5.1 损失函数的两个目标

DeepQTMT的损失函数需要同时满足两个目标:

  1. 分类准确:预测正确的划分模式
  2. RD代价敏感:当预测错误时,偏好RD代价更低的错误答案

例如(纯假设),如果GT是QT,但模型错误地预测成了HBT或VTT,两种错误的RD代价可能不同------我们当然希望模型犯"代价更小的错误"。

因此,论文提出了组合损失函数:

L = L C E + β ⋅ L R D \mathcal{L} = \mathcal{L}{CE} + \beta \cdot \mathcal{L}{RD} L=LCE+β⋅LRD

5.2 修正交叉熵损失 L C E \mathcal{L}_{CE} LCE

数据集中各类划分模式的分布极不均衡(如64×64 CU中NS占比远高于VTT),直接使用标准交叉熵会导致模型偏向多数类。为此,引入类别权重

L C E = − 1 N ∑ n = 1 N ∑ m = 1 6 ( 1 p m ) α ⋅ y n , m ⋅ log ⁡ ( y ^ n , m ) \mathcal{L}{CE} = -\frac{1}{N}\sum{n=1}^{N} \sum_{m=1}^{6} \left(\frac{1}{p_m}\right)^\alpha \cdot y_{n,m} \cdot \log(\hat{y}_{n,m}) LCE=−N1n=1∑Nm=1∑6(pm1)α⋅yn,m⋅log(y^n,m)

其中:

  • N N N:batch中的样本数
  • m ∈ { 0 , 1 , 2 , 3 , 4 , 5 } m \in \{0,1,2,3,4,5\} m∈{0,1,2,3,4,5}:6种划分模式
  • p m p_m pm:第m类在训练集中的出现比例
  • α \alpha α:控制惩罚力度的超参数(论文中 α = 0.5 \alpha=0.5 α=0.5)
  • y n , m y_{n,m} yn,m:one-hot标签
  • y ^ n , m \hat{y}_{n,m} y^n,m:模型预测概率

权重项 ( 1 / p m ) α (1/p_m)^\alpha (1/pm)α 确保少数类获得更大的损失惩罚,缓解类别不均衡问题。

5.3 RD代价损失 L R D \mathcal{L}_{RD} LRD

这是论文最具创新性的设计之一。当模型将GT划分模式 m ∗ m^* m∗ 错误预测为模式 m m m 时,RD代价损失为:

L R D = 1 N ∑ n = 1 N ∑ m = 1 6 y ^ n , m ⋅ ( r n , m r n , min ⁡ − 1 ) \mathcal{L}{RD} = \frac{1}{N}\sum{n=1}^{N} \sum_{m=1}^{6} \hat{y}{n,m} \cdot \left(\frac{r{n,m}}{r_{n,\min}} - 1\right) LRD=N1n=1∑Nm=1∑6y^n,m⋅(rn,minrn,m−1)

其中:

  • r n , m r_{n,m} rn,m:第n个样本选择模式m时的RD代价
  • r n , min ⁡ = min ⁡ m r n , m r_{n,\min} = \min_m r_{n,m} rn,min=minmrn,m:所有模式中的最小RD代价(即最优模式的RD代价)

解读 :当 y ^ n , m \hat{y}{n,m} y^n,m 较大(模型强烈预测模式m)而 r n , m / r n , min ⁡ r{n,m}/r_{n,\min} rn,m/rn,min 也较大(该模式的RD代价远大于最优)时,损失会很大。这迫使模型避免给出高置信度的"昂贵错误"。

⚠️ 实践发现: 开源实现(raulkviana)的实验表明, L R D \mathcal{L}_{RD} LRD 项在训练中的实际贡献有限,最终模型仅用修正交叉熵训练也达到了良好效果。这说明RD代价的区分度在某些情况下可能不够显著,或者 β \beta β 参数的调节需要更精细的策略。


6. 多阈值决策机制

6.1 从软输出到硬决策

MSE-CNN的每个Stage输出一个6维Softmax概率向量 y ^ \hat{\mathbf{y}} y^。最简单的方式是取argmax:

m pred = arg ⁡ max ⁡ m y ^ m m_{\text{pred}} = \arg\max_m \hat{y}_m mpred=argmmaxy^m

但argmax只输出一个模式------如果模型预测NS的概率是0.51,QT是0.49,argmax强制选NS,可能在边缘情况下频繁出错。

6.2 多阈值决策方案

作者提出了更灵活的方案:对于每个Stage s ∈ { 2 , 3 , 4 , 5 , 6 } s \in \{2,3,4,5,6\} s∈{2,3,4,5,6},设置一个阈值 θ s \theta_s θs。将所有满足以下条件的预测模式都作为候选:

y ^ m ≥ θ s ⋅ max ⁡ m ′ y ^ m ′ \hat{y}m \geq \theta_s \cdot \max{m'} \hat{y}_{m'} y^m≥θs⋅m′maxy^m′

然后进行早期退出

  • 如果候选模式中包含NS(不分割),则当前CU停止划分
  • 否则,对候选模式中概率最高的模式(排除NS后)进行划分

阈值设计的关键特性:

  • 分阶段阈值:不同Stage使用不同阈值,因为各Stage的预测准确率不同
  • 阈值越大→候选越少→编码越快→但RD质量损失越大
  • 论文提供了"moderate"和"faster"两套阈值配置
配置 Stage 2 Stage 3 Stage 4 Stage 5 Stage 6 ∆T BD-BR
moderate 0.3 0.5 0.5 0.7 0.7 -44.65% 1.322%
faster 0.3 0.5 0.5 0.7 0.9 -66.88% 3.188%

7. 开源代码拆解

论文的开源实现由 Raul Kevin Viana 在GitHub维护:

7.1 核心类架构

#mermaid-svg-TJt9tz8ZNolVMAzE{font-family:"trebuchet ms",verdana,arial,sans-serif;font-size:16px;fill:#333;}@keyframes edge-animation-frame{from{stroke-dashoffset:0;}}@keyframes dash{to{stroke-dashoffset:0;}}#mermaid-svg-TJt9tz8ZNolVMAzE .edge-animation-slow{stroke-dasharray:9,5!important;stroke-dashoffset:900;animation:dash 50s linear infinite;stroke-linecap:round;}#mermaid-svg-TJt9tz8ZNolVMAzE .edge-animation-fast{stroke-dasharray:9,5!important;stroke-dashoffset:900;animation:dash 20s linear infinite;stroke-linecap:round;}#mermaid-svg-TJt9tz8ZNolVMAzE .error-icon{fill:#552222;}#mermaid-svg-TJt9tz8ZNolVMAzE .error-text{fill:#552222;stroke:#552222;}#mermaid-svg-TJt9tz8ZNolVMAzE .edge-thickness-normal{stroke-width:1px;}#mermaid-svg-TJt9tz8ZNolVMAzE .edge-thickness-thick{stroke-width:3.5px;}#mermaid-svg-TJt9tz8ZNolVMAzE .edge-pattern-solid{stroke-dasharray:0;}#mermaid-svg-TJt9tz8ZNolVMAzE .edge-thickness-invisible{stroke-width:0;fill:none;}#mermaid-svg-TJt9tz8ZNolVMAzE .edge-pattern-dashed{stroke-dasharray:3;}#mermaid-svg-TJt9tz8ZNolVMAzE .edge-pattern-dotted{stroke-dasharray:2;}#mermaid-svg-TJt9tz8ZNolVMAzE .marker{fill:#333333;stroke:#333333;}#mermaid-svg-TJt9tz8ZNolVMAzE .marker.cross{stroke:#333333;}#mermaid-svg-TJt9tz8ZNolVMAzE svg{font-family:"trebuchet ms",verdana,arial,sans-serif;font-size:16px;}#mermaid-svg-TJt9tz8ZNolVMAzE p{margin:0;}#mermaid-svg-TJt9tz8ZNolVMAzE g.classGroup text{fill:#9370DB;stroke:none;font-family:"trebuchet ms",verdana,arial,sans-serif;font-size:10px;}#mermaid-svg-TJt9tz8ZNolVMAzE g.classGroup text .title{font-weight:bolder;}#mermaid-svg-TJt9tz8ZNolVMAzE .cluster-label text{fill:#333;}#mermaid-svg-TJt9tz8ZNolVMAzE .cluster-label span{color:#333;}#mermaid-svg-TJt9tz8ZNolVMAzE .cluster-label span p{background-color:transparent;}#mermaid-svg-TJt9tz8ZNolVMAzE .cluster rect{fill:#ffffde;stroke:#aaaa33;stroke-width:1px;}#mermaid-svg-TJt9tz8ZNolVMAzE .cluster text{fill:#333;}#mermaid-svg-TJt9tz8ZNolVMAzE .cluster span{color:#333;}#mermaid-svg-TJt9tz8ZNolVMAzE .nodeLabel,#mermaid-svg-TJt9tz8ZNolVMAzE .edgeLabel{color:#131300;}#mermaid-svg-TJt9tz8ZNolVMAzE .edgeLabel .label rect{fill:#ECECFF;}#mermaid-svg-TJt9tz8ZNolVMAzE .label text{fill:#131300;}#mermaid-svg-TJt9tz8ZNolVMAzE .labelBkg{background:#ECECFF;}#mermaid-svg-TJt9tz8ZNolVMAzE .edgeLabel .label span{background:#ECECFF;}#mermaid-svg-TJt9tz8ZNolVMAzE .classTitle{font-weight:bolder;}#mermaid-svg-TJt9tz8ZNolVMAzE .node rect,#mermaid-svg-TJt9tz8ZNolVMAzE .node circle,#mermaid-svg-TJt9tz8ZNolVMAzE .node ellipse,#mermaid-svg-TJt9tz8ZNolVMAzE .node polygon,#mermaid-svg-TJt9tz8ZNolVMAzE .node path{fill:#ECECFF;stroke:#9370DB;stroke-width:1px;}#mermaid-svg-TJt9tz8ZNolVMAzE .divider{stroke:#9370DB;stroke-width:1;}#mermaid-svg-TJt9tz8ZNolVMAzE g.clickable{cursor:pointer;}#mermaid-svg-TJt9tz8ZNolVMAzE g.classGroup rect{fill:#ECECFF;stroke:#9370DB;}#mermaid-svg-TJt9tz8ZNolVMAzE g.classGroup line{stroke:#9370DB;stroke-width:1;}#mermaid-svg-TJt9tz8ZNolVMAzE .classLabel .box{stroke:none;stroke-width:0;fill:#ECECFF;opacity:0.5;}#mermaid-svg-TJt9tz8ZNolVMAzE .classLabel .label{fill:#9370DB;font-size:10px;}#mermaid-svg-TJt9tz8ZNolVMAzE .relation{stroke:#333333;stroke-width:1;fill:none;}#mermaid-svg-TJt9tz8ZNolVMAzE .dashed-line{stroke-dasharray:3;}#mermaid-svg-TJt9tz8ZNolVMAzE .dotted-line{stroke-dasharray:1 2;}#mermaid-svg-TJt9tz8ZNolVMAzE #compositionStart,#mermaid-svg-TJt9tz8ZNolVMAzE .composition{fill:#333333!important;stroke:#333333!important;stroke-width:1;}#mermaid-svg-TJt9tz8ZNolVMAzE #compositionEnd,#mermaid-svg-TJt9tz8ZNolVMAzE .composition{fill:#333333!important;stroke:#333333!important;stroke-width:1;}#mermaid-svg-TJt9tz8ZNolVMAzE #dependencyStart,#mermaid-svg-TJt9tz8ZNolVMAzE .dependency{fill:#333333!important;stroke:#333333!important;stroke-width:1;}#mermaid-svg-TJt9tz8ZNolVMAzE #dependencyStart,#mermaid-svg-TJt9tz8ZNolVMAzE .dependency{fill:#333333!important;stroke:#333333!important;stroke-width:1;}#mermaid-svg-TJt9tz8ZNolVMAzE #extensionStart,#mermaid-svg-TJt9tz8ZNolVMAzE .extension{fill:transparent!important;stroke:#333333!important;stroke-width:1;}#mermaid-svg-TJt9tz8ZNolVMAzE #extensionEnd,#mermaid-svg-TJt9tz8ZNolVMAzE .extension{fill:transparent!important;stroke:#333333!important;stroke-width:1;}#mermaid-svg-TJt9tz8ZNolVMAzE #aggregationStart,#mermaid-svg-TJt9tz8ZNolVMAzE .aggregation{fill:transparent!important;stroke:#333333!important;stroke-width:1;}#mermaid-svg-TJt9tz8ZNolVMAzE #aggregationEnd,#mermaid-svg-TJt9tz8ZNolVMAzE .aggregation{fill:transparent!important;stroke:#333333!important;stroke-width:1;}#mermaid-svg-TJt9tz8ZNolVMAzE #lollipopStart,#mermaid-svg-TJt9tz8ZNolVMAzE .lollipop{fill:#ECECFF!important;stroke:#333333!important;stroke-width:1;}#mermaid-svg-TJt9tz8ZNolVMAzE #lollipopEnd,#mermaid-svg-TJt9tz8ZNolVMAzE .lollipop{fill:#ECECFF!important;stroke:#333333!important;stroke-width:1;}#mermaid-svg-TJt9tz8ZNolVMAzE .edgeTerminals{font-size:11px;line-height:initial;}#mermaid-svg-TJt9tz8ZNolVMAzE .classTitleText{text-anchor:middle;font-size:18px;fill:#333;}#mermaid-svg-TJt9tz8ZNolVMAzE .label-icon{display:inline-block;height:1em;overflow:visible;vertical-align:-0.125em;}#mermaid-svg-TJt9tz8ZNolVMAzE .node .label-icon path{fill:currentColor;stroke:revert;stroke-width:revert;}#mermaid-svg-TJt9tz8ZNolVMAzE :root{--mermaid-font-family:"trebuchet ms",verdana,arial,sans-serif;} 继承
子网络中使用
子网络中使用
MseCnnStg1
+first_simple_conv: Sequential
+simple_conv_stg1, stg2
+sub_net: Sequential
+residual_unit_stg1(x, nr)
+residual_unit_stg2(x, nr)
+nr_calc(ac, ap)
+split(cu, coords, sizes, split)
+forward(cu, sizes, coords)
MseCnnStgX
+conv_32_64, conv_32_32
+conv_16_16, conv_16_32, conv_16_64
+conv_8_8, conv_8_16, conv_8_32, conv_8_64
+conv_4_4, conv_4_8, conv_4_16, conv_4_32
+sub_net_min_32, min_16, min_8, min_4
+residual_unit(x, nr)
+pass_through_subnet(x)
+forward(cu, ap, splits, sizes, coords)
QP_half_mask
+QP: int
+normalize_QP(QP)
+forward(feature_maps)
LossFunctionMSE
+beta: float
+MAX_RD: float
+forward(pred, labels, RD)
-get_min_RDs(RDs)
LossFunctionMSE_Ratios
+alpha: float
+pm: Tensor
+forward(pred, labels, RD)

图4: 核心类继承关系与依赖

7.2 模型初始化与推理流程

python 复制代码
import torch
import msecnn
import train_model_utils

# === 参数配置 ===
device = "cuda:0"
qp = 32

# === 模型初始化:5个模型对应6个Stage ===
stg1_2 = msecnn.MseCnnStg1(device=device, QP=qp).to(device)  # Stage 1&2
stg3   = msecnn.MseCnnStgX(device=device, QP=qp).to(device)  # Stage 3
stg4   = msecnn.MseCnnStgX(device=device, QP=qp).to(device)  # Stage 4
stg5   = msecnn.MseCnnStgX(device=device, QP=qp).to(device)  # Stage 5
stg6   = msecnn.MseCnnStgX(device=device, QP=qp).to(device)  # Stage 6
model = (stg1_2, stg3, stg4, stg5, stg6)

# === 加载预训练权重 ===
model = train_model_utils.load_model_parameters_eval(
    model, "model_coefficients/best_coefficients", device
)

# === Stage 1&2 前向传播 ===
CTU = torch.rand(1, 1, 128, 128).to(device)  # 128x128亮度CTU
cu_size = torch.tensor([[64, 64]]).to(device)
cu_pos = torch.tensor([[0, 0]]).to(device)
pred1_2, CUs, ac = model[0](CTU, cu_size, cu_pos)

7.3 QP半掩码源码解读

python 复制代码
class QP_half_mask(nn.Module):
    def __init__(self, QP=32):
        super(QP_half_mask, self).__init__()
        self.QP = QP

    def normalize_QP(self, QP):
        """
        QP归一化:q_tilde = QP / 51 + 0.5
        确保归一化后的值在 [0.5, 1.5] 范围内
        """
        q_tilde = QP / 51 + 0.5
        return q_tilde

    def forward(self, feature_maps):
        q_tilde = self.normalize_QP(self.QP)

        # 将特征图沿通道维度(C=1) 平分为两半
        half_num = feature_maps.size(1) // 2
        half_1, half_2 = torch.split(feature_maps, (half_num, half_num), dim=1)

        # 后半部分乘以归一化QP值
        half_2 = half_2 * q_tilde

        # 拼接回来
        new_feature_maps = torch.cat((half_1, half_2), dim=1)
        return new_feature_maps

💡 源码亮点: 该实现通过 torch.split 沿通道维度分半,然后用标量乘法广播到后半部分的所有通道和空间位置。这是一种简洁高效的特征调制方式------没有引入额外可训练参数。

7.4 条件卷积的残差单元源码

python 复制代码
def nr_calc(self, ac, ap):
    """
    根据父CU和当前CU的最小轴长计算所需残差单元数量
    ac: 当前CU最小轴长度
    ap: 父CU最小轴长度

    关键公式: nr = log2(ap / ac)
    """
    nr = 0
    if ac == 128:
        nr = 1  # CTU级别只用一个残差单元
    elif ap != 0:
        if 4 <= ac <= 64:
            nr = int(math.log2(ap / ac))
    return nr

def residual_unit(self, x, nr):
    """
    执行nr个残差单元
    每个单元: x = PReLU(x + Conv(PReLU(Conv(x))))
    """
    x_shortcut = x
    if nr == 1:
        x = self.simple_conv(x)              # Conv3x3 + PReLU
        x = self.simple_conv_no_activation(x) # Conv3x3(无激活)
        x = torch.add(x_shortcut, x)          # 跳跃连接
        x = self.activation_PRelu(x)          # PReLU
    elif nr == 2:
        # 第1个残差单元
        x = self.simple_conv(x)
        x = self.simple_conv_no_activation(x)
        x = torch.add(x_shortcut, x)
        x = self.activation_PRelu(x)
        # 第2个残差单元
        x_shortcut = x
        x = self.simple_conv2(x)
        x = self.simple_conv_no_activation2(x)
        x = torch.add(x_shortcut, x)
        x = self.activation_PRelu2(x)
    return x

7.5 自适应子网络路由源码

python 复制代码
def pass_through_subnet(self, x):
    """
    根据输入CU的形状自动选择合适的子网络
    核心逻辑:按CU的最小轴长度(min dim)路由到对应子网络
    """
    # 确保 Height <= Width (转置处理)
    if x.shape[-2] < x.shape[-1]:
        input_shape = (x.shape[-2], x.shape[-1])
    else:
        input_shape = (x.shape[-1], x.shape[-2])

    if min(input_shape) == 64:
        logits = self.sub_net(x)  # 使用Stage2子网络

    elif min(input_shape) == 32:
        # 根据宽高比选择不同kernel的预处理卷积
        if input_shape == (32, 64):
            logits = self.conv_32_64(x)
        else:
            logits = self.conv_32_32(x)
        # 统一送入min_32子网络
        logits = self.sub_net_min_32(logits)

    elif min(input_shape) == 16:
        if input_shape == (16, 64):
            logits = self.conv_16_64(x)
        elif input_shape == (16, 32):
            logits = self.conv_16_32(x)
        else:
            logits = self.conv_16_16(x)
        logits = self.sub_net_min_16(logits)

    # ... min=8, min=4 类似处理 ...

    return logits

7.6 损失函数源码解读

python 复制代码
class LossFunctionMSE(nn.Module):
    def __init__(self, use_mod_cross_entropy=True, beta=1):
        super(LossFunctionMSE, self).__init__()
        self.beta = beta
        self.MAX_RD = 1E10  # RD代价截断上限,防止梯度爆炸

    def forward(self, pred, labels, RD):
        # === 第1部分:修正交叉熵损失 ===
        if self.use_mod_cross_entropy:
            # L_CE = -mean(sum(y * log(pred + ε)))
            loss_CE = torch.mul(torch.log(pred + 1e-17), labels)
            loss_CE = torch.sum(loss_CE, dim=1)
            loss_CE = -torch.mean(loss_CE, dim=0)

        # === 第2部分:RD代价损失 ===
        # 1. 获取每个样本的最小RD代价
        min_RDs = self.get_min_RDs(RD)

        # 2. RD预处理:处理inf、零值和异常大值
        RD_mod = self.remove_inf_values(RD)
        RD_mod = self.remove_zero(RD_mod)          # 零→1E10
        RD_mod = self.remove_values_above(RD_mod, self.MAX_RD)  # >1E10→1E10

        # 3. 计算 RD损失 = pred * (RD_mod/min - 1)
        loss_RD = torch.mul(pred, torch.sub(torch.div(RD_mod, min_RDs), 1))
        loss_RD = self.remove_values_lower(loss_RD, 0, 0)  # 移除负值
        loss_RD = torch.sum(loss_RD, dim=1)
        loss_RD = torch.mean(loss_RD, dim=0)

        # 4. 损失截断:防止过大损失
        if loss_RD.item() > 20:
            temp = torch.div(20, loss_RD)
            loss_RD = torch.mul(loss_RD, temp)

        # === 组合损失 ===
        loss = torch.add(loss_CE, torch.mul(loss_RD, self.beta))
        return loss, loss_CE, loss_RD

7.7 迁移学习权重加载源码

python 复制代码
def load_model_stg_12_stg_3(model, path, dev):
    """
    Stage 1&2 → Stage 3 的迁移学习权重加载
    将Stage2的条件卷积权重复制到Stage3,实现知识迁移
    """
    # 加载Stage 1&2权重
    model[0].load_state_dict(stg_1_2)

    # 将Stage2的条件卷积层权重复制到Stage3对应层
    with torch.no_grad():
        model[1].simple_conv[0].weight.copy_(stg_1_2["simple_conv_stg2.0.weight"])
        model[1].simple_conv[0].bias.copy_(stg_1_2["simple_conv_stg2.0.bias"])
        model[1].simple_conv[1].weight.copy_(stg_1_2["simple_conv_stg2.1.weight"])
        model[1].simple_conv_no_activation[0].weight.copy_(
            stg_1_2["simple_conv_no_activation_stg2.0.weight"])
        # ... 更多层复制 ...

7.8 多阈值决策源码

python 复制代码
def obtain_best_modes(rs, pred):
    """
    多阈值候选模式筛选
    rs: 阈值
    pred: 模型输出概率

    规则: 所有满足 pred >= rs * max(pred) 的模式都作为候选
    """
    # 获取每个预测的最大概率值
    y_max = torch.reshape(torch.max(pred, dim=1)[0], shape=(-1, 1))
    y_max_rs = y_max * rs  # 动态阈值 = 阈值系数 × 最大概率

    # 筛选所有超过动态阈值的模式
    search_RD_logic = pred >= y_max_rs
    search_RD = torch.nonzero(
        torch.squeeze(search_RD_logic.int()), as_tuple=False
    )
    return search_RD

8. 实验结果与消融分析

8.1 主实验结果

在VTM-7.0 All-Intra配置下,使用JCT-VC标准22条测试序列(Class A1~E),评估编码时间节省(∆T)和BD-BR/BD-PSNR:

配置 BD-BR (%) BD-PSNR (dB) ∆T @ QP22 ∆T @ QP27 ∆T @ QP32 ∆T @ QP37
moderate 1.322 -0.053 -46.77% -44.68% -43.77% -43.39%
faster 3.188 -0.134 -66.88% -65.67% -63.05% -59.57%

🎯 关键结论: moderate模式以仅**1.322%的BD-BR代价换来了44.65%**的编码时间节省,这一trade-off在视频编码领域是极其优秀的。faster模式在3.188% BD-BR下达到66.88%时间节省,适合实时性要求极高的场景。

8.2 与现有方法的对比

方法 BD-BR (%) ∆T (%) 类型
Fu et al. (ICME 2019) 1.45 -38.2 启发式
Yang et al. (TCSVT 2020) 1.60 -43.5 启发式
Amestoy et al. (TIP 2020) 1.70 -44.0 机器学习
DeepQTMT (moderate) 1.322 -44.65 深度学习
DeepQTMT (faster) 3.188 -66.88 深度学习

8.3 消融实验(Ablation Study)

消融实验从最简单的版本逐步添加MSE-CNN的关键组件,验证每个设计的有效性:

实验 多阶段 RD代价 分阶段阈值 BD-BR ∆T@Q22
Ablation 1 (SSE-CNN) 6.539% -59.11%
Ablation 2 3.571% -65.48%
Ablation 3 3.328% -66.33%
Ablation 4 (faster) 3.188% -66.88%

从Ablation 1到Ablation 4,BD-BR从6.539%逐步降至3.188%(同时时间节省更大),每个组件都对性能有正向贡献。

🔬 最有启发的结果: 从Ablation 1(单阶段CNN)到Ablation 2(多阶段MSE-CNN),BD-BR直接从6.539%降到3.571%,缩水近一半。这说明多阶段设计是MSE-CNN最核心的性能来源------它充分利用了VVC划分的层次结构特性。

8.4 单独模式预测消融

作者还做了有趣的消融:如果MSE-CNN只预测单一划分模式(如只预测QT),效果如何?

结果表明:只预测QT时BD-BR约2.5%但∆T不到-35%;只预测某个BT/TT模式时虽然∆T可达-55%但BD-BR高达6%以上。这证明同时建模多种划分模式是必要的------因为只优化一种模式能缩减的时间空间有限,而错误代价很大。


9. 总结与思考

9.1 论文核心价值

  1. 问题驱动架构设计:MSE-CNN不是"为用CNN而用CNN",而是深刻理解了VVC QTMT的层次结构特点后,为每一层量身定制的网络结构------条件卷积对应CU尺寸多样性、多阶段对应划分层次、早退机制对应递归终止。
  2. 损失函数的工业思维:将RD代价融入损失函数体现了"不仅关注准确率,更关注错误代价"的工程思维------这在视频编码这种对性能极度敏感的场景中尤为重要。
  3. 可控的复杂度-RD trade-off:多阈值方案让用户可以根据实际需求在"快"和"好"之间自由调节,极大提升了方法的实用性。

9.2 局限性分析

  1. 仅支持All-Intra模式:实际应用中,Random Access和Low Delay模式更为常见,扩展到帧间编码需要处理运动信息等额外维度。
  2. GPU推理开销:虽然减少了编码时间,但引入了CNN推理的额外计算,在纯CPU环境下优势可能缩小。
  3. QP依赖:每个QP值需要独立训练一个模型(或使用QP mask感知),增加了部署复杂度。
  4. 开源实现的差距:开源代码中RD损失项的效果存疑,说明论文中的某些细节可能在实际复现时需要重新验证。

9.3 可借鉴的设计范式

DeepQTMT的几个设计模式可以迁移到其他视频编码加速任务中:

设计模式 适用场景 核心思想
多阶段出口 具有天然层次结构的决策问题 早期阶段快速排除低概率选项
条件计算 输入维度多变的任务 根据输入属性动态调整网络深度
代价敏感损失 不同错误的实际代价差异大的分类问题 将领域知识编码为损失函数
分阶段阈值 需要在多个粒度级别做决策的系统 精粗粒度不同置信度要求

9.4 未来方向

论文提出了几个有前景的扩展方向:

  • 扩展到帧间编码(Inter-mode),处理运动信息
  • 加速VVC编码器其他组件(如帧内角度选择、运动矢量估计)
  • 网络加速技术(量化、剪枝、知识蒸馏
  • FPGA硬件部署,实现实时的深度学习辅助编码

参考文献

  1. T. Li, M. Xu, R. Tang, Y. Chen and Q. Xing, "DeepQTMT: A Deep Learning Approach for Fast QTMT-based CU Partition of Intra-mode VVC," IEEE Transactions on Image Processing, vol. 30, pp. 5377-5390, 2021.
  2. A. Tissier et al., "Complexity reduction opportunities in the future VVC intra encoder," IEEE MMSP, 2019.
  3. R. K. Viana, "MSE-CNN Implementations," GitHub, 2022.
  4. T. Li, "CPIV: A large-scale database for QTMT-based CU partition of VVC," GitHub.
  5. JVET, "VVC Test Model (VTM)," Fraunhofer HHI.
  6. T. Fu et al., "Fast CU partitioning algorithm for H.266/VVC intra-frame coding," IEEE ICME, 2019.
  7. H. Yang et al., "Low-complexity CTU partition structure decision for VVC," IEEE TCSVT, 2020.
  8. T. Amestoy et al., "Tunable VVC frame partitioning based on lightweight machine learning," IEEE TIP, 2020.

技术博客 · 2026年8月2日 · 基于 IEEE TIP 2021 论文及 raulkviana/MSE-CNN-Implementations 源码

相关推荐
shxjnpl1 小时前
Qwen3-ASR 从 PyTorch 迁移到 vLLM:一次信创环境下的推理路径改造实录
人工智能·pytorch·vllm
阿童木写作2 小时前
Python实现Temu图片批量翻译自动化教程
运维·人工智能·python·自动化
冬奇Lab2 小时前
代码库知识库系列(06):把调用图编码进 Embedding——结构增强有效,但不够
人工智能
冬奇Lab3 小时前
开源项目第175期:Buzz — Jack Dorsey 的 Block 用 Nostr 重新定义团队协作,AI Agent 拥有自己的加密身份
人工智能·开源·资讯
AI分享猿3 小时前
游戏原画与建筑灵感:AI图像生成如何服务前期设计
人工智能·游戏
字节跳动视频云技术团队3 小时前
为什么 AI 视频,需要“懂生成”的画质增强
人工智能
后端小肥肠4 小时前
我做了个能一键搭建个人工作台的 Skill,已开源
人工智能·aigc·agent
码云之上4 小时前
AI Agent 工程化总览篇:从 Prompt 到 Harness
前端·人工智能
2501_926978334 小时前
以说明书 DNA 为模板——完整 AGI 的结构图景
前端·人工智能·经验分享·笔记·ai写作