ControlNet 深度剖析:条件控制扩散模型的结构化控制系统

7001 字
35 分钟
ControlNet 深度剖析:条件控制扩散模型的结构化控制系统

1. ControlNet 的核心问题#

1.1 扩散模型缺乏可控性#

尽管 Stable Diffusion 等扩散模型能生成高质量图像,但”无条件生成”在实际应用中远远不够:

实际应用的需求:
✓ "生成一只猫" → 只需文本
✗ "生成一只坐着的猫" → 需要姿态控制
✗ "生成一只猫,放在我的书桌上" → 需要空间布局
✗ "按照我画的草图生成完整图像" → 需要边缘/轮廓控制
✗ "用我给的深度图作为参考" → 需要深度条件

核心矛盾:文本描述的是语义内容,但无法精确控制空间布局、物体姿态、边缘轮廓等结构性信息。

1.2 现有的四种失败方案#

ControlNet 论文(Adding Conditional Control to Text-to-Image Diffusion Models, Zhang et al., ICCV 2023)系统分析了已有的失败方案:

方案 1: 在 UNet 中加入额外输入通道
思想: 把条件图像拼接到潜变量通道
问题: ❌ 需要重新训练整个 UNet,显存爆炸
参数量: +86M 额外参数需要从头训练
方案 2: 只微调额外输入层
思想: 新增输入分支,只训练新层
问题: ❌ 新层参数随机初始化,微调效果差
原因: 随机初始化的层在前向传播中产生噪声,破坏预训练 UNet 的特征
方案 3: 注意力机制
思想: 用 cross-attention 注入条件
问题: ❌ 文本注意力不够强,结构条件信息丢失
原因: 文本 token 数量有限,无法精细编码空间信息
方案 4: 分类器引导
思想: 训练一个条件分类器,引导生成
问题: ❌ 需要额外的分类器,数据标注成本高
原因: 分类器的质量直接限制生成效果

1.3 ControlNet 的核心洞察#

把预训练的 UNet 视为强大的”主脑”,在旁边训练一个可拆卸的”小脑”——这个”小脑”学会理解额外的条件输入(边缘、姿态、深度等),并通过渐进式零卷积初始化保证训练稳定性——小脑产生的信号从零开始逐步增长,不破坏主脑的预训练知识。

1.4 论文信息#

论文: "Adding Conditional Control to Text-to-Image Diffusion Models"
作者:Lvmin Zhang, Maneesh Agrawala
单位: Stanford University
发表于: ICCV 2023
引用: > 5,000 次 ★★★★★
开源: https://github.com/lllyasviel/ControlNet
在线体验: https://modelscope.cn/studios/iic/ControlNet

1.5 一句话概括 ControlNet#

ControlNet 的核心是”复制并冻结 UNet 编码器 + 训练一个可拆卸的副本”——冻结的编码器保证预训练知识不被破坏,副本在条件输入的监督下学会精确的结构控制,然后通过零卷积初始化的跳跃连接把控制信号渐进注入 UNet 解码器,实现对图像生成的结构级控制,同时保留文本到图像的质量。

2. 架构设计:复制 + 冻结 + 零卷积#

2.1 整体架构#

┌────────────────────────────────────────────────────────────────┐
│ │
│ [条件图像 c] → 条件编码器 │
│ (Canny边缘/深度图/姿态) ↓ │
│ 可训练的副本分支 │
│ (ControlNet Block) │
│ ↓ │
│ [噪声潜变量 z_t] ───────────────────┐ │
│ │ ↓ │
│ [时间步 t] ─→ UNet 编码器 ─→ [冻结 ❄️] ──────────────────┐ │
│ ↓ │ │
│ UNet 解码器 │ │
│ ↑ │ │
│ 跳跃连接 ← ← ← ← ← ← ← ← ← ← ← ← ← ← ┘ │
│ (零卷积 ZeroConv) │
│ ↑ │
│ ControlNet 副本编码器 ← 条件图像 c │
│ (训练 🔥) │
│ │
│ [文本条件] ──→ CLIP Text Encoder ──→ 交叉注意力注入 │
│ │
└────────────────────────────────────────────────────────────────┘
关键设计:
1. UNet 编码器完整复制 → 冻结 → 预训练权重
2. 条件编码器: 与 UNet 编码器结构相同的副本 → 训练
3. 零卷积: 跳跃连接处的 1×1 卷积,初始为零 → 渐进增长

2.2 核心组件:ControlNet Block#

class ControlNetBlock(nn.Module):
"""
ControlNet 的基本模块。
与 UNet 的 ResBlock 结构相同,但包含跳跃连接的零卷积。
"""
def __init__(self, in_channels, out_channels, context_dim=None):
super().__init__()
# 主体: ResBlock
self.norm1 = nn.GroupNorm(32, in_channels)
self.conv1 = nn.Conv2d(in_channels, out_channels, 3, padding=1)
if context_dim is not None:
# 如果需要,添加交叉注意力(ControlNet 不需要,这里预留给扩展)
self.has_attn = False # ControlNet 主体不用 attention
self.norm2 = nn.GroupNorm(32, out_channels)
self.conv2 = nn.Conv2d(out_channels, out_channels, 3, padding=1)
# 残差连接
if in_channels != out_channels:
self.shortcut = nn.Conv2d(in_channels, out_channels, 1)
else:
self.shortcut = nn.Identity()
# ===== 关键: 跳跃连接的零卷积 =====
# 初始为零,训练时渐进增长
self.zero_conv = ZeroConv2d(in_channels, out_channels)
class ZeroConv2d(nn.Module):
"""
零卷积层。
权重初始化为零 → 前向传播输出为零 → 不破坏预训练网络。
训练时通过梯度更新逐渐增长。
"""
def __init__(self, in_channels, out_channels):
super().__init__()
self.conv = nn.Conv2d(in_channels, out_channels, 1, padding=0)
# 关键: 权重初始化为零
nn.init.constant_(self.conv.weight, 0.0)
nn.init.constant_(self.conv.bias, 0.0)
def forward(self, x):
return self.conv(x)

