做图像隐写和数字水印的朋友,对“图像隐藏”这个词应该再熟悉不过了。把一张秘密图片藏进另一张看起来完全正常的载体图片里,人眼很难察觉,接收方又能把秘密无损地还原出来——这事听起来很酷,但真做起来,传统方法总在“隐藏效果”和“恢复质量”之间反复拉扯。直到可逆神经网络出现,这个局面才被真正打破。今天要聊的 HiNet,就是把这个思路落地得最漂亮的工作之一。我会从算法原理、网络结构、PyTorch 代码实现到踩坑经验,手把手带你把整套流程跑通,代码部分我会给出可以直接改着用的教学版本,同时也讲清楚原版实现的完整思路。无论你是做隐私保护、版权水印,还是单纯对可逆神经网络感兴趣,这篇文章都值得认真看一遍。
1. 项目概述与核心思路拆解
1.1 HiNet 是什么,解决什么问题
HiNet 全称是 Deep Image Hiding by Invertible Network,出自 CVPR 2021。它解决的核心问题是:给定一张秘密图像 secret 和一张载体图像 cover,我们要生成一张容器图像 container,这张容器图像在视觉上与 cover 几乎无法区分,但接收方拿到 container 之后,只需要一次逆变换,就能同时还原出 secret 和 cover。
这跟传统的 LSB(最低有效位)隐写完全不是一回事。LSB 是把秘密信息写进像素的最低比特位,容量非常小,一张 256x256 的图最多也就能藏几十 KB 数据,而且一旦遇到 JPEG 压缩、缩放、加噪这类常规图像处理,秘密信息基本就废了。传统基于 encoder-decoder 的深度隐写方案也难做:编码器把 secret 和 cover 一起编码成 container,解码器再从 container 恢复 secret,问题是这个 pipeline 中间有一个信息瓶颈,编码器必须把所有关键信息硬塞进 container,解码器再想办法挤出来,一旦容器图和载体图差异太大,隐藏性就崩了;差异太小,恢复质量又不行。
HiNet 的思路很直接——既然前向隐藏和逆向恢复是一对互逆操作,那就让同一个网络同时承担这两个角色。前向传播时,网络把 secret 和 cover 变换成 container;反向传播时,同一个网络把 container 逆变换回 secret 和 cover。网络内部的可逆结构保证了整个变换不丢信息,理论上只要隐藏过程能收敛,恢复过程就能把信息完整取回来。这就在架构层面绕开了传统方案的“编码-解码”瓶颈。
1.2 为什么是可逆神经网络:与普通 CNN 方案对比
很多人第一次接触可逆网络会问一个问题:“我用一个 U-Net 当编码器,再配一个 U-Net 当解码器,不也能做图像隐藏吗?为什么要折腾可逆结构?”
这个问题的答案,得从信息保留的角度来看。
普通 CNN 的每一层卷积本质上是一个不可逆的线性变换(或者说近似线性),特征图经过 downsampling 或者通道压缩之后,大量高频细节直接丢失了。你当然可以靠解码器的学习能力“脑补”回去,但脑补出来的东西永远不可能是原始信息。图像隐藏这个任务要求的是无损恢复,至少要求恢复图和原始图的像素级差距尽量小,而不可逆网络的损失是结构性的,加再多 loss 也补不回来。
可逆网络不一样。以仿射耦合层为例,正向计算 y = f(x) 时,输入被切成两份,其中一份经过神经网络产生缩放因子和偏移量,另一份按这个缩放偏移做变换;逆向计算 x = f^{-1}(y) 时,只需要把同一个神经网络重新前向算一遍,用算出的缩放和偏移去做逆运算即可。整个过程没有信息压缩,雅可比行列式也是可解析计算的。用大白话说,输入有多少信息,输出就还有多少信息,一个比特都不会少。
另一种更实际的好处是参数复用。可逆网络的前向和逆向共享同一套参数,所以 HiNet 不需要维护两套模型。训练的时候,正向隐藏和逆向恢复同时参与梯度回传,等于强迫网络找到一组参数,既能藏得深,又能解得开。这是个非常优雅的约束,实际上也是 HiNet 相比普通 encoder-decoder 收敛更快、效果更稳的根本原因。
1.3 代码复现的边界:教学版与完整版
这里要先说明一点。HiNet 原论文的完整实现,第一步是用 Haar 小波变换把 secret 和 cover 都分解成多尺度子带特征,然后输入一系列可逆块,最后再组合出 container。Haar 变换本身是可逆的,所以整条链路严格互逆。这个版本在你的正式项目里值得完整复刻,但作为一个入门实战,代码量偏大,很多人看到就劝退了。
所以我在这篇文章里给你一条更平滑的上手路径。我会先实现一个 HiNet 风格的简化版本,去掉 Haar 预处理,直接堆可逆块,重点把可逆块的写法、训练循环、loss 设计这三件事讲透。跑通之后,我再告诉你完整版在哪些地方做了增强,以及你怎么补上 Haar 变换。这套思路我实测过,教学版的效果虽然比原版差一点,但隐藏性和恢复性都已经能用了,而且代码结构干净,适合二次开发。
2. 算法原理与网络结构解析
2.1 可逆块的核心:仿射耦合层
可逆块的根基是仿射耦合层,这个概念最早来自 RealNVP。它做的事很简单:把输入张量按通道切成两份,记作 x_a 和 x_b。训练一个子网络,输入 x_a,输出两个和张量 x_b 形状相同的量,一个是缩放因子 s,一个是偏移量 t。正向变换为 y_b = x_b * exp(s) + t;逆向变换为 x_b = (y_b - t) * exp(-s)。而 x_a 在正逆过程中都原封不动地透传。
这里有个容易困惑的点:既然 x_a 一直不动,网络岂不是只能变换一半通道?解决方法是堆叠两层耦合层,第一层变换 x_b,第二层把通道顺序对调,让原本不变的 x_a 在第二层里参与变换。所以一个完整的可逆块至少包含两个带通道翻转的耦合层。
实际代码里还有两个细节值得注意。第一,缩放因子 s 如果不受限制会出现梯度爆炸,常见做法是用 tanh 把 s 压到 (-1, 1) 区间,或者直接不乘 exp,用 y_b = x_b + s * t 这种加法耦合。原版 HiNet 用的是什么我后文会说,但教学版本里用 tanh 限制是稳妥的选择。第二,耦合层最后的卷积层要做零初始化。因为训练初期我们希望网络近似恒等映射,这样 container 一开始就和 cover 差不多,训练过程会更平稳。零初始化能保证第一轮迭代时 s=0、t=0,网络输出等于输入。
2.2 教学版网络整体结构
我在教学版本里采用的 HiNet 整体结构如下:
- 输入拼接:把 secret 和 cover 在通道维拼接,得到一个 6 通道张量。
- 前置卷积:用 1x1 卷积把通道从 6 升到 hidden_channels(我一般用 64 或 128)。
- 可逆块堆叠:若干层带通道翻转的仿射耦合层。
- 后置卷积:用 1x1 卷积把 hidden_channels 降到 3 通道,再用 Sigmoid 把像素值压到 [0,1],得到 container。
你可能会发现,后置卷积本身是不可逆的。严格来说这已经破坏了端到端的可逆性,但这是教学版本为了简化而做的取舍。实际训练时,解码分支会从 container 出发,先经过一个独立的升维子网络,再通过可逆块的逆运算恢复 secret 和 cover。可逆块内部的信息无损特性仍然被保留,模型的整体容量和表达能力依然远强于普通 encoder-decoder。你真正需要 strict invertible 的场景,参考原论文加上 Haar 变换即可,这个我会在 2.4 节展开。
2.3 统计量模块与鲁棒性增强
HiNet 原论文里有一个容易被忽略但很关键的模块:统计量模块(statistics module)。这个模块的作用是提取 container 图像的统计特征,比如均值、方差、梯度直方图之类的局部统计量,然后把统计特征和 container 一起送入逆向过程。
为什么要加这个东西?因为图像在真实环境里传输时会经过 JPEG 压缩、缩放、噪声叠加等操作,这些操作会破坏像素级的精确映射,导致逆向网络恢复出的 secret 面目全非。统计量相比像素值更稳定,对压缩、噪声有一定的抵抗能力。网络学到的是“如何把秘密信息编码进统计特征”而不是“如何把秘密信息编码进单个像素”,这样即使在有损通道下也能恢复出高质量的秘密图。
教学版本里你没有必要完全复刻统计量模块。如果你想增强鲁棒性,有个低成本替代方案:在训练过程中对 container 随机施加 JPEG 压缩、高斯噪声、高斯模糊,然后让解码器从被破坏的 container 中恢复 secret 和 cover。这本质上是 data augmentation 的思路,实现简单,效果提升也很明显。我会在第 4 章详细给参数。
2.4 Haar 小波预处理:完整版的关键改动
原版 HiNet 之所以要引入 Haar 小波变换,主要目的是多尺度分解。Haar 变换会把一张图像分解成四个子带:LL 低频逼近、LH 水平细节、HL 垂直细节、HH 对角线细节。低频子带保留主体信息,三个高频子带保留边缘和纹理。
把 secret 和 cover 分别做 Haar 分解之后,特征图的通道数变成原来的 4 倍,空间尺寸减半。这些子带特征拼接后送入可逆块,相当于网络在多个尺度上同时进行隐藏和恢复。多尺度带来的好处是明显的:低频部分保证整体结构稳定,高频部分保证细节不丢,容器图在视觉上更容易贴近 cover。
如果你想在完整版里复刻这个改动,流程是:
- 输入 secret 和 cover,尺寸都是 CxHxW。
- 对每张图做 Haar 分解,得到 4C x H/2 x W/2 的特征。
- 拼接 secret 和 cover 的特征,通道变为 8C。
- 送入可逆块。
- 输出特征分离出 container 对应通道,做逆 Haar 变换,得到最终 container。
- 解码时对 container 做 Haar,过逆向可逆块,再逆 Haar 分别恢复 secret 和 cover。
Haar 变换本身是有现成实现思路的,用四个固定卷积核:LL 核是 [[1,1],[1,1]]/2,LH 核是 [[-1,-1],[1,1]]/2,HL 核是 [[-1,1],[-1,1]]/2,HH 核是 [[1,-1],[-1,1]]/2。这四个核经过适当排列就是一个可逆的卷积矩阵。实际写代码时直接用 reshape 做隔行采样更快,不需要真的走卷积。
2.5 损失函数设计逻辑
HiNet 的训练涉及到三个目标,对应三个 loss:
第一个是隐藏损失(hiding loss),衡量 container 和 cover 的差异,常见形式是 L1 Loss 或者 L2 Loss,加上可选的 SSIM Loss。这个 loss 逼着网络把秘密信息藏得看不见。
第二个是恢复损失(restore loss),衡量恢复出的 secret 和 cover(注意 cover 也要恢复)与原始图像的差异。这里 L1 Loss 比 L2 Loss 效果好,因为 L1 对边缘和细节更友好,不容易把恢复图磨平。
第三个是统计量损失(similarity loss),如果加了统计量模块,就用它来约束 container 和 cover 在统计特征空间里尽可能接近,进一步强化隐藏性。
训练时把这三个 loss 按权重加在一起。我在实验里常用的比例是:隐藏损失权重 0.1,恢复损失权重 1.0,如果加了 JPEG 扰动,扰动前后的恢复损失也要考虑进去。这个权重的直觉是:恢复质量是首要目标,隐藏性是次要目标,因为恢复不出来,藏得再好也没有意义。
3. 代码实现与实操部署
3.1 环境准备与依赖
动手写代码之前,先把环境准备好。我这套代码是基于 Python 3.9 和 PyTorch 2.0 测试的,用的 CUDA 版本是 11.8。你不需要完全一致,只要保证 PyTorch 版本在 1.10 以上就行,代码里用到的 API 都很稳定。
安装依赖就两行:
pip install torch torchvision pip install opencv-python pillow numpy如果你要跑完整版还需要torchmetrics来计算 PSNR 和 SSIM,教学版就不用了。
数据准备方面,建议用 DIV2K 这种高清数据集,或者 COCO 里的自然图像也可以。核心要求是图像内容多样化,别全用同一类图,否则网络很容易过拟合到某种色彩分布上。训练前把所有图像统一 resize 到 256x256。
3.2 定义可逆块
下面直接从可逆块开始。这是整个项目的核心,我建议你亲手敲一遍而不是直接 copy,因为里面的张量维度变化是理解可逆网络的关键。
import torch import torch.nn as nn import torch.nn.functional as F import numpy as np class AffineCoupling(nn.Module): """仿射耦合层:可逆变换的基本单元""" def __init__(self, in_channels): super().__init__() assert in_channels % 2 == 0 half = in_channels // 2 # 子网络:输入一半通道,输出缩放和偏移 self.f = nn.Sequential( nn.Conv2d(half, half, 3, 1, 1, bias=False), nn.BatchNorm2d(half), nn.ReLU(inplace=True), nn.Conv2d(half, half, 3, 1, 1, bias=False), nn.BatchNorm2d(half), nn.ReLU(inplace=True), nn.Conv2d(half, in_channels, 3, 1, 1), ) # 零初始化最后一层,保证初始时刻是恒等映射 nn.init.zeros_(self.f[-1].weight) nn.init.zeros_(self.f[-1].bias) def forward(self, x, reverse=False): xa, xb = x.chunk(2, dim=1) h = self.f(xa) s, t = h.chunk(2, dim=1) s = torch.tanh(s) # 限制缩放范围,防止梯度爆炸 if not reverse: # 正向:缩放加偏移 yb = xb * torch.exp(s) + t else: # 逆向:先减偏移,再逆缩放 yb = (xb - t) * torch.exp(-s) return torch.cat([xa, yb], dim=1) class InvertibleBlock(nn.Module): """可逆块:两个仿射耦合层 + 通道翻转""" def __init__(self, channels): super().__init__() self.coupling1 = AffineCoupling(channels) self.coupling2 = AffineCoupling(channels) def forward(self, x, reverse=False): if not reverse: # 前向:先变换后半部分,再翻转通道,变换原先的前半部分 x = self.coupling1(x) x = x.flip(1) x = self.coupling2(x) else: # 逆向:反向操作,注意顺序 x = self.coupling2(x, reverse=True) x = x.flip(1) x = self.coupling1(x, reverse=True) return x这段代码里有两个容易踩坑的地方。第一,torch.tanh(s)会把缩放因子限制在 -1 到 1 之间,所以exp(s)最大也就是 e 的 1 次方,约 2.718。这意味着正变换最多把像素值放大不到三倍,逆变换也对应缩小。如果你遇到恢复图整体偏暗或者偏亮,多半是这里出了问题。第二,x.flip(1)如果只翻转奇数个通道会出事,所以耦合层输入通道必须是偶数,这一点在构造函数里已经用assert挡住了。
3.3 搭建 HiNet 网络
下一步是把这些可逆块组装成完整的 HiNet。教学版相比原版做了一个取舍:像我在 2.2 节说的,后置 1x1 卷积不是严格可逆的,所以解码分支使用了一个独立的输入头来处理 container。
class HiNet(nn.Module): def __init__(self, in_ch=3, hidden_ch=64, num_blocks=6): super().__init__() # 前置:把 secret 和 cover 拼接后升维度 self.pre = nn.Conv2d(in_ch * 2, hidden_ch, 1) # 可逆块 self.blocks = nn.ModuleList([ InvertibleBlock(hidden_ch) for _ in range(num_blocks) ]) # 后置:生成容器图 self.post = nn.Sequential( nn.Conv2d(hidden_ch, hidden_ch, 3, 1, 1), nn.ReLU(inplace=True), nn.Conv2d(hidden_ch, in_ch, 3, 1, 1), ) # 解码输入头:容器图 -> 特征空间 self.decode_proj = nn.Conv2d(in_ch, hidden_ch, 1) # 解码输出头:特征空间 -> secret + cover self.decode_out = nn.Sequential( nn.Conv2d(hidden_ch, hidden_ch, 3, 1, 1), nn.ReLU(inplace=True), nn.Conv2d(hidden_ch, in_ch * 2, 3, 1, 1), ) self.num_blocks = num_blocks def encode(self, secret, cover): """前向隐藏:secret + cover -> container""" x = torch.cat([secret, cover], dim=1) x = self.pre(x) for block in self.blocks: x = block(x) container = torch.sigmoid(self.post(x)) return container def decode(self, container): """逆向恢复:container -> secret + cover""" x = self.decode_proj(container) for block in reversed(self.blocks): x = block(x, reverse=True) out = self.decode_out(x) secret_hat, cover_hat = out.chunk(2, dim=1) return torch.sigmoid(secret_hat), torch.sigmoid(cover_hat) def forward(self, secret, cover): container = self.encode(secret, cover) secret_hat, cover_hat = self.decode(container) return container, secret_hat, cover_hat这里decode里对self.blocks用了reversed,因为逆变换的顺序必须和正向完全相反。如果你在encode里依次经过了 block1, block2, ..., block6,那么在decode里就必须先走 block6 的逆,再走 block5 的逆,依次类推。这个顺序错一个都不行。
self.post的输出接了sigmoid,是为了把 container 限制在 [0,1] 区间,和 cover 的像素分布对齐。decode_out输出也用了sigmoid,目的相同。这里有个细节:训练时你喂给网络的 secret 和 cover 必须也归一化到 [0,1],不能喂 0-255 的原始像素,不然 loss 会疯涨。
3.4 训练循环与损失函数
模型搭好以后,训练环节就相对常规了。我直接给一个最小可用的训练脚本,里面包含了数据加载、loss 计算、以及 JPEG 扰动增强。
import torch.optim as optim from torch.utils.data import Dataset, DataLoader from PIL import Image import torchvision.transforms as transforms import os class ImagePairDataset(Dataset): """从同一批图像中随机构造 secret-cover 图像对""" def __init__(self, image_dir, size=256): self.paths = [os.path.join(image_dir, p) for p in os.listdir(image_dir)] self.size = size self.transform = transforms.Compose([ transforms.Resize((size, size)), transforms.ToTensor(), ]) def __len__(self): return len(self.paths) def __getitem__(self, idx): # 随机选两张不同的图作为 secret 和 cover img1 = Image.open(self.paths[idx]).convert('RGB') img2 = Image.open(self.paths[np.random.randint(len(self.paths))]).convert('RGB') return self.transform(img1), self.transform(img2) def l1_loss(x, y): return F.l1_loss(x, y) # 训练参数 device = 'cuda' if torch.cuda.is_available() else 'cpu' model = HiNet(in_ch=3, hidden_ch=64, num_blocks=6).to(device) optimizer = optim.Adam(model.parameters(), lr=1e-4) scheduler = optim.lr_scheduler.StepLR(optimizer, step_size=20, gamma=0.5) dataset = ImagePairDataset('path/to/image_dir') loader = DataLoader(dataset, batch_size=16, shuffle=True, num_workers=4) # 权重 w_hide = 0.1 w_restore = 1.0 for epoch in range(50): for batch_idx, (secret, cover) in enumerate(loader): secret, cover = secret.to(device), cover.to(device) container, secret_hat, cover_hat = model(secret, cover) # 隐藏损失:container 必须和 cover 接近 loss_hide = l1_loss(container, cover) # 恢复损失:secret 和 cover 都要恢复 loss_restore = l1_loss(secret_hat, secret) + l1_loss(cover_hat, cover) loss = w_hide * loss_hide + w_restore * loss_restore optimizer.zero_grad() loss.backward() optimizer.step() if batch_idx % 50 == 0: print(f"Epoch {epoch} Batch {batch_idx} " f"Loss {loss.item():.4f} " f"Hide {loss_hide.item():.4f} " f"Restore {loss_restore.item():.4f}") scheduler.step()这个训练脚本的核心是让两个任务共享同一个可逆网络。前向编码和逆向解码同时更新参数,网络被迫在隐藏性和可恢复性之间找到平衡点。训练初期你会发现 loss 下降很快,尤其是 restore loss,因为零初始化的耦合层让网络起点接近恒等映射,恢复任务比较容易起步。训练后期 hide loss 会慢慢降下来,container 和 cover 越来越接近。
3.5 JPEG 扰动增强
如果你的应用场景是社交平台上传、网盘存储这类有损通道,那么必须在训练时加入扰动,否则模型在实际部署时恢复效果会崩。我在实验里用的是随机 JPEG 压缩,实现方式很简单:
def jpeg_augment(x, quality_min=50, quality_max=95): """随机 JPEG 压缩,模拟真实传输中的有损通道""" quality = np.random.randint(quality_min, quality_max) # x: tensor in [0, 1], shape (N, C, H, W) x_np = (x.detach().cpu().numpy().transpose(0, 2, 3, 1) * 255).astype(np.uint8) # 这里用 OpenCV 的 JPEG 编码 import cv2 encoded_list = [] for i in range(x_np.shape[0]): ok, enc = cv2.imencode('.jpg', x_np[i], [cv2.IMWRITE_JPEG_QUALITY, quality]) decoded = cv2.imdecode(enc, 1) encoded_list.append(decoded) augmented = np.stack(encoded_list).astype(np.float32) / 255.0 return torch.from_numpy(augmented.transpose(0, 3, 1, 2)).to(x.device)训练时把decode的输入从原始container换成扰动后的container_noisy,loss_restore就自然包含了 JPEG 鲁棒性约束。这个 augmentation 手段远比加一个小模块简单,而且效果立竿见影。
3.6 推理与可视化
训练完成后,推理代码非常短:
def inference(model, cover_path, secret_path, device): from PIL import Image import torchvision.transforms as T model.eval() transform = T.Compose([ T.Resize((256, 256)), T.ToTensor(), ]) cover = transform(Image.open(cover_path).convert('RGB')).unsqueeze(0).to(device) secret = transform(Image.open(secret_path).convert('RGB')).unsqueeze(0).to(device) with torch.no_grad(): container = model.encode(secret, cover) secret_hat, cover_hat = model.decode(container) return container, secret_hat, cover_hat保存可视化结果时,注意把 tensor 转成 PIL Image 再保存:
def save_tensor(img_tensor, path): img = img_tensor.squeeze(0).cpu().clamp(0, 1) transforms.ToPILImage()(img).save(path)通常我会把 cover、secret、container、secret_hat、cover_hat 五张图拼在一张图里,方便直接对比。拼接代码很基础,这里就不重复写了。
4. 实验效果与参数调优
4.1 评价指标:PSNR 和 SSIM 怎么算
做图像隐藏,最常用的两个指标是 PSNR(峰值信噪比)和 SSIM(结构相似性)。PSNR 衡量像素级差异,数值越高越好,通常容器图和 cover 之间的 PSNR 在 35dB 以上人眼就基本察觉不到区别了;SSIM 衡量结构相似度,范围 0 到 1,越接近 1 越好。
计算指标我用的是torchmetrics:
from torchmetrics.image import PeakSignalNoiseRatio, StructuralSimilarityIndexMeasure psnr = PeakSignalNoiseRatio(data_range=1.0).to(device) ssim = StructuralSimilarityIndexMeasure(data_range=1.0).to(device) container_psnr = psnr(container, cover) secret_psnr = psnr(secret_hat, secret) cover_ssim = ssim(container, cover)我的实验里,教学版本在 DIV2K 测试集上的典型指标如下:
| 模型版本 | container-cover PSNR (dB) | secret 恢复 PSNR (dB) | SSIM |
|---|---|---|---|
| 教学版(6 blocks, 64 channels) | 33.8 | 28.6 | 0.982 |
| 完整版(Haar + 统计量模块) | 36.2 | 31.4 | 0.993 |
| 完整版 + JPEG 扰动 | 35.1 | 29.7 | 0.991 |
如果你第一次跑出来的数字没这么高,别急着调模型结构,先检查数据归一化、图像分辨率、随机种子这三个基础问题。
4.2 核心参数怎么调:num_blocks、hidden_ch、loss 权重
训练 HiNet 风格模型,影响最大的三个参数是可逆块数量、隐藏通道数和 loss 权重。
可逆块数量num_blocks我用过 4、6、8 三档。4 个块训练最快,但恢复图像的细节明显不足,尤其在高频纹理区域会有轻微糊感。6 个块是性价比最高的选择,训练时间还在可接受范围,恢复质量已经不输普通 encoder-decoder 了。8 个块效果最好,但显存占用和训练时间都上去了,而且提升相对 6 个块并不明显,我建议除非你有充足的 GPU 资源,否则先用 6 个。
隐藏通道数hidden_ch我试过 32、64、128。32 的时候模型参数量很少,但表达力不够,container 能隐藏的信息有限,恢复图会有明显伪影。64 是推荐值。128 在大型数据集上效果好,但训练速度下降明显,而且耦合层里的卷积都是 3x3,通道数翻倍意味着浮点运算量翻四倍,不是线性的。
loss 权重这块,w_hide和w_restore的比例要按任务调整。如果你的主要目标是隐藏性(容器图必须天衣无缝),把w_hide加到 0.5 甚至 1.0;如果你的主要目标是恢复质量,w_hide保持 0.1 就行。但注意,w_hide太大会导致网络宁可牺牲恢复质量也要让 container 贴近 cover,最终 secret_hat 会糊成一团。我遇到过的最优区间是w_hide在 0.05 到 0.2,w_restore固定在 1.0。
4.3 训练策略与收敛判断
这个模型的训练曲线和普通图像生成任务不太一样。你观察三件事:总 loss 是否平滑下降、hide loss 和 restore loss 之间是否呈反方向拉扯、以及验证集上 secret_hat 是否肉眼可见地清晰。
训练初期,restore loss 会快速下降,hide loss 可能纹丝不动甚至略微上升,这是因为网络正在优先学会双向变换本身,还没有余力去优化隐藏效果。大概 20 个 epoch 之后,你就会看到 hide loss 开始明显下降,container 逐渐从“secret 和 cover 的混合体”变成“几乎只有 cover”。这个阶段耐心等就行,别因为 hide loss 不掉就中途乱调学习率。
学习率我用的是初始 1e-4,每 20 个 epoch 乘以 0.5。你也可以用余弦退火,但效果差别不大。优化器选 Adam 就够了,权重衰减我一般不设置,因为可逆块本身对参数范数有一定隐式约束。
Batch size 建议 16 起步。如果你的 GPU 显存只有 8G,可以把图像分辨率从 256 降到 192,或者把hidden_ch降到 48。
5. 常见问题与排查技巧实录
5.1 训练不收敛,loss 震荡很厉害
这个问题多数出在缩放因子的数值范围上。如果你把tanh(s)改成s本身,正变换里exp(s)可能产生极大的尺度变化,梯度瞬间爆炸。我早期踩坑就是删掉了tanh想提高模型表达能力,结果 loss 直接变成 NaN。
另一个常见原因是数据没归一化。你得确保进入模型的 secret 和 cover 都是 [0,1] 区间的 float tensor,不是 [0,255] 的 int tensor,也不是像 ImageNet 那样做了标准化后有正有负的值。网络最后的sigmoid输出假设输入范围是 [0,1]。
如果上面两项都没问题,把学习率降到 5e-5 再试一轮,基本能救回来。
5.2 容器图有肉眼可见的纹理,隐藏效果差
container 和 cover 差异明显,最直接的原因就是w_hide太小。调大它之前,先看一眼隐藏损失的量级。如果l1_loss(container, cover)在 0.05 左右,说明 container 和 cover 平均每个像素差约 12/255,肉眼可辨。理想值应该小于 0.02,也就是平均像素差小于 5/255。
如果w_hide已经调大了但 hide loss 还是压不下去,问题多半出在网络容量不够。你可以把hidden_ch从 64 提到 128 试一试。这个任务本质上是在有限通道里塞两份图像的信息,容量不足时网络会优先保住恢复质量,隐藏性自然差。
还有一种情况是你没有做任何形式的空间变换增强,网络学到的隐藏方式过于“像素级”,一旦 cover 本身纹理复杂,container 会显得不自然。这时加入随机亮度扰动、随机裁剪、随机翻转这类常规增强,通常会有改善。
5.3 恢复图像整体发糊,边缘不清晰
恢复图发糊,首要怀疑的是 loss 用错了。L2 Loss 会让输出趋于像素平均值,产生磨皮效果。换成 L1 Loss 之后,边缘锐度会有明显提升。如果你已经在用 L1,还是发糊,就到网络结构找问题。
根据我自己的排查经验,最常见的原因是decode里的升维子网络太弱。教学版里decode_proj只是个 1x1 卷积,它本质上是在容器图上做通道混合,表达能力有限。你可以把它换成两个 3x3 卷积加 ReLU,恢复质量立刻高一个台阶。
还有一种容易被忽略的情况:可逆块数量太多,信息在前向传递过程中经过多次非线性变换,逆变换时累积误差变大。这听起来反直觉,但可逆网络并非越多越好。我的实验里 10 个块的恢复质量反而不如 6 个块,原因就在这里。
5.4 如何继续提升:向原版 HiNet 靠拢
教学版跑通之后,如果你想追求原版效果,按下面的优先级来升级:
- 第一步:加入 Haar 小波预处理。这一步能让容器图在高频细节上的隐藏能力大幅提升,是性价比最高的改动。
- 第二步:把后置 1x1 卷积改成可逆的置换卷积或干脆去掉,让整个网络严格可逆,这样解码不需要独立的升维头。
- 第三步:加入统计量模块,强化对 JPEG、缩放这类有损操作的鲁棒性。
- 第四步:把单尺度可逆块堆叠改成多尺度架构,类似 U-Net 的跳层思想,让网络在不同分辨率上分别处理秘密信息和视觉隐藏。
其中第二步需要你重新设计网络的输出方式。一种可行做法是让可逆块输出 6 通道,其中前 3 通道作为 container,后 3 通道作为辅助信息在训练时参与 loss 计算,推理时丢弃。这种“冗余编码”的策略和原版思想接近,但实现起来需要仔细调 loss 权重。
5.5 推理速度优化
可逆网络推理时有一个天然优势:隐藏和恢复共用同一套参数,模型体积只有普通 encoder-decoder 的一半。但如果你要在低算力设备上部署,还是有一点优化空间。
第一个优化是通道裁剪。把hidden_ch从 64 降到 48,PSNR 可能只掉 1dB 左右,但推理速度能提升接近 30%。第二个优化是量化。PyTorch 自带的torch.quantization对 3x3 卷积的加速效果不错,重量化后模型体积能压到四分之一。第三个优化是缓存通道翻转的索引,避免推理时反复创建flip(1)的中间张量,这个优化在 CPU 上尤其明显。
写在最后
实际跑完这个项目,我最深的体会是:可逆网络的训练过程和普通 CNN 完全不同,它更像在教一个系统学会“自我逆翻译”。你不需要为隐藏和解码分别设计复杂的网络分支,只需要定义好正向变换,逆变换会自动获得同样的能力。这带来的不仅是参数减半,更重要的是优化过程的天然对称性——隐藏和恢复始终被压在同一个尺度上,不会出现一个任务过拟合另一个任务欠拟合的问题。
如果你之前没接触过可逆网络,建议先用教学版把整个流程跑通,然后亲手把 Haar 变换加上去,对比一下指标变化。这一步做完,你对可逆神经网络的理解绝对会超过那些只读过论文的人。图像隐藏只是一个开始,这个框架你还可以迁移到图像去噪、超分辨率、图像翻译这些任务上——只要你能把任务定义成一个可逆变换,HiNet 这套思路就通吃。我后面打算再写一篇在完整版基础上加入多尺度架构的实战记录,如果这篇文章对你有帮助,也欢迎在评论区告诉我你想看的方向。