ARTICLE DETAIL

资讯详情

深耕郑州网站建设与运营推广的一线实战洞察。

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

[特殊字符] Diffusers 条件二维 UNet(UNet2DConditionModel)全面解析:架构、配置与实战 Diffusers 条件二维 UNetUNet2DConditionModel全面解析架构、配置与实战【免费下载链接】diffusers Diffusers: State-of-the-art diffusion models for image, video, and audio generation in PyTorch.项目地址: https://gitcode.com/GitHub_Trending/di/diffusersUNet2DConditionModel是 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 无条件模型UNet2DModel2D UNet 条件模型UNet2DConditionModel——本文主角支持文本、类别、图像等多种条件注入3D UNet 条件模型UNet3DConditionModel用于视频扩散从源码的类定义可见UNet2DConditionModel同时继承了多个 Mixin使其具备丰富的生态能力unet_2d_condition.py#L76-L78class 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条件约束的隐藏状态输出即模型最后一层的输出当forward的return_dictTrue默认时返回该对象置为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 - 1flip_sin_to_cosTrue时间嵌入中是否将 sin 翻转为 cosfreq_shift0时间嵌入的频率偏移down_block_types(CrossAttnDownBlock2D, CrossAttnDownBlock2D, CrossAttnDownBlock2D, DownBlock2D)下采样块类型元组mid_block_typeUNetMidBlock2DCrossAttn中间块类型可选UNetMidBlock2DCrossAttn、UNetMidBlock2D、UNetMidBlock2DSimpleCrossAttn设为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下采样卷积的 paddingmid_block_scale_factor1.0中间块的比例因子dropout0.0dropout 概率act_fnsilu激活函数SiLUnorm_num_groups32归一化分组数设为None则跳过后处理的归一化与激活层norm_eps1e-5归一化的 epsiloncross_attention_dim1280交叉注意力特征维度即文本嵌入的维度transformer_layers_per_block1每个交叉注意力块中BasicTransformerBlock的数量仅对CrossAttnDownBlock2D、CrossAttnUpBlock2D、UNetMidBlock2DCrossAttn有效reverse_transformer_layers_per_blockNone非对称 UNet 中上采样块使用的 Transformer 层数当transformer_layers_per_block为嵌套元组时必须提供encoder_hid_dimNone若定义了encoder_hid_dim_typeencoder_hidden_states会从该维度投影到cross_attention_dimencoder_hid_dim_typeNone编码器隐藏状态的投影方式如text_proj、text_image_proj、image_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_projectionaddition_embed_typeNone额外嵌入类型None或text使用TextTimeEmbedding层等addition_time_embed_dimNone额外时间步嵌入的维度num_class_embedsNone类别条件嵌入矩阵的输入维度upcast_attentionFalse是否对注意力上转型计算resnet_time_scale_shiftdefaultResNet 块的时间尺度偏移配置可选default或scale_shiftresnet_skip_time_actFalseResNet 是否跳过时间激活resnet_out_scale_factor1.0ResNet 输出缩放因子time_embedding_typepositional时间步位置嵌入类型positional或fouriertime_embedding_dimNone投影后时间嵌入维度的覆盖值time_embedding_act_fnNone时间嵌入的激活函数silu、mish、gelu、swishtimestep_post_actNone时间步嵌入中的第二个激活函数silu、mish、gelutime_cond_proj_dimNone时间步嵌入中cond_proj层的维度conv_in_kernel3输入卷积核大小conv_out_kernel3输出卷积核大小projection_class_embeddings_input_dimNoneclass_embed_typeprojection时class_labels的输入维度必须提供attention_typedefault注意力类型如gatedGLIGENclass_embeddings_concatFalse是否将时间嵌入与类别嵌入拼接拼接时传入块的时间嵌入维度翻倍mid_block_only_cross_attentionNone中间块是否仅用交叉注意力only_cross_attention为单个 bool 且此值为None时继承该值cross_attention_normNone交叉注意力的归一化方式addition_embed_type_num_heads64addition_embed_typetext时TextTimeEmbedding的头数3.1 参数校验与广播机制__init__中首先调用_check_config做一致性校验unet_2d_condition.py#L498-L548核心规则包括down_block_types与up_block_types长度必须一致block_out_channels、only_cross_attention、attention_head_dim、cross_attention_dim、layers_per_block等列表型参数的长度必须等于 down block 数量若transformer_layers_per_block为嵌套列表且未提供reverse_transformer_layers_per_block会直接报错非对称 UNet 场景。随后是标量广播逻辑only_cross_attention、num_attention_heads、attention_head_dim、cross_attention_dim、layers_per_block、transformer_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可以看出完整结构输入层conv_in一个将in_channels映射到block_out_channels[0]的nn.Conv2dpadding 由(conv_in_kernel - 1) // 2自动计算时间/条件嵌入层time_projtime_embeddingTimestepEmbedding以及可选的类别嵌入、附加嵌入、编码器隐藏状态投影下采样路径down_blocks按down_block_types逐块构建每个块由若干层组成除最后一个块外都带add_downsampleunet_2d_condition.py#L361-L394中间块mid_block通过get_mid_block按mid_block_type构建unet_2d_condition.py#L396-L418上采样路径up_blocks按up_block_types构建块类型与 down blocks 逆序对称除最终块外都带add_upsample并记录num_upsamplers用于计算整体上采样因子见 unet_2d_condition.py#L420-L477输出层conv_norm_outGroupNormconv_actconv_outunet_2d_condition.py#L479-L494。块类型集中在 unet_2d_blocks.py 中UNetMidBlock2DCrossAttn带交叉注意力的中间块CrossAttnDownBlock2D带交叉注意力的下采样块含 ResNet 块 BasicTransformerBlockDownBlock2D纯 ResNet 的下采样块CrossAttnUpBlock2D带交叉注意力的上采样块UpBlock2D纯 ResNet 的上采样块。以 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丢弃会被转换为 biascross_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添加到中间块残差的张量ControlNetdown_intrablock_additional_residualstuple添加到 down 块内部的残差T2I-Adapter 专用encoder_attention_masktorch.Tensor交叉注意力掩码形状(batch, seq_len)return_dictbool是否返回UNet2DConditionOutput整个流程可划分为六个阶段尺寸与掩码预处理计算default_overall_up_factor 2 ** num_upsamplers若输入空间尺寸不是该因子的整数倍则开启forward_upsample_size以在解码阶段动态插值对齐attention_mask与encoder_attention_mask均被转换为(1 - mask) * -10000.0的注意力 bias并增加单例 query 维度unet_2d_condition.py#L1039-L1074。条件嵌入合成先由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。输入预处理conv_in卷积unet_2d_condition.py#L1107-L1108若attention_type为gatedGLIGEN还会经position_net处理边界框条件。下采样down逐个遍历down_blocks带交叉注意力的块会额外接收encoder_hidden_states、attention_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。中间块mid若配置了 mid block 则执行ControlNet 场景下将mid_block_additional_residual直接加到输出上unet_2d_condition.py#L1170-L1193。上采样up与后处理解码器逐块消费跳连接res_samples非最终块且开启forward_upsample_size时按对应 down 块尺寸插值最后经conv_norm_out、conv_act、conv_out输出与输入同尺寸的预测结果unet_2d_condition.py#L1195-L1230。值得注意的是forward上方的apply_lora_scale(cross_attention_kwargs)装饰器unet_2d_condition.py#L978会在每次前向时应用 LoRA 缩放这就是 LoRA 微调权重能无缝作用于推理的原因。六、条件注入的四种通道源码级解读从forward的实现可以归纳出模型支持的四类条件时间步条件必选timestep经get_time_embed→Timesteps正弦投影 →TimestepEmbedding是扩散过程的进度指示器文本/编码器条件核心encoder_hidden_states经process_encoder_hidden_states处理unet_2d_condition.py#L942-L976按encoder_hid_dim_type支持text_proj线性投影、text_image_projKandinsky 2.1 风格、image_projKandinsky 2.2 风格、ip_image_projImage Prompt 风格四种投影然后注入各交叉注意力块类别条件class_labels经get_class_embedunet_2d_condition.py#L874-L888按class_embed_type支持nn.Embedding、timestep、identity、projection、simple_projection五种实现最终与时间嵌入相加或拼接附加条件added_cond_kwargsget_aug_embedunet_2d_condition.py#L890-L940覆盖五种模式——textTextTimeEmbedding、text_imageKandinsky 2.1、text_timeSDXL 的text_embedstime_ids、imageKandinsky 2.2、image_hintKandinsky 2.2 ControlNet。例如 SDXL 要求added_cond_kwargs必须包含text_embeds与time_ids否则会抛ValueError。正是这套时间 文本 类别 附加的多通道条件注入机制让同一个UNet2DConditionModel能服务于文本到图像Stable Diffusion、图像到图像img2img、修复inpainting、ControlNet 引导、T2I-Adapter 等多种扩散管线。七、进阶 API注意力切片、FreeU 与 QKV 融合除前向推理外模型还内置了一系列实用方法7.1 set_attention_slice低显存注意力set_attention_slice(slice_sizeauto)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_PROCESSORS或CROSS_ATTENTION_PROCESSORS自动回退到AttnAddedKVProcessor或AttnProcessor用于清除自定义处理器。7.5 其他模型级属性源码还声明了若干辅助属性_supports_gradient_checkpointing True支持梯度检查点训练、_no_split_modulesFSDP/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, subfolderunet) # 加载到指定设备与精度 unet.to(cuda, dtypetorch.float16)由于UNet2DConditionModel继承自ModelMixin与ConfigMixinfrom_pretrained、save_pretrained、push_to_hub等通用能力开箱即用。8.2 从零实例化自定义配置from diffusers import UNet2DConditionModel unet UNet2DConditionModel( sample_size64, # 训练分辨率 in_channels4, # 潜在空间通道数 out_channels4, down_block_types( CrossAttnDownBlock2D, CrossAttnDownBlock2D, CrossAttnDownBlock2D, DownBlock2D, ), up_block_types( UpBlock2D, CrossAttnUpBlock2D, CrossAttnUpBlock2D, CrossAttnUpBlock2D, ), block_out_channels(320, 640, 1280, 1280), layers_per_block2, cross_attention_dim768, # 文本编码器嵌入维度CLIP 为 768 attention_head_dim8, norm_num_groups32, )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(samplesample, timesteptimestep, encoder_hidden_statesencoder_hidden_states) # output 为 UNet2DConditionOutputoutput.sample 形状为 (2, 4, 64, 64)8.4 测试覆盖功能验证的说明书仓库测试 test_models_unet_2d_condition.py 提供了官方验证逻辑是理解模型行为的极佳参考TestUNet2DCondition 继承ModelTesterMixin与UNetTesterMixin覆盖输入输出形状、前向稳定性、训练模式等通用检查专项测试包括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_blockTestUNet2DConditionHubLoading 覆盖分片 checkpoint 的 Hub 加载、subfolder 加载等场景。此外tests/lora/utils.py 与 tests/models/test_modeling_common.py 也引用了该模型分别验证 LoRA 适配器与通用建模能力。九、生态中的位置它支撑了哪些管线在 Diffusers 中UNet2DConditionModel是大量 2D 图像扩散管线的去噪核心例如Stable Diffusion 系列文本条件扩散使用CrossAttnDownBlock2D/CrossAttnUpBlock2D注入 CLIP 文本嵌入Stable Diffusion XL通过addition_embed_typetext_time注入 pooled text 嵌入与尺寸时间 ID并启用use_linear_projectionKandinsky 2.x通过addition_embed_type的text_image/image/image_hint与encoder_hid_dim_type的image_proj实现图像条件注入ControlNet / T2I-Adapter通过down_block_additional_residuals与down_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.项目地址: https://gitcode.com/GitHub_Trending/di/diffusers创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表