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/ControlNet1.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 analysis2.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 主体设计为冻结 SD3. 训练策略:渐进式解冻
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 canvas4.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_pred5. 推理:精确控制生成
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 image5.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 controlnet6. 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 output6.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 output7.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_features8.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 torchimport torch.nn as nnimport torch.nn.functional as Ffrom dataclasses import dataclassimport 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文章分享
如果这篇文章对你有帮助,欢迎分享给更多人!