2.3 零卷积初始化的数学保证#

def zero_conv_guarantee():
"""
零卷积的数学保证。
"""
analysis = {
"初始化时": {
"权重": "W = zero",
"偏置": "b = zero",
"输出": "y = W·x + b = zero",
"效果": "ControlNet 对主网络的影响 = 0",
"主网络行为": "完全等价于原始 SD,生成质量不变",
},
"训练过程中": {
"梯度": "∂L/∂W ≠ 0 → W 开始更新",
"更新方向": "由条件损失驱动,向条件方向调整",
"增长速度": "初期慢(因为 W 接近零,梯度被放大)→ 后期稳定",
"避免破坏": "W 从零增长,不会产生大的破坏性跳变",
},
"收敛时": {
"权重": "W → W* (最优)",
"效果": "ControlNet 学会了精确的条件控制",
"主网络": "仍然冻结,预训练知识完整保留",
},
}
return analysis

2.4 完整的 ControlNet 架构#

class ControlNet(nn.Module):
"""
完整的 ControlNet。
复制 SD 的 UNet 编码器为可训练的副本。
"""
def __init__(self, sd_unet):
super().__init__()
self.sd_unet = sd_unet
# 冻结 SD UNet
for param in sd_unet.parameters():
param.requires_grad_(False)
# 复制编码器结构为 ControlNet
self.control_scales = [1.0] * 12 # 12 层控制强度
# 每一层对应一个可训练的 ControlNet Block
self.control_blocks = nn.ModuleList([
self._make_control_block(ch) for ch in sd_unet.channel_list
])
# 零卷积跳跃连接
self.zero_convs = nn.ModuleList([
ZeroConv2d(sd_unet.channel_list[i], sd_unet.channel_list[i])
for i in range(len(sd_unet.channel_list))
])
# 额外的条件图像处理
self.condition_encoder = nn.Sequential(
nn.Conv2d(3, 16, 3, padding=1), # 假设条件图为 RGB
nn.SiLU(),
nn.Conv2d(16, 32, 3, stride=2, padding=1),
nn.SiLU(),
nn.Conv2d(32, 64, 3, stride=2, padding=1),
nn.SiLU(),
nn.Conv2d(64, 64, 3, padding=1),
)
def _make_control_block(self, channels):
"""创建一个 ControlNet Block。"""
return nn.Sequential(
ResBlock(channels, channels),
ResBlock(channels, channels),
)
def forward(self, x_t, t, cond_image, context):
"""
前向传播。
参数:
x_t: (B, 4, H, W) 噪声潜变量
t: (B,) 时间步
cond_image: (B, 3, H, W) 条件图像(边缘/深度等)
context: (B, seq_len, d) 文本 embedding
"""
# 1) 编码条件图像
cond = self.condition_encoder(cond_image) # (B, 64, H/4, W/4)
# 2) 获取 SD UNet 编码器特征
# (冻结的编码器用于提供特征,ControlNet 副本用于学习条件)
features = self.sd_unet.get_encoder_features(x_t, t)
# 3) 条件注入
controlled_features = []
for i, (feat, ctrl_block, zero_conv) in enumerate(zip(
features, self.control_blocks, self.zero_convs
)):
# ControlNet 副本处理特征
ctrl_feat = ctrl_block(feat + cond) # 条件信息注入
# 零卷积跳跃连接
# 初始时 zero_conv 输出为 0 → 不影响主网络
# 训练后: 输出控制信号
skip = zero_conv(ctrl_feat)
controlled_features.append(skip)
# 4) SD UNet 解码器
output = self.sd_unet.decode_with_skip(
x_t, t, context, skip_connections=controlled_features
)
return output
def unfix_layers(self):
"""
解冻部分 SD UNet 层(可选的高级用法)。
渐进解冻: 先训练零卷积,再解冻更多层。
"""
pass # ControlNet 主体设计为冻结 SD

3. 训练策略:渐进式解冻#

3.1 零卷积的渐进式解冻#

ControlNet 的训练分为两个阶段:固定阶段解冻阶段

┌──────────────────────────────────────────────────────────────┐
│ 阶段 1: 固定阶段 (Locked Stage) │
│ │
│ SD UNet: 冻结 ❄️ (权重不变) │
│ ZeroConv: 冻结 🔒 (梯度计算,但权重不更新) │
│ │
│ 效果: │
│ - 零卷积从零开始,输出 = 0 │
│ - 只有 ControlNet Block 更新 │
│ - 解码器收到的跳跃连接 = 0 │
│ - 主网络完全保持预训练状态 │
│ │
├──────────────────────────────────────────────────────────────┤
│ 阶段 2: 解冻阶段 (Unlocked Stage) │
│ │
│ ZeroConv: 解冻 🔥 (零卷积开始学习) │
│ │
│ 效果: │
│ - 零卷积从零开始增长 │
│ - ControlNet 的控制信号逐渐注入 │
│ - 避免突然的大跳变 │
│ │
│ 为什么需要解冻 ZeroConv? │
│ - ZeroConv 需要学习最优的控制强度缩放 │
│ - 某些层可能需要更强的控制,某些层需要更弱 │
│ - 解冻后 ZeroConv 可以独立调整每层的贡献 │
└──────────────────────────────────────────────────────────────┘

3.2 ControlNet 训练循环#

