基于深度学习的风格转换通常有两条实现路线:一种是以Gatys等人提出的神经风格迁移为代表的慢速优化方法,它不训练生成器,而是直接对输入图像做梯度下降;另一种是Johnson等人提出的感知损失前馈网络,训练一个编码器-残差块-解码器结构的生成器,推理时一次前向即可得到风格化结果。两种路线虽然训练方式不同,但共享一个关键组件——用预训练VGG网络提取特征,并在特征空间里定义内容损失与风格损失。理解这个特征空间如何被切分和利用,是看懂风格转换网络结构的基础。

一、VGG19特征提取:深层管内容,浅层管风格
VGG19是一个经典的卷积网络,最初用于图像分类。风格转换不会使用它的全连接分类层,只保留卷积部分。经验上,内容损失通常从conv4_2这一层提取,因为深层特征已经丢弃了像素级细节,只保留空间布局和语义信息;风格损失则从conv1_1、conv2_1、conv3_1、conv4_1、conv5_1这五个浅层到中层响应中提取,它们对纹理和颜色更敏感。
内容损失很直接:把内容图和生成图通过VGG同一层得到的特征图做均方误差。风格损失则要先把特征图转成Gram矩阵。Gram矩阵描述了一个特征图内部各通道之间的相关性,它不记录空间位置,因此天然适合表达纹理、笔触这类与位置无关的视觉属性。计算时把特征图从C × H × W展平成C × (H*W),再让这个矩阵与自己的转置相乘,最后除以元素总数。
import torch
import torch.nn as nn
import torch.nn.functional as F
def gram_matrix(feature_map):
# feature_map形状: (batch, C, H, W)
batch_size, channels, height, width = feature_map.size()
features = feature_map.view(batch_size, channels, height * width)
gram = torch.bmm(features, features.transpose(1, 2))
return gram / (channels * height * width)
def content_loss(content_features, generated_features):
return F.mse_loss(generated_features, content_features)
def style_loss(style_gram, generated_features):
generated_gram = gram_matrix(generated_features)
return F.mse_loss(generated_gram, style_gram)
注意torch.bmm要求输入是三维张量,所以先保留batch维度。如果只处理单张图像,也可以使用torch.mm。预先对风格图像计算好每个目标层的Gram矩阵,训练时直接和生成图的Gram矩阵比较即可。
二、快速风格迁移的生成器结构:下采样、残差块、上采样
慢速优化方法虽然能生成很好的结果,但每次换风格都要重新优化一张图,速度太慢。快速风格迁移把生成过程交给一个前馈卷积网络。生成器通常由三部分组成:编码器用步长卷积缩小空间尺寸,残差块在保持尺寸的同时完成风格特征融合,解码器用转置卷积恢复分辨率。输出层使用Tanh把像素值限制在-1到1之间。
一个典型的轻量结构是:输入3通道图像,先经过一个9×9卷积和两个步长为2的卷积,把空间尺寸缩小到原来的四分之一;然后堆叠5个残差块;随后用两个转置卷积上采样回原始尺寸;最后接一个9×9卷积输出3通道。所有非输出层后面都使用InstanceNorm和ReLU。
这里有个容易踩的坑:不要用BatchNorm替换InstanceNorm。BatchNorm会按照一个小批次样本计算均值和方差,如果批次内包含风格差异较大的图,统计量会被彼此污染。InstanceNorm只对单张图的每个通道做归一化,更适合风格迁移这种要求保持单图统计特性的任务。另外,反射填充比零填充更能减少边缘伪影。
import torch
import torch.nn as nn
class ResidualBlock(nn.Module):
def __init__(self, channels):
super(ResidualBlock, self).__init__()
self.conv1 = nn.Conv2d(channels, channels, 3, 1, 1, padding_mode='reflect')
self.in1 = nn.InstanceNorm2d(channels)
self.conv2 = nn.Conv2d(channels, channels, 3, 1, 1, padding_mode='reflect')
self.in2 = nn.InstanceNorm2d(channels)
def forward(self, x):
identity = x
out = torch.relu(self.in1(self.conv1(x)))
out = self.in2(self.conv2(out))
return identity + out
class FastStyleGenerator(nn.Module):
def __init__(self, in_channels=3, out_channels=3, base_channels=32):
super(FastStyleGenerator, self).__init__()
self.down1 = nn.Conv2d(in_channels, base_channels, 9, 1, 4, padding_mode='reflect')
self.in1 = nn.InstanceNorm2d(base_channels)
self.down2 = nn.Conv2d(base_channels, base_channels * 2, 3, 2, 1)
self.in2 = nn.InstanceNorm2d(base_channels * 2)
self.down3 = nn.Conv2d(base_channels * 2, base_channels * 4, 3, 2, 1)
self.in3 = nn.InstanceNorm2d(base_channels * 4)
self.res_blocks = nn.Sequential(*[ResidualBlock(base_channels * 4) for _ in range(5)])
self.up1 = nn.ConvTranspose2d(base_channels * 4, base_channels * 2, 3, 2, 1, output_padding=1)
self.in4 = nn.InstanceNorm2d(base_channels * 2)
self.up2 = nn.ConvTranspose2d(base_channels * 2, base_channels, 3, 2, 1, output_padding=1)
self.in5 = nn.InstanceNorm2d(base_channels)
self.output_conv = nn.Conv2d(base_channels, out_channels, 9, 1, 4, padding_mode='reflect')
def forward(self, x):
x = torch.relu(self.in1(self.down1(x)))
x = torch.relu(self.in2(self.down2(x)))
x = torch.relu(self.in3(self.down3(x)))
x = self.res_blocks(x)
x = torch.relu(self.in4(self.up1(x)))
x = torch.relu(self.in5(self.up2(x)))
return torch.tanh(self.output_conv(x))
这个结构的目的不是像分类网络那样减少参数量提高语义抽象,而是在一个紧凑的特征空间里交换风格信息。残差块的恒等连接能保证生成器不会丢掉内容图的整体结构。若风格过强导致内容变形,可以减少残差块数量;若纹理过于平淡,可以适当增加残差块。
三、训练循环:用VGG计算感知损失
前馈生成器的训练目标不是让输出图和风格图逐像素接近,而是在VGG特征空间里保持内容接近、风格统计接近。训练前需要把VGG19的卷积部分加载出来,并冻结所有参数,只把它当作损失计算器。内容目标仍然从conv4_2提取,风格目标从5个层提取并提前计算Gram矩阵。总损失由内容损失、风格损失和总变差损失加权求和。
总变差损失用来抑制生成图中的高频噪点,它计算相邻像素之间的差异。训练图像一般在0到1之间,VGG输入前还需要做标准化,使用ImageNet的均值和标准差。优化器常用Adam,学习率设置成0.001或更小。每个训练步同时给生成器送入内容图和风格图,生成器只接收内容图作为输入,风格信息是通过损失函数间接注入的。
import torch
import torch.nn as nn
import torchvision.models as models
import torchvision.transforms as transforms
class VGGFeatureExtractor(nn.Module):
def __init__(self):
super(VGGFeatureExtractor, self).__init__()
vgg = models.vgg19(pretrained=True).features.eval()
self.slices = nn.ModuleList()
self.content_index = None
self.style_indices = []
current = 0
# 按模块索引截取VGG19卷积部分
for idx, layer in enumerate(vgg):
if isinstance(layer, nn.Conv2d):
current += 1
if current == 4 and self.content_index is None:
self.content_index = len(self.slices)
if current in [1, 2, 3, 4, 5] and isinstance(layer, nn.ReLU):
self.style_indices.append(len(self.slices))
self.slices.append(layer)
def forward(self, x):
outputs = []
for layer in self.slices:
x = layer(x)
outputs.append(x)
content = outputs[self.content_index]
styles = [outputs[i] for i in self.style_indices]
return content, styles
上面的截取逻辑只是为了说明如何从torchvision的VGG模型中拿到不同层的输出。实际项目中更常见的做法是直接用register_forward_hook注册钩子,或者按名称保存named_children里的层。训练循环里,先把内容图和风格图都经过特征提取器,得到目标内容特征和目标风格Gram,再对生成器输出重复一次前向,计算当前特征与目标特征之间的距离。
def train_step(generator, optimizer, vgg_extractor, content_img, style_img, style_weight, content_weight, tv_weight):
generator.train()
optimizer.zero_grad()
target_content, target_styles = vgg_extractor(content_img)
target_style_grams = [gram_matrix(s) for s in target_styles]
generated_img = generator(content_img)
generated_content, generated_styles = vgg_extractor(generated_img)
c_loss = content_loss(generated_content, target_content)
s_loss = 0.0
for target_gram, generated_style in zip(target_style_grams, generated_styles):
s_loss = s_loss + style_loss(target_gram, generated_style)
tv_loss = torch.mean(torch.abs(generated_img[:, :, :-1, :] - generated_img[:, :, 1:, :])) + \
torch.mean(torch.abs(generated_img[:, :, :, :-1] - generated_img[:, :, :, 1:]))
total_loss = content_weight * c_loss + style_weight * s_loss + tv_weight * tv_loss
total_loss.backward()
optimizer.step()
return total_loss.item()
注意上面代码中的反斜杠用于换行连接两个张量差,这是Python语法中的续行符,运行时会原样保留。训练初期内容损失下降较快,风格损失下降相对慢一些。可以把内容权重设为1,风格权重从高一点开始,比如1e5到1e7数量级,总变差权重取1e-6到1e-3,具体需要根据图像尺寸和VGG特征幅值调整。
四、网络结构差异与调参建议
慢速优化和快速前馈并不是替代关系。慢速方法更灵活,一张内容图配合一张风格图就能开始,风格强度可以通过迭代次数和权重精细控制;缺点是耗时,适合生成少量高质量结果。快速方法训练一次生成器需要较大的风格数据集,但推理速度快,部署到移动端或Web端更现实。
如果要做未配对风格转换,CycleGAN的生成器也采用编码器、残差块、解码器结构,但训练目标引入了循环一致性损失和对抗损失。此时网络不仅要重建内容,还要骗过判别器。生成器中的InstanceNorm同样重要,判别器通常由几层步长卷积组成,输出一个局部感知的真假判断。
调参方面,残差块数量、基础通道数和下采样次数共同决定生成器的表达能力。基础通道太大会让模型推理变慢,太小会导致纹理不足。内容权重过高会让生成图接近原图,风格权重过高则可能产生重复纹理和伪影。建议先固定内容权重为1,在一个小尺寸图像上快速试验不同风格权重,再逐步提高分辨率训练。
风格转换神经网络结构Python深度学习修改时间:2026-09-25 09:03:24