[特殊字符] Diffusers 条件二维 UNet(UNet2DConditionModel)全面解析:架构、配置与实战

🤗 Diffusers 条件二维 UNet(UNet2DConditionModel)全面解析:架构、配置与实战

【免费下载链接】diffusers 🤗 Diffusers: State-of-the-art diffusion models for image, video, and audio generation in PyTorch. 【免费下载链接】diffusers 项目地址: https://gitcode.com/GitHub_Trending/di/diffusers

UNet2DConditionModel 是 🤗 Diffusers 中图像、视频生成类扩散模型(如 Stable Diffusion 系列、SDXL、Kandinsky、ControlNet、T2I-Adapter 等)的核心去噪主干网络。本文以 docs/source/en/api/models/unet2d-cond.md 为骨架,结合 unet_2d_condition.py 的完整源码与 test_models_unet_2d_condition.py 测试用例,系统讲解其设计原理、全部构造参数、前向推理流程、条件注入机制以及注意力切片、FreeU、QKV 融合等进阶用法,帮助读者从"会用"进阶到"懂原理、能调参、可二次开发"。

一、从 U-Net 到条件二维 UNet:为什么扩散系统离不开它

U-Net 架构最初由 Ronneberger 等人在 2015 年提出,用于生物医学图像分割。其核心设计是一个"收缩路径(contracting path)"用于捕获上下文信息,配以一个"对称的扩张路径(expanding path)"用于精确定位。该论文摘要原文如下:

There is large consent that successful training of deep networks requires many thousand annotated training samples. In this paper, we present a network and training strategy that relies on the strong use of data augmentation to use the available annotated samples more efficiently. The architecture consists of a contracting path to capture context and a symmetric expanding path that enables precise localization. We show that such a network can be trained end-to-end from very few images and outperforms the prior best method (a sliding-window convolutional network) on the ISBI challenge for segmentation of neuronal structures in electron microscopic stacks. Using the same network trained on transmitted light microscopy images (phase contrast and DIC) we won the ISBI cell tracking challenge 2015 in these categories by a large margin. Moreover, the network is fast. Segmentation of a 512x512 image takes less than a second on a recent GPU. The full implementation (based on Caffe) and the trained networks are available at http://lmb.informatik.uni-freiburg.de/people/ronneber/u-net.

🤗 Diffusers 之所以大量采用 UNet 结构,是因为它能够输出与输入相同尺寸的图片——这一"输入输出同分辨率"的特性天然契合扩散模型的去噪任务:在扩散过程中,网络需要在每一步把加噪的潜在表示(latent)逐步还原为清晰表示,而不改变空间尺寸。