def train_controlnet(controlnet, sd_unet, vae, text_encoder, dataloader,
optimizer, device, cfg_drop_prob=0.1):
"""
ControlNet 训练循环。
"""
controlnet.train()
sd_unet.eval() # 冻结
for batch in tqdm(dataloader):
images, cond_images, texts = (
batch["image"].to(device),
batch["condition"].to(device), # Canny边缘/深度图/姿态
batch["text"],
)
# ============ 1) 编码图像到潜空间 ============
with torch.no_grad():
z = vae.encode(images) # (B, 4, 64, 64)
# ============ 2) 采样时间步 ============
batch_size = z.shape[0]
t = torch.randint(0, 1000, (batch_size,), device=device)
# ============ 3) 加噪 ============
noise = torch.randn_like(z)
alpha_bar = get_diffusion_schedule(t) # 预设噪声调度
z_t = alpha_bar.sqrt().view(-1, 1, 1, 1) * z + \
(1 - alpha_bar).sqrt().view(-1, 1, 1, 1) * noise
# ============ 4) 文本编码 + CFG Drop ============
context = text_encoder.encode(texts)
uncond_context = text_encoder.encode([""] * batch_size)
# 随机丢弃条件(Classifier-Free Guidance)
drop_mask = torch.rand(batch_size, device=device) < cfg_drop_prob
context = context.masked_fill(drop_mask.unsqueeze(1), 0.0)
# ============ 5) SD UNet 前向(冻结,提供特征) ============
with torch.no_grad():
# SD 的无条件预测
noise_pred_uncond = sd_unet(z_t, t, uncond_context)
# ============ 6) ControlNet 前向(训练) ============
noise_pred_cond = controlnet(z_t, t, cond_images, context)
# ============ 7) CFG 合并 ============
# 训练时也用 CFG: 预测的噪声 = uncond + w * (cond - uncond)
guidance_scale = 7.5
noise_pred = noise_pred_uncond + guidance_scale * (noise_pred_cond - noise_pred_uncond)
# ============ 8) 计算损失 ============
loss = F.mse_loss(noise_pred, noise)
# ============ 9) 反向传播 ============
optimizer.zero_grad()
loss.backward()
torch.nn.utils.clip_grad_norm_(controlnet.parameters(), max_norm=1.0)
optimizer.step()
return loss.item()

3.3 条件预处理#

不同的控制类型需要不同的预处理方法:

class ConditionPreprocessor:
"""
条件图像预处理。
把各种类型的条件图像标准化到模型需要的格式。
"""
@staticmethod
def canny_edge(image, low_threshold=100, high_threshold=200):
"""
Canny 边缘检测 → 二值边缘图。
"""
import cv2
gray = cv2.cvtColor(image, cv2.COLOR_RGB2GRAY)
edges = cv2.Canny(gray, low_threshold, high_threshold)
# 转为 RGB (3通道,与 SD 输入对齐)
edges_rgb = np.stack([edges, edges, edges], axis=-1)
return edges_rgb / 255.0
@staticmethod
def depth_map(image, model="MiDaS"):
"""
单目深度估计 → 深度图。
"""
# 使用 MiDaS 或 ZoeDepth 估计深度
# 返回归一化的深度图
pass
@staticmethod
def human_pose(image, model="OpenPose"):
"""
人体姿态估计 → 关键点骨架图。
"""
# 使用 OpenPose 或 DWPose 检测 17/25 个关键点
# 绘制骨架连接
pass
@staticmethod
def normal_map(image, model="MiDaS"):
"""
表面法线估计 → 法线图。
"""
# 估计表面法线,用于几何控制
pass
@staticmethod
def semantic_seg(image, model="U2Net"):
"""
语义分割 → 分割掩码图。
"""
# 像素级类别分割
pass
def normalize_condition(cond_image, target_size=(512, 512)):
"""
标准化条件图像到 [0, 1] 范围,并 resize 到目标尺寸。
"""
import torchvision.transforms as T
transform = T.Compose([
T.ToTensor(),
T.Resize(target_size),
T.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5]),
])
return transform(cond_image)

4. 八种控制类型#

4.1 控制类型总览#

ControlNet 官方支持八种条件控制:

def control_types():
"""
ControlNet 支持的八种控制类型。
"""
return {
"Canny Edge": {
"输入": "任意图像",
"预处理": "Canny 边缘检测",
"效果": "精确的边缘轮廓控制",
"适用": "草图到图像、线稿上色",
"预训练模型大小": "~1.4GB",
},
"Depth": {
"输入": "任意图像",
"预处理": "单目深度估计 (MiDaS)",
"效果": "空间深度结构控制",
"适用": "深度感知生成、3D 场景构建",
"预训练模型大小": "~1.4GB",
},
"Normal Map": {
"输入": "任意图像",
"预处理": "表面法线估计",
"效果": "物体表面几何控制",
"适用": "纹理映射、几何细节",
"预训练模型大小": "~1.4GB",
},
"OpenPose": {
"输入": "人物图像",
"预处理": "人体姿态估计 (17 关键点)",
"效果": "人物姿态控制",
"适用": "角色动作生成、舞蹈/运动",
"预训练模型大小": "~1.4GB",
},
"Semantic Segmentation": {
"输入": "任意图像",
"预处理": "语义分割 (ADE20K)",
"效果": "物体类别和位置控制",
"适用": "场景布局、物体放置",
"预训练模型大小": "~1.4GB",
},
"M-LSD": {
"输入": "任意图像",
"预处理": "直线检测",
"效果": "建筑/室内线条结构控制",
"适用": "建筑草图、室内设计",
"预训练模型大小": "~1.4GB",
},
"HED Soft Edge": {
"输入": "任意图像",
"预处理": "HED 边缘检测 (软边缘)",
"效果": "柔和的边缘轮廓,比 Canny 更自然",
"适用": "艺术风格边缘、风格迁移",
"预训练模型大小": "~1.4GB",
},
"Scribble": {
"输入": "手绘涂鸦",
"预处理": "涂鸦转二值图",
"效果": "手绘草图控制",
"适用": "从涂鸦生成完整图像",
"预训练模型大小": "~1.4GB",
},
}

