论文 :IEEE TIP 2021 | 作者 :Tianyi Li, Mai Xu 等 | 机构 :北京航空航天大学
链接 :arXiv:2006.13125 | 代码 :raulkviana/MSE-CNN-Implementations | 数据库 :tianyili2017/CPIV
目录
- 背景与问题动机
- [VVC QTMT划分机制深度解析](#VVC QTMT划分机制深度解析)
- 大规模数据库与数据洞察
- MSE-CNN网络架构详解
- 自适应损失函数设计
- 多阈值决策机制
- 开源代码拆解
- 实验结果与消融分析
- 总结与思考
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的四大贡献
- 大规模数据库 :建立了CPIH-Intra数据库(后序发展为CPIV),包含充足且多样化的QTMT CU划分模式,为数据驱动的VVC加速研究奠定基础。
- 多阶段出口CNN(MSE-CNN) :提出了一种带有提前退出(early-exit)机制的模型架构,与QTMT多层次结构天然契合------深度越深的CU,特征越丰富,决策越精确。
- 自适应损失函数 :在交叉熵损失中融入RD代价(Rate-Distortion cost),使得模型不仅关注划分模式的分类准确率,同时也敏感于错误分类对编码效率的实际损害。
- 多阈值决策方案 :通过在不同划分阶段设置不同的置信度阈值,实现了编码复杂度与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的划分模式分布差异巨大。
- 64×64 CU:非分割(NS)占多数(~36%),说明大块CU倾向于不细分
- 32×32 CU:四叉树(QT)占比最高(~35%),这是最常见的大块划分方式
- 16×16 CU:N种划分模式分布相对均匀,决策难度最大
- 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的各阶段不是独立训练的,而是采用级联迁移学习策略:
- 训练Stage 1&2 :在64×64 CU数据上训练
MseCnnStg1 - 迁移到Stage 3 :将Stage 2的条件卷积权重复制到Stage 3的
MseCnnStgX作为初始化,然后在32×32 CU数据上微调 - 迁移到Stage 4:将Stage 3的权重作为Stage 4的初始化,在16×16级别的CU数据上微调
- 依次类推到Stage 5和Stage 6
这种策略利用了低层纹理特征在不同阶段之间的可迁移性,加速了深层Stage的训练收敛。
5. 自适应损失函数设计
5.1 损失函数的两个目标
DeepQTMT的损失函数需要同时满足两个目标:
- 分类准确:预测正确的划分模式
- 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维护:
- 仓库地址 :raulkviana/MSE-CNN-Implementations
- 许可证 :MIT | 框架:PyTorch
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 论文核心价值
- 问题驱动架构设计:MSE-CNN不是"为用CNN而用CNN",而是深刻理解了VVC QTMT的层次结构特点后,为每一层量身定制的网络结构------条件卷积对应CU尺寸多样性、多阶段对应划分层次、早退机制对应递归终止。
- 损失函数的工业思维:将RD代价融入损失函数体现了"不仅关注准确率,更关注错误代价"的工程思维------这在视频编码这种对性能极度敏感的场景中尤为重要。
- 可控的复杂度-RD trade-off:多阈值方案让用户可以根据实际需求在"快"和"好"之间自由调节,极大提升了方法的实用性。
9.2 局限性分析
- 仅支持All-Intra模式:实际应用中,Random Access和Low Delay模式更为常见,扩展到帧间编码需要处理运动信息等额外维度。
- GPU推理开销:虽然减少了编码时间,但引入了CNN推理的额外计算,在纯CPU环境下优势可能缩小。
- QP依赖:每个QP值需要独立训练一个模型(或使用QP mask感知),增加了部署复杂度。
- 开源实现的差距:开源代码中RD损失项的效果存疑,说明论文中的某些细节可能在实际复现时需要重新验证。
9.3 可借鉴的设计范式
DeepQTMT的几个设计模式可以迁移到其他视频编码加速任务中:
| 设计模式 | 适用场景 | 核心思想 |
|---|---|---|
| 多阶段出口 | 具有天然层次结构的决策问题 | 早期阶段快速排除低概率选项 |
| 条件计算 | 输入维度多变的任务 | 根据输入属性动态调整网络深度 |
| 代价敏感损失 | 不同错误的实际代价差异大的分类问题 | 将领域知识编码为损失函数 |
| 分阶段阈值 | 需要在多个粒度级别做决策的系统 | 精粗粒度不同置信度要求 |
9.4 未来方向
论文提出了几个有前景的扩展方向:
- 扩展到帧间编码(Inter-mode),处理运动信息
- 加速VVC编码器其他组件(如帧内角度选择、运动矢量估计)
- 网络加速技术(量化、剪枝、知识蒸馏)
- FPGA硬件部署,实现实时的深度学习辅助编码
参考文献
- 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.
- A. Tissier et al., "Complexity reduction opportunities in the future VVC intra encoder," IEEE MMSP, 2019.
- R. K. Viana, "MSE-CNN Implementations," GitHub, 2022.
- T. Li, "CPIV: A large-scale database for QTMT-based CU partition of VVC," GitHub.
- JVET, "VVC Test Model (VTM)," Fraunhofer HHI.
- T. Fu et al., "Fast CU partitioning algorithm for H.266/VVC intra-frame coding," IEEE ICME, 2019.
- H. Yang et al., "Low-complexity CTU partition structure decision for VVC," IEEE TCSVT, 2020.
- 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 源码