在 🤗 Diffusers 中存在多个 UNet 变体,按维度是否条件化划分:

  • 2D UNet 无条件模型(UNet2DModel
  • 2D UNet 条件模型(UNet2DConditionModel)——本文主角,支持文本、类别、图像等多种条件注入
  • 3D UNet 条件模型(UNet3DConditionModel,用于视频扩散)

从源码的类定义可见,UNet2DConditionModel 同时继承了多个 Mixin,使其具备丰富的生态能力(unet_2d_condition.py#L76-L78):

class UNet2DConditionModel(
    ModelMixin, AttentionMixin, ConfigMixin, FromOriginalModelMixin, UNet2DConditionLoadersMixin, PeftAdapterMixin
):

其中 ModelMixin 提供通用的模型下载、保存、from_pretrained 等能力;ConfigMixin 提供 @register_to_config 配置注册机制;UNet2DConditionLoadersMixin 提供单文件(single-file)权重加载;PeftAdapterMixin 提供 LoRA 适配器支持;AttentionMixin 提供注意力处理器(attention processor)管理能力。

二、类与输出定义

2.1 UNet2DConditionOutput

模型的输出由 UNet2DConditionOutput 定义,它是一个基于 BaseOutput 的 dataclass,只有一个字段:

字段类型说明
sampletorch.Tensor,形状 (batch_size, num_channels, height, width)encoder_hidden_states 条件约束的隐藏状态输出,即模型最后一层的输出

forwardreturn_dict=True(默认)时返回该对象;置为 False 时返回普通元组 (sample,)unet_2d_condition.py#L1232-L1235)。

2.2 UNet2DConditionModel 的定位

按官方文档描述,该模型是一个条件 2D UNet:接收"带噪样本 + 条件状态 + 时间步(timestep)",返回与输入形状一致的样本。文档原文的 autodoc 指令([[autodoc]] UNet2DConditionModel)会直接从源码 docstring 生成 API 文档,因此源码类 docstring 中列出的全部构造参数即是该模型最权威的配置清单,详见下文第三节。

三、构造参数全解:读懂模型的"骨架配置"

UNet2DConditionModel.__init__ 使用 @register_to_config 装饰(unet_2d_condition.py#L177-L238),所有参数都会持久化到模型的 config.json 中。下表整理了全部参数及其默认值:

参数默认值说明
sample_sizeNone输入/输出样本的高和宽,可为 int(h, w) 元组
in_channels4输入样本的通道数(Stable Diffusion 的 VAE 潜在空间为 4 通道)
out_channels4输出通道数
center_input_sampleFalse是否将输入样本中心化(sample = 2 * sample - 1
flip_sin_to_cosTrue时间嵌入中是否将 sin 翻转为 cos
freq_shift0时间嵌入的频率偏移
down_block_types("CrossAttnDownBlock2D", "CrossAttnDownBlock2D", "CrossAttnDownBlock2D", "DownBlock2D")下采样块类型元组
mid_block_type"UNetMidBlock2DCrossAttn"中间块类型,可选 UNetMidBlock2DCrossAttnUNetMidBlock2DUNetMidBlock2DSimpleCrossAttn,设为 None 则跳过中间块
up_block_types("UpBlock2D", "CrossAttnUpBlock2D", "CrossAttnUpBlock2D", "CrossAttnUpBlock2D")上采样块类型元组
only_cross_attentionFalse基础 Transformer 块中是否只使用交叉注意力(不含自注意力),可为 bool 或逐块 bool 元组
block_out_channels(320, 640, 1280, 1280)每个块的输出通道数
layers_per_block2每个块的层数
downsample_padding1下采样卷积的 padding
mid_block_scale_factor1.0中间块的比例因子
dropout0.0dropout 概率
act_fn"silu"激活函数(SiLU)
norm_num_groups32归一化分组数;设为 None 则跳过后处理的归一化与激活层
norm_eps1e-5归一化的 epsilon
cross_attention_dim1280交叉注意力特征维度(即文本嵌入的维度)
transformer_layers_per_block1每个交叉注意力块中 BasicTransformerBlock 的数量,仅对 CrossAttnDownBlock2DCrossAttnUpBlock2DUNetMidBlock2DCrossAttn 有效
reverse_transformer_layers_per_blockNone非对称 UNet 中上采样块使用的 Transformer 层数(当 transformer_layers_per_block 为嵌套元组时必须提供)
encoder_hid_dimNone若定义了 encoder_hid_dim_typeencoder_hidden_states 会从该维度投影到 cross_attention_dim
encoder_hid_dim_typeNone编码器隐藏状态的投影方式,如 text_projtext_image_projimage_proj
attention_head_dim8注意力头维度
num_attention_headsNone注意力头数量;未定义时默认取 attention_head_dim 的值
dual_cross_attentionFalse是否使用双重交叉注意力
use_linear_projectionFalse交叉注意力是否使用线性投影(SDXL 启用)
class_embed_typeNone类别嵌入类型:None"timestep""identity""projection""simple_projection"
addition_embed_typeNone额外嵌入类型:None"text"(使用 TextTimeEmbedding 层)等
addition_time_embed_dimNone额外时间步嵌入的维度
num_class_embedsNone类别条件嵌入矩阵的输入维度
upcast_attentionFalse是否对注意力上转型计算
resnet_time_scale_shift"default"ResNet 块的时间尺度偏移配置,可选 defaultscale_shift
resnet_skip_time_actFalseResNet 是否跳过时间激活
resnet_out_scale_factor1.0ResNet 输出缩放因子
time_embedding_type"positional"时间步位置嵌入类型:positionalfourier
time_embedding_dimNone投影后时间嵌入维度的覆盖值
time_embedding_act_fnNone时间嵌入的激活函数:silumishgeluswish
timestep_post_actNone时间步嵌入中的第二个激活函数:silumishgelu
time_cond_proj_dimNone时间步嵌入中 cond_proj 层的维度
conv_in_kernel3输入卷积核大小
conv_out_kernel3输出卷积核大小
projection_class_embeddings_input_dimNoneclass_embed_type="projection"class_labels 的输入维度(必须提供)
attention_type"default"注意力类型,如 gated(GLIGEN)
class_embeddings_concatFalse是否将时间嵌入与类别嵌入拼接(拼接时传入块的时间嵌入维度翻倍)
mid_block_only_cross_attentionNone中间块是否仅用交叉注意力;only_cross_attention 为单个 bool 且此值为 None 时继承该值
cross_attention_normNone交叉注意力的归一化方式
addition_embed_type_num_heads64addition_embed_type="text"TextTimeEmbedding 的头数

3.1 参数校验与广播机制

__init__ 中首先调用 _check_config 做一致性校验(unet_2d_condition.py#L498-L548),核心规则包括:

  • down_block_typesup_block_types 长度必须一致;
  • block_out_channelsonly_cross_attentionattention_head_dimcross_attention_dimlayers_per_block 等列表型参数的长度必须等于 down block 数量
  • transformer_layers_per_block 为嵌套列表且未提供 reverse_transformer_layers_per_block,会直接报错(非对称 UNet 场景)。

随后是标量广播逻辑:only_cross_attentionnum_attention_headsattention_head_dimcross_attention_dimlayers_per_blocktransformer_layers_per_block 若传入单个值,都会按 down block 数量展开为逐块元组(unet_2d_condition.py#L329-L351)。

3.2 关于 num_attention_heads 的历史遗留问题

源码中有段值得注意的"兼容代码"(unet_2d_condition.py#L243-L254):传入 num_attention_heads 会直接抛出 ValueError,提示该参数因历史命名问题暂不支持,只能通过 attention_head_dim 控制头维度;随后 num_attention_heads = num_attention_heads or attention_head_dim,即用 attention_head_dim 的值兜底。这是 diffusers 早期版本命名不一致留下的向后兼容设计,读者在自定义配置时应只使用 attention_head_dim

四、网络结构:编码器—中间块—解码器三段式

__init__ 的构建逻辑(unet_2d_condition.py#L270-L496)可以看出完整结构:

  1. 输入层conv_in,一个将 in_channels 映射到 block_out_channels[0]nn.Conv2d,padding 由 (conv_in_kernel - 1) // 2 自动计算;
  2. 时间/条件嵌入层time_proj + time_embeddingTimestepEmbedding),以及可选的类别嵌入、附加嵌入、编码器隐藏状态投影;
  3. 下采样路径(down_blocks):按 down_block_types 逐块构建,每个块由若干层组成,除最后一个块外都带 add_downsampleunet_2d_condition.py#L361-L394);
  4. 中间块(mid_block):通过 get_mid_blockmid_block_type 构建(unet_2d_condition.py#L396-L418);
  5. 上采样路径(up_blocks):按 up_block_types 构建,块类型与 down blocks 逆序对称,除最终块外都带 add_upsample,并记录 num_upsamplers(用于计算整体上采样因子,见 unet_2d_condition.py#L420-L477);
  6. 输出层conv_norm_out(GroupNorm)+ conv_act + conv_outunet_2d_condition.py#L479-L494)。

块类型集中在 unet_2d_blocks.py 中:

以 Stable Diffusion 默认配置为例:down_block_types 为三个 CrossAttnDownBlock2D 加一个 DownBlock2Dup_block_types 对称地为一个 UpBlock2D 加三个 CrossAttnUpBlock2Dblock_out_channels=(320, 640, 1280, 1280)。这也是文档中"几个不同的 UNet 变体取决于维度和是否条件化"的具体体现——条件化正是通过交叉注意力块注入文本/图像条件实现的

五、forward 前向流程:一次完整去噪的内部旅程

forward 方法(unet_2d_condition.py#L978-L1235)接受以下参数:

参数类型说明
sampletorch.Tensor带噪输入,形状 (batch, channel, height, width)
timesteptorch.Tensor/float/int去噪时间步
encoder_hidden_statestorch.Tensor编码器隐藏状态,形状 (batch, seq_len, feature_dim)(如文本嵌入)
class_labelstorch.Tensor类别标签(其嵌入会与时间步嵌入相加)
timestep_condtorch.Tensor时间步的条件嵌入(若有,与经 time_embedding 的样本相加)
attention_masktorch.Tensor形状 (batch, key_tokens) 的注意力掩码,1 保留、0 丢弃,会被转换为 bias
cross_attention_kwargsdict透传给 AttentionProcessor 的 kwargs(含 LoRA 缩放与 GLIGEN 参数)
added_cond_kwargsdict附加条件嵌入字典(SDXL 的 text_embeds/time_ids 等)
down_block_additional_residualstuple添加到 down 块残差的张量(ControlNet 专用
mid_block_additional_residualtorch.Tensor添加到中间块残差的张量(ControlNet)
down_intrablock_additional_residualstuple添加到 down 块内部的残差(T2I-Adapter 专用
encoder_attention_masktorch.Tensor交叉注意力掩码,形状 (batch, seq_len)
return_dictbool是否返回 UNet2DConditionOutput

整个流程可划分为六个阶段:

  1. 尺寸与掩码预处理:计算 default_overall_up_factor = 2 ** num_upsamplers,若输入空间尺寸不是该因子的整数倍,则开启 forward_upsample_size 以在解码阶段动态插值对齐;attention_maskencoder_attention_mask 均被转换为 (1 - mask) * -10000.0 的注意力 bias,并增加单例 query 维度(unet_2d_condition.py#L1039-L1074)。
  2. 条件嵌入合成:先由 get_time_embed 生成时间嵌入 t_emb(内部经 Timesteps 正弦编码并广播到 batch 维度),再经 time_embedding 得到 emb;随后把类别嵌入(get_class_embed,拼接或相加)、附加嵌入(get_aug_embed,覆盖 text / text_image / text_time / image / image_hint 五种模式)逐步合成进 embunet_2d_condition.py#L1076-L1101)。
  3. 输入预处理conv_in 卷积(unet_2d_condition.py#L1107-L1108);若 attention_typegated(GLIGEN),还会经 position_net 处理边界框条件。
  4. 下采样(down):逐个遍历 down_blocks,带交叉注意力的块会额外接收 encoder_hidden_statesattention_mask 等;收集 down_block_res_samples 作为跳连接(skip connection)。这里同时支持两条旁路:ControlNet 的 down_block_additional_residuals 与 T2I-Adapter 的 down_intrablock_additional_residuals,其中 T2I-Adapter 的旧式传参(走 down_block_additional_residuals)会触发 deprecation 警告(unet_2d_condition.py#L1116-L1168)。
  5. 中间块(mid):若配置了 mid block 则执行;ControlNet 场景下将 mid_block_additional_residual 直接加到输出上(unet_2d_condition.py#L1170-L1193)。
  6. 上采样(up)与后处理:解码器逐块消费跳连接 res_samples,非最终块且开启 forward_upsample_size 时按对应 down 块尺寸插值;最后经 conv_norm_outconv_actconv_out 输出与输入同尺寸的预测结果(unet_2d_condition.py#L1195-L1230)。

值得注意的是,forward 上方的 @apply_lora_scale("cross_attention_kwargs") 装饰器(unet_2d_condition.py#L978)会在每次前向时应用 LoRA 缩放,这就是 LoRA 微调权重能无缝作用于推理的原因。

六、条件注入的四种通道:源码级解读

forward 的实现可以归纳出模型支持的四类条件

  1. 时间步条件(必选):timestepget_time_embedTimesteps 正弦投影 → TimestepEmbedding,是扩散过程的"进度指示器";
  2. 文本/编码器条件(核心):encoder_hidden_statesprocess_encoder_hidden_states 处理(unet_2d_condition.py#L942-L976),按 encoder_hid_dim_type 支持 text_proj(线性投影)、text_image_proj(Kandinsky 2.1 风格)、image_proj(Kandinsky 2.2 风格)、ip_image_proj(Image Prompt 风格)四种投影,然后注入各交叉注意力块;
  3. 类别条件class_labelsget_class_embedunet_2d_condition.py#L874-L888),按 class_embed_type 支持 nn.Embeddingtimestepidentityprojectionsimple_projection 五种实现,最终与时间嵌入相加或拼接;
  4. 附加条件(added_cond_kwargs)get_aug_embedunet_2d_condition.py#L890-L940)覆盖五种模式——textTextTimeEmbedding)、text_image(Kandinsky 2.1)、text_time(SDXL 的 text_embeds + time_ids)、image(Kandinsky 2.2)、image_hint(Kandinsky 2.2 ControlNet)。例如 SDXL 要求 added_cond_kwargs 必须包含 text_embedstime_ids,否则会抛 ValueError

正是这套"时间 + 文本 + 类别 + 附加"的多通道条件注入机制,让同一个 UNet2DConditionModel 能服务于文本到图像(Stable Diffusion)、图像到图像(img2img)、修复(inpainting)、ControlNet 引导、T2I-Adapter 等多种扩散管线。

七、进阶 API:注意力切片、FreeU 与 QKV 融合

除前向推理外,模型还内置了一系列实用方法:

7.1 set_attention_slice:低显存注意力

set_attention_slice(slice_size="auto")unet_2d_condition.py#L725-L788)将注意力计算按切片分批执行,以少量速度损失换取显存节省:

  • "auto":注意力头输入减半,注意力分两步计算(默认推荐,速度/显存平衡好);
  • "max":每次只运行一个切片,最大化显存节省;
  • 传入数值:按 attention_head_dim // slice_size 切分,要求 attention_head_dim 能被 slice_size 整除。

实现上会递归遍历所有子模块,收集各层的 sliceable_head_dim,再反向递归应用切片,因此兼容任意深度的嵌套块结构。

7.2 enable_freeu / disable_freeu:免训练画质增强

enable_freeu(s1, s2, b1, b2)unet_2d_condition.py#L790-L812)启用 FreeU 机制:s1/s2 衰减两个阶段跳连接特征的贡献(缓解"过度平滑"),b1/b2 放大主干特征的贡献,四个系数直接挂载到每个上采样块上;disable_freeu()unet_2d_condition.py#L814-L820)则将其全部置为 None 关闭。不同管线(SD v1/v2/SDXL)有各自调好的系数组合。

7.3 fuse_qkv_projections / unfuse_qkv_projections:注意力融合加速

fuse_qkv_projections()unet_2d_condition.py#L822-L841)将自注意力的 Q/K/V 三个投影矩阵融合为一个,交叉注意力的 K/V 融合,并用 FusedAttnProcessor2_0 替代原处理器;unfuse_qkv_projections()unet_2d_condition.py#L843-L850)恢复原始处理器。该方法标记为 🧪 实验性 API,且不支持带 Added KV 投影的模型(会主动抛错)。

7.4 set_default_attn_processor 与注意力处理器体系

set_default_attn_processor()unet_2d_condition.py#L710-L723)根据当前处理器类型(ADDED_KV_ATTENTION_PROCESSORSCROSS_ATTENTION_PROCESSORS)自动回退到 AttnAddedKVProcessorAttnProcessor,用于清除自定义处理器。

7.5 其他模型级属性

源码还声明了若干辅助属性:_supports_gradient_checkpointing = True(支持梯度检查点训练)、_no_split_modules(FSDP/DeepSpeed 分片时不可切分的模块)、_skip_layerwise_casting_patterns = ["norm"](跳过 norm 层的逐层类型转换)、_repeated_blocks = ["BasicTransformerBlock"](用于设备卸载与参数共享识别)。

八、训练与推理实践:如何实例化与加载

8.1 从预训练仓库加载

使用统一的 from_pretrained 接口即可加载任意基于该架构的预训练模型:

from diffusers import UNet2DConditionModel

# 从 Hugging Face Hub 加载(以 Stable Diffusion v1.5 的 UNet 为例)
unet = UNet2DConditionModel.from_pretrained("runwayml/stable-diffusion-v1-5", subfolder="unet")

# 加载到指定设备与精度
unet.to("cuda", dtype=torch.float16)

由于 UNet2DConditionModel 继承自 ModelMixinConfigMixinfrom_pretrainedsave_pretrainedpush_to_hub 等通用能力开箱即用。

8.2 从零实例化自定义配置

from diffusers import UNet2DConditionModel

unet = UNet2DConditionModel(
    sample_size=64,                      # 训练分辨率
    in_channels=4,                       # 潜在空间通道数
    out_channels=4,
    down_block_types=(
        "CrossAttnDownBlock2D",
        "CrossAttnDownBlock2D",
        "CrossAttnDownBlock2D",
        "DownBlock2D",
    ),
    up_block_types=(
        "UpBlock2D",
        "CrossAttnUpBlock2D",
        "CrossAttnUpBlock2D",
        "CrossAttnUpBlock2D",
    ),
    block_out_channels=(320, 640, 1280, 1280),
    layers_per_block=2,
    cross_attention_dim=768,             # 文本编码器嵌入维度(CLIP 为 768)
    attention_head_dim=8,
    norm_num_groups=32,
)

8.3 前向调用示例

import torch

sample = torch.randn(2, 4, 64, 64)            # 带噪潜在表示
timestep = torch.tensor([500, 500])           # 去噪时间步
encoder_hidden_states = torch.randn(2, 77, 768)  # 文本嵌入 (batch, seq, dim)

output = unet(sample=sample, timestep=timestep, encoder_hidden_states=encoder_hidden_states)
# output 为 UNet2DConditionOutput,output.sample 形状为 (2, 4, 64, 64)

8.4 测试覆盖:功能验证的"说明书"

仓库测试 test_models_unet_2d_condition.py 提供了官方验证逻辑,是理解模型行为的极佳参考:

  • TestUNet2DCondition 继承 ModelTesterMixinUNetTesterMixin,覆盖输入输出形状、前向稳定性、训练模式等通用检查;
  • 专项测试包括:test_model_with_attention_head_dim_tuple(注意力头维度元组)、test_model_with_use_linear_projection(线性投影)、test_model_with_cross_attention_dim_tuple(交叉注意力维度元组)、test_model_with_simple_projection(简单投影)、test_model_with_class_embeddings_concat(类别嵌入拼接)、test_model_xattn_padding(交叉注意力 padding)、test_asymmetrical_unet(非对称 UNet,验证 reverse_transformer_layers_per_block);
  • TestUNet2DConditionHubLoading 覆盖分片 checkpoint 的 Hub 加载、subfolder 加载等场景。

此外,tests/lora/utils.pytests/models/test_modeling_common.py 也引用了该模型,分别验证 LoRA 适配器与通用建模能力。

九、生态中的位置:它支撑了哪些管线

在 🤗 Diffusers 中,UNet2DConditionModel 是大量 2D 图像扩散管线的去噪核心,例如:

  • Stable Diffusion 系列:文本条件扩散,使用 CrossAttnDownBlock2D/CrossAttnUpBlock2D 注入 CLIP 文本嵌入;
  • Stable Diffusion XL:通过 addition_embed_type="text_time" 注入 pooled text 嵌入与尺寸时间 ID,并启用 use_linear_projection
  • Kandinsky 2.x:通过 addition_embed_typetext_image/image/image_hintencoder_hid_dim_typeimage_proj 实现图像条件注入;
  • ControlNet / T2I-Adapter:通过 down_block_additional_residualsdown_intrablock_additional_residuals 注入外部引导信号。

十、小结

UNet2DConditionModel 以经典的对称编码器—解码器 U-Net 结构为骨架,通过交叉注意力块与多通道嵌入注入机制,将"带噪样本 + 时间步 + 任意条件"映射为同尺寸的去噪输出。理解其构造参数(块类型、通道数、注意力维度、嵌入模式)与前向流程(时间/类别/附加条件合成 → down → mid → up → 后处理),是定制扩散模型、接入 ControlNet/T2I-Adapter、以及训练 LoRA 的基础。若需进一步深入,可继续研读 unet_2d_blocks.py 中各类块的实现、attention_processor.py 的注意力处理器体系,以及 test_models_unet_2d_condition.py 中覆盖各类配置组合的测试用例。

【免费下载链接】diffusers 🤗 Diffusers: State-of-the-art diffusion models for image, video, and audio generation in PyTorch. 【免费下载链接】diffusers 项目地址: https://gitcode.com/GitHub_Trending/di/diffusers

创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

实付
使用余额支付
点击重新获取
扫码支付
钱包余额 0

抵扣说明:

1.余额是钱包充值的虚拟货币,按照1:1的比例进行支付金额的抵扣。
2.余额无法直接购买下载,可以购买VIP、付费专栏及课程。

余额充值