4.2 Canny Edge 控制详解#

Canny 是最常用的控制类型之一,适用于从线稿生成图像:

def canny_control_workflow():
"""
Canny Edge 控制的工作流程。
"""
return {
"输入图像": "任意照片或生成图",
"↓": "",
"Canny 边缘检测": "低阈值 100, 高阈值 200",
"↓": "得到精确的边缘轮廓图 (512×512×3)",
"ControlNet (Canny)": "学习"边缘 → 图像内容"的映射",
"↓": "输出条件噪声预测",
"SD UNet + CFG": "结合文本条件生成图像",
"↓": "",
"输出图像": "保留原始边缘结构的生成图",
}
def canny_training_data():
"""
Canny ControlNet 的训练数据构建。
"""
return {
"数据来源": "LAION-5B / COCO / ADE20K",
"样本数": "约 300 万对",
"条件图像": "原始图像经过 Canny 边缘检测",
"目标图像": "原始图像本身",
"文本": "图像的描述 (BLIP-2 生成)",
"训练技巧": "使用 ImageNet 预训练 Canny 检测器",
}

4.3 OpenPose 姿态控制详解#

OpenPose 是 ControlNet 在人物生成中最受欢迎的控制类型:

OpenPose 骨架:
┌─────────────────────────────────────────────────┐
│ │
│ [0] 鼻子 │
│ │ │
│ [1]─[2]─[3]─[4]─[15]─[14]─[16]─[17] │
│ │ │
│ [5] [6] [7] [8] │
│ │ │
│ [9] [10] [11] │
│ │ │
│ [12] [13] │
│ │
│ 17 个关键点: │
│ 0: 鼻子 1-2: 眼睛 3-4: 耳朵 │
│ 5-8: 肩膀、手肘、手腕 (右) │
│ 9-12: 髋、膝盖、脚踝 (右) │
│ 13-16: 肩膀、手肘、手腕 (左) │
│ 17-20: 髋、膝盖、脚踝 (左) │
│ │
└─────────────────────────────────────────────────┘
关键点连线 (骨骼):
头部: 0-1, 0-2, 1-3, 2-4 (脸)
躯干: 5-6, 6-7, 7-8 (右臂)
躯干: 5-11, 11-12, 12-13 (右腿)
躯干: 5-9, 9-10, 10-11 (左臂)
躯干: 5-14, 14-15, 15-16 (左腿)
class OpenPoseVisualizer:
"""
OpenPose 关键点可视化。
"""
def __init__(self, keypoint_confidence=0.3):
self.confidence = keypoint_confidence
# 骨骼连接
self.pairs = [
(0, 1), (0, 2), (1, 3), (2, 4), # 脸
(5, 6), (6, 7), (7, 8), # 右臂
(5, 9), (9, 10), (10, 11), # 左臂
(5, 11), (11, 12), (12, 13), # 右腿
(5, 14), (14, 15), (15, 16), # 左腿
]
# 颜色
self.keypoint_color = (255, 0, 85) # 红色
self.limb_color = (255, 85, 0) # 橙色
def draw_pose(self, keypoints, image_size):
"""
在空白图上绘制姿态骨架。
参数:
keypoints: (17, 3) [x, y, confidence] 关键点
image_size: (H, W) 输出图像大小
"""
canvas = np.zeros((image_size[0], image_size[1], 3), dtype=np.uint8)
# 绘制关键点
for kp in keypoints:
x, y, conf = kp
if conf > self.confidence:
cv2.circle(canvas, (int(x), int(y)), 4, self.keypoint_color, -1)
# 绘制骨骼连线
for i, j in self.pairs:
if keypoints[i][2] > self.confidence and keypoints[j][2] > self.confidence:
x1, y1 = int(keypoints[i][0]), int(keypoints[i][1])
x2, y2 = int(keypoints[j][0]), int(keypoints[j][1])
cv2.line(canvas, (x1, y1), (x2, y2), self.limb_color, 2)
return canvas

4.4 多控制组合#

def multi_control_inference(controlnet_models, sd_unet, vae, text_encoder,
z_t, t, conditions, texts, guidance_scale=7.5):
"""
多控制组合推理。
多个 ControlNet 可以组合使用。
"""
noise_preds = []
for ctrl_model, cond_img in zip(controlnet_models, conditions):
# 每个 ControlNet 预测
ctrl_pred = ctrl_model(z_t, t, cond_img, text_encoder.encode(texts))
noise_preds.append(ctrl_pred)
# 平均所有 ControlNet 的预测
noise_pred_cond = torch.stack(noise_preds).mean(dim=0)
# 无条件预测
noise_pred_uncond = sd_unet(z_t, t, text_encoder.encode([""]))
# CFG 合并
noise_pred = noise_pred_uncond + guidance_scale * (noise_pred_cond - noise_pred_uncond)
return noise_pred

5. 推理:精确控制生成#

5.1 ControlNet 推理流程#

@torch.no_grad()
def controlnet_inference(controlnet, sd_unet, vae, text_encoder,
prompt, cond_image, num_steps=50, guidance_scale=7.5):
"""
ControlNet 推理。
参数:
prompt: str 文本描述
cond_image: PIL.Image 条件图像(边缘/深度/姿态等)
num_steps: 推理步数 (DDIM)
guidance_scale: CFG 强度
"""
device = next(controlnet.parameters()).device
# ============ 1) 预处理条件图像 ============
cond_tensor = preprocess_condition(cond_image) # (1, 3, 512, 512)
cond_tensor = cond_tensor.to(device)
# ============ 2) 编码文本 ============
context = text_encoder.encode([prompt]) # 条件文本
uncond_context = text_encoder.encode([""]) # 无条件文本
# ============ 3) 初始化噪声 ============
z_T = torch.randn(1, 4, 64, 64, device=device)
# ============ 4) DDIM 采样 ============
timesteps = get_ddim_timesteps(num_steps) # 递减时间步
for i, t in enumerate(tqdm(timesteps)):
t_tensor = torch.tensor([t], device=device)
# 条件预测 (ControlNet)
ctrl_pred = controlnet(z_T, t_tensor, cond_tensor, context)
# 无条件预测 (SD)
uncond_pred = sd_unet(z_T, t_tensor, uncond_context)
# CFG 合并
noise_pred = uncond_pred + guidance_scale * (ctrl_pred - uncond_pred)
# DDIM 去噪步
z_T = ddim_step(z_T, noise_pred, t, timesteps[i - 1] if i > 0 else 0)
# ============ 5) VAE 解码 ============
image = vae.decode(z_T)
return image

5.2 ControlNet 的控制强度#

def control_strength():
"""
控制强度参数 (control_scale) 的效果。
"""
results = {
"0.0 (无控制)": {
"效果": "完全由文本决定,忽略条件图像",
"生成": "与普通 SD 无异",
},
"0.3-0.5 (弱控制)": {
"效果": "条件提供参考,但允许灵活变化",
"生成": "保留条件的部分结构,加入创意",
},
"0.7-0.9 (标准控制)": {
"效果": "条件主导结构,细节和风格由文本决定",
"生成": "精确的结构 + 丰富的细节",
},
"1.0 (强控制)": {
"效果": "条件主导一切,文本仅影响颜色/风格",
"生成": "结构与条件高度一致",
},
"1.2+ (过强)": {
"效果": "可能产生 artifacts 或过度平滑",
"生成": "结构过拟合,缺乏多样性",
},
}
return results
def set_control_scale(controlnet, scale=0.8):
"""
设置 ControlNet 的控制强度。
推理时可以动态调整。
"""
for module in controlnet.modules():
if isinstance(module, ZeroConv2d):
# ZeroConv 权重乘以 scale
with torch.no_grad():
module.conv.weight.mul_(scale)
return controlnet

6. ControlNet-XS:效率优化#

6.1 ControlNet 的效率问题#

原始 ControlNet 的问题是:需要完整复制 UNet 编码器,导致参数量大(与 SD UNet 相当):

原始 ControlNet:
- SD UNet: ~860M 参数
- ControlNet 副本: ~860M 参数 (相同大小!)
- 总计: ~1.7B 参数
- 问题: 部署成本高,推理慢

6.2 ControlNet-XS 的解法#

ControlNet-XS(ControlNet-XS: Lightweight ControlNet for Faster and Better Image Editing, Patil et al., 2024)提出了两个关键改进:

改进 1: 只复制编码器的后半部分
- 原始 ControlNet: 复制全部 12 层编码器
- ControlNet-XS: 只复制后 4 层 (下采样最深的部分)
- 参数量: ~860M → ~300M
原理:
- 高层语义特征 (深层) 才需要条件信息
- 低层纹理特征 (浅层) 对条件控制贡献小
改进 2: 更轻量的条件注入
- 不再完整复制编码器结构
- 使用浅层的旁路连接
- 通过 tiny 的调节器(modulator)注入条件
class ControlNetXS(nn.Module):
"""
ControlNet-XS: 轻量级 ControlNet。
只复制编码器的后半部分。
"""
def __init__(self, sd_unet, num_layers_to_copy=4):
super().__init__()
self.sd_unet = sd_unet
self.num_layers = num_layers_to_copy # 只需复制 4 层而非 12 层
# 只复制编码器的深层部分
self.control_blocks = nn.ModuleList([
self._make_control_block(ch)
for ch in sd_unet.channel_list[-num_layers_to_copy:]
])
# 轻量级条件编码器
self.cond_encoder = nn.Sequential(
nn.Conv2d(3, 32, 3, padding=1),
nn.SiLU(),
nn.Conv2d(32, 64, 3, stride=2, padding=1),
nn.SiLU(),
nn.Conv2d(64, 64, 3, padding=1),
)
# 轻量级跳跃连接
self.zero_convs = nn.ModuleList([
ZeroConv2d(sd_unet.channel_list[-num_layers_to_copy + i],
sd_unet.channel_list[-num_layers_to_copy + i])
for i in range(num_layers_to_copy)
])
def forward(self, x_t, t, cond_image, context):
# 条件编码
cond = self.cond_encoder(cond_image)
# 只在深层注入控制
features = self.sd_unet.get_deep_features(x_t, t) # 只取后 N 层
controlled_skips = []
for i, (feat, ctrl_block, zero_conv) in enumerate(zip(
features, self.control_blocks, self.zero_convs
)):
ctrl_feat = ctrl_block(feat + cond)
controlled_skips.append(zero_conv(ctrl_feat))
# UNet 解码(与原始相同)
output = self.sd_unet.decode_with_skip(x_t, t, context, controlled_skips)
return output

6.3 ControlNet-XS vs 原始 ControlNet#

def controlnet_comparison():
"""
ControlNet vs ControlNet-XS 对比。
"""
return {
"参数量": {
"ControlNet": "~860M (完整复制)",
"ControlNet-XS": "~300M (部分复制)",
"节省": "65%",
},
"推理速度": {
"ControlNet": "1x (慢)",
"ControlNet-XS": "~2x (快)",
},
"生成质量": {
"ControlNet": "最佳",
"ControlNet-XS": "略有下降,但在可接受范围",
"差距": "FID 差值 < 2%",
},
"适用场景": {
"ControlNet": "追求最高质量",
"ControlNet-XS": "边缘部署、实时应用",
},
}

7. ControlLoRA:更轻量的方案#

7.1 ControlLoRA 的核心思路#

ControlLoRA(2024)用 LoRA 低秩适配 替代完整的 ControlNet 副本:

原始 ControlNet:
- 完整复制 UNet 编码器: ~860M 参数
- 需要完整微调
ControlLoRA:
- 只微调编码器的 Low-rank 分解矩阵: ~5-20M 参数
- 冻结原始权重,通过低秩矩阵注入控制
- 训练更高效,部署更轻量
class ControlLoRA(nn.Module):
"""
ControlLoRA: 用 LoRA 替代完整复制。
"""
def __init__(self, sd_unet, rank=16, alpha=16):
super().__init__()
self.sd_unet = sd_unet
self.rank = rank
self.alpha = alpha
# 为每层添加 LoRA
self.lora_layers = nn.ModuleDict()
for name, param in sd_unet.named_parameters():
if "conv" in name and "weight" in name:
# 对卷积层添加 LoRA
self.lora_layers[name] = LoRALinear(param, rank, alpha)
# 冻结原始权重
for param in sd_unet.parameters():
param.requires_grad_(False)
# 条件处理
self.cond_encoder = CondEncoder(rank)
def forward(self, x_t, t, cond_image, context):
# 处理条件
cond = self.cond_encoder(cond_image) # 低秩条件表示
# 用 LoRA 注入条件
for name, lora_layer in self.lora_layers.items():
original = self.sd_unet.state_dict()[name]
lora_weight = lora_layer(cond) # 基于条件的低秩调整
self.sd_unet.set_param(name, original + lora_weight)
# 前向
output = self.sd_unet(x_t, t, context)
# 恢复原始权重
for name in self.lora_layers:
original = self.sd_unet.state_dict()[name]
self.sd_unet.set_param(name, original)
return output

7.2 ControlNet vs ControlNet-XS vs ControlLoRA#

def control_variants():
"""
三种控制方案对比。
"""
return {
"ControlNet": {
"参数量": "~860M",
"训练方式": "全量微调",
"质量": "最高",
"速度": "慢",
"适用": "追求最佳效果",
},
"ControlNet-XS": {
"参数量": "~300M",
"训练方式": "部分层微调",
"质量": "略低 (~2% FID差距)",
"速度": "中等",
"适用": "效率敏感的部署",
},
"ControlLoRA": {
"参数量": "~10-20M",
"训练方式": "LoRA 微调",
"质量": "低于 XS",
"速度": "快",
"适用": "轻量部署、个性化控制",
},
}

8. T2I-Adapter:另一种条件控制方案#

8.1 T2I-Adapter 的设计#

T2I-Adapter(Mou et al., 2024)与 ControlNet 不同:不是复制 UNet,而是训练一个轻量适配器

ControlNet:
- 完整复制 UNet 编码器
- 每个控制类型需要一个完整的 ControlNet
- 参数量大
T2I-Adapter:
- 不复制 UNet
- 训练一个轻量适配器 (约 80M 参数)
- 通过特征注入影响 UNet
- 可共享基础 UNet,多个 Adapter 组合使用
class T2IAdapter(nn.Module):
"""
T2I-Adapter: 轻量级条件适配器。
"""
def __init__(self, channels=[320, 640, 1280, 1280], num_blocks=4):
super().__init__()
# 条件图像编码器
self.cond_encoder = nn.Sequential(
nn.Conv2d(3, 64, 3, padding=1),
*[ResBlock(64, 64) for _ in range(num_blocks)],
nn.Conv2d(64, channels[-1], 1),
)
# 适配块: 与 UNet 各层对齐
self.adapters = nn.ModuleList([
AdapterBlock(in_ch, out_ch)
for in_ch, out_ch in zip(channels[:-1], channels[1:])
])
def forward(self, cond_image, unet_features):
"""
注入条件信息到 UNet 特征。
参数:
cond_image: (B, 3, H, W) 条件图像
unet_features: List[Tensor] UNet 各层特征
"""
# 编码条件
cond = self.cond_encoder(cond_image)
# 逐层适配
adapted_features = []
for feat, adapter in zip(unet_features, self.adapters):
adapted = adapter(feat, cond)
adapted_features.append(adapted)
return adapted_features

8.2 ControlNet vs T2I-Adapter#

def controlnet_vs_adapter():
"""
ControlNet vs T2I-Adapter 对比。
"""
return {
"架构": {
"ControlNet": "完整复制 UNet 编码器",
"T2I-Adapter": "独立轻量适配器",
},
"参数量": {
"ControlNet": "~860M",
"T2I-Adapter": "~80M",
},
"可组合性": {
"ControlNet": "每个控制类型独立模型",
"T2I-Adapter": "多个 Adapter 可叠加",
},
"质量": {
"ControlNet": "更高(更多参数)",
"T2I-Adapter": "略低,但足够好",
},
"灵活性": {
"ControlNet": "固定与 SD 配合",
"T2I-Adapter": "更灵活的注入点",
},
}

9. 完整实现:简化版 ControlNet#

import torch
import torch.nn as nn
import torch.nn.functional as F
from dataclasses import dataclass
import math
# ============ 基础模块 ============
class ResBlock(nn.Module):
def __init__(self, channels):
super().__init__()
self.norm1 = nn.GroupNorm(32, channels)
self.conv1 = nn.Conv2d(channels, channels, 3, padding=1)
self.norm2 = nn.GroupNorm(32, channels)
self.conv2 = nn.Conv2d(channels, channels, 3, padding=1)
def forward(self, x):
h = F.silu(self.norm1(x))
h = self.conv1(h)
h = F.silu(self.norm2(h))
h = self.conv2(h)
return h + x
class ZeroConv2d(nn.Module):
"""零卷积: 初始为零,渐进增长。"""
def __init__(self, in_ch, out_ch):
super().__init__()
self.conv = nn.Conv2d(in_ch, out_ch, 1)
nn.init.constant_(self.conv.weight, 0.0)
nn.init.constant_(self.conv.bias, 0.0)
def forward(self, x):
return self.conv(x)
# ============ 条件编码器 ============
class ConditionEncoder(nn.Module):
"""条件图像编码器。"""
def __init__(self, in_ch=3, mid_ch=64, out_ch=320):
super().__init__()
self.encoder = nn.Sequential(
nn.Conv2d(in_ch, mid_ch, 3, padding=1),
nn.SiLU(),
nn.Conv2d(mid_ch, mid_ch, 3, stride=2, padding=1), # 256
nn.SiLU(),
nn.Conv2d(mid_ch, mid_ch * 2, 3, padding=1),
nn.SiLU(),
nn.Conv2d(mid_ch * 2, mid_ch * 2, 3, stride=2, padding=1), # 128
nn.SiLU(),
nn.Conv2d(mid_ch * 2, out_ch, 3, padding=1),
)
def forward(self, x):
return self.encoder(x)
# ============ ControlNet Block ============
class ControlNetBlock(nn.Module):
"""ControlNet 基本模块。"""
def __init__(self, in_ch, out_ch):
super().__init__()
self.res1 = ResBlock(in_ch)
self.res2 = ResBlock(out_ch)
self.zero_conv = ZeroConv2d(out_ch, out_ch)
if in_ch != out_ch:
self.shortcut = nn.Conv2d(in_ch, out_ch, 1)
else:
self.shortcut = nn.Identity()
def forward(self, x, cond):
h = self.res1(x)
h = h + cond # 条件注入
h = self.res2(h)
return self.zero_conv(h)
# ============ ControlNet ============
class ControlNet(nn.Module):
"""
简化版 ControlNet。
"""
def __init__(self, sd_unet):
super().__init__()
self.sd_unet = sd_unet
self.channels = sd_unet.channel_list # [320, 640, 1280, 1280]
# 冻结 SD UNet
for param in sd_unet.parameters():
param.requires_grad_(False)
# 条件编码器
self.cond_encoder = ConditionEncoder(out_ch=self.channels[0])
# ControlNet Blocks(与 UNet 编码器结构对应)
self.control_blocks = nn.ModuleList([
nn.ModuleList([
ControlNetBlock(self.channels[i], self.channels[i + 1])
for _ in range(2) # 每层 2 个 block
])
for i in range(len(self.channels) - 1)
])
# 下采样
self.downsample = nn.ModuleList([
nn.Conv2d(self.channels[i], self.channels[i], 3, stride=2, padding=1)
for i in range(len(self.channels) - 1)
])
def forward(self, x_t, t, cond_image, context):
"""
参数:
x_t: (B, 4, 64, 64) 噪声潜变量
t: (B,) 时间步
cond_image: (B, 3, 512, 512) 条件图像
context: (B, seq_len, d) 文本 embedding
"""
# 编码条件
cond = self.cond_encoder(cond_image) # (B, 320, 64, 64)
# 获取 SD UNet 的特征
unet_features = self.sd_unet.get_encoder_features(x_t, t)
# 通过 ControlNet 处理并返回跳跃连接
skips = []
h = x_t
for i, (feat, blocks, down) in enumerate(zip(
unet_features, self.control_blocks, self.downsample
)):
for block in blocks:
feat = block(feat, cond if i == 0 else torch.zeros_like(cond))
skips.append(feat)
if i < len(self.downsample) - 1:
h = down(h)
# SD UNet 解码
output = self.sd_unet.decode_with_skips(x_t, t, context, skips)
return output
# ============ SD UNet (简化版) ============
class SDUNet(nn.Module):
"""简化版 Stable Diffusion UNet。"""
def __init__(self):
super().__init__()
self.channel_list = [320, 640, 1280, 1280]
# 简化实现
self.encoder_blocks = nn.ModuleList([
nn.ModuleList([ResBlock(ch) for _ in range(2)])
for ch in self.channel_list
])
self.downsample = nn.ModuleList([
nn.Conv2d(ch, ch, 3, stride=2, padding=1)
for ch in self.channel_list[:-1]
])
self.bottleneck = nn.ModuleList([ResBlock(1280) for _ in range(2)])
self.decoder_blocks = nn.ModuleList([
nn.ModuleList([ResBlock(ch) for _ in range(2)])
for ch in reversed(self.channel_list)
])
self.upsample = nn.ModuleList([
nn.ConvTranspose2d(ch, ch, 4, stride=2, padding=1)
for ch in reversed(self.channel_list[:-1])
])
self.out_conv = nn.Conv2d(320, 4, 3, padding=1)
self.time_embed = TimeEmbedding(320)
def get_encoder_features(self, x, t):
"""返回编码器各层特征(用于 ControlNet)。"""
features = []
t_emb = self.time_embed(t)
h = x
for i, (blocks, down) in enumerate(zip(self.encoder_blocks, self.downsample)):
for block in blocks:
h = block(h)
features.append(h)
h = down(h)
return features
def decode_with_skips(self, x, t, context, skips):
"""用 ControlNet 的跳跃连接解码。"""
t_emb = self.time_embed(t)
h = x
for i, (blocks, up) in enumerate(zip(
self.decoder_blocks, self.upsample
)):
# 拼接跳跃连接(ControlNet 提供)
if i < len(skips):
h = torch.cat([h, skips[-(i + 1)]], dim=1)
for block in blocks:
h = block(h)
h = up(h)
return self.out_conv(h)
def forward(self, x, t, context):
features = self.get_encoder_features(x, t)
return self.decode_with_skips(x, t, context, [torch.zeros_like(x)] * 4)
class TimeEmbedding(nn.Module):
def __init__(self, dim):
super().__init__()
self.mlp = nn.Sequential(
nn.Linear(128, dim * 4),
nn.SiLU(),
nn.Linear(dim * 4, dim),
)
def forward(self, t):
half = 64
emb = math.log(10000) / half
emb = torch.exp(torch.arange(half, device=t.device) * -emb)
emb = t[:, None] * emb[None, :]
emb = torch.cat([emb.sin(), emb.cos()], dim=-1)
return self.mlp(emb)

10. 应用场景#

10.1 ControlNet 的典型应用#

def controlnet_applications():
"""
ControlNet 的典型应用场景。
"""
return {
"线稿上色": {
"输入": "黑白线稿",
"条件": "Canny / HED 边缘检测",
"输出": "保留线稿结构的彩色图像",
"模型": "ControlNet Canny / HED",
},
"姿态控制": {
"输入": "姿态骨架图",
"条件": "OpenPose 关键点",
"输出": "符合姿态的人物图像",
"模型": "ControlNet OpenPose",
},
"深度感知生成": {
"输入": "深度图",
"条件": "MiDaS 深度估计",
"输出": "保持深度结构的三维图像",
"模型": "ControlNet Depth",
},
"建筑/室内设计": {
"输入": "建筑线稿",
"条件": "M-LSD 直线检测",
"输出": "保留建筑结构的渲染图",
"模型": "ControlNet M-LSD",
},
"分割引导生成": {
"输入": "语义分割图",
"条件": "ADE20K 分割",
"输出": "符合分割布局的图像",
"模型": "ControlNet Seg",
},
"风格迁移": {
"输入": "参考图像",
"条件": "法线图 / 深度图",
"输出": "保持结构、改变风格",
"模型": "ControlNet Normal / Depth",
},
}

10.2 组合多个 ControlNet#

def combine_multiple_controlnets():
"""
组合多个 ControlNet 实现更精细的控制。
"""
return {
"场景": "生成一个坐在室内的女性",
"控制1": "OpenPose (姿态) + ControlNet OpenPose",
"控制2": "深度图 (空间) + ControlNet Depth",
"控制3": "语义分割 (布局) + ControlNet Seg",
"文本": "a beautiful woman sitting in a modern living room",
"方法": "多个 ControlNet 的噪声预测求平均",
"效果": "精确的姿态 + 空间结构 + 场景布局",
}

11. 总结#

11.1 核心要点#

维度关键要点
核心问题扩散模型缺乏空间/结构级可控性
解决方案复制 UNet 编码器为可训练副本
核心机制零卷积初始化保证训练稳定,不破坏预训练权重
训练策略先冻结零卷积训练 ControlNet Block,再解冻零卷积
八种控制Canny/深度/姿态/分割/M-LSD/HED/Scribble/Normal
效率优化ControlNet-XS(部分复制)/ ControlLoRA(LoRA 替代)
开源影响成为 SD 最重要的控制扩展,催生 ComfyUI 等生态

11.2 ControlNet vs 其他条件控制方案#

ControlNet T2I-Adapter 分类器引导 额外通道
架构 复制编码器 轻量适配器 训练分类器 修改输入
参数量 ~860M ~80M ~100M ~86M
训练难度 低(零卷积) 中等 高 高
控制精度 高 中 中 中
兼容性 最佳 好 差 差
开源生态 最活跃 一般 无 一般

11.3 一句话总结#

ControlNet 的核心创新是”复制 + 冻结 + 零卷积”——把预训练的 SD UNet 编码器完整复制一份作为可训练副本,冻结原始权重保证预训练知识不丢失,通过零卷积初始化的跳跃连接把控制信号从零开始渐进注入解码器。这使得任意条件(边缘、姿态、深度、分割等)都能以插件形式精确控制图像生成结构,同时完全保留 Stable Diffusion 的文本生成质量——成为 AI 图像生成从”创意工具”升级为”可控设计工具”的关键技术。

11.4 推荐资源#

论文:
- ControlNet (Zhang & Agrawala, 2023): "Adding Conditional Control to Text-to-Image Diffusion Models"
- ControlNet-XS (Patil et al., 2024): "ControlNet-XS: Lightweight ControlNet for Faster and Better Image Editing"
- ControlLoRA (2024): LoRA-based ControlNet
- T2I-Adapter (Mou et al., 2024): "T2I-Adapter: Adapting Diffusion Models for Text-to-Image Generation"
代码:
- lllyasviel/ControlNet (官方实现)
- lllyasviel/sd-webui-controlnet (AUTOMATIC1111 WebUI 插件)
- comfyanonymous/ComfyUI (节点式 ControlNet 工作流)
预训练模型:
- HuggingFace: lllyasviel/sd-controlnet-canny
- HuggingFace: lllyasviel/sd-controlnet-depth
- HuggingFace: lllyasviel/sd-controlnet-openpose
- ModelScope: damo/control_v11p_sd15_canny

文章分享

如果这篇文章对你有帮助,欢迎分享给更多人!

ControlNet 深度剖析:条件控制扩散模型的结构化控制系统
https://aiattnstudio.link/posts/controlnet/
作者
Federico
发布于
2026-07-16
许可协议
CC BY-NC-SA 4.0
Profile Image of the Author

Federico

AI Research Lab

Hello, I'm Federico.

关于实验室 / About
公告

欢迎来到Federico的个人博客

分类
标签
站点统计
57文章
7分类
404标签