拓冰建站拓冰建站
首页 / 资讯中心 / 正文

Diffusers 中的 SanaControlNetModel:为 Sana 文本生成图像模型添加空间条件控制

Diffusers 中的 SanaControlNetModel为 Sana 文本生成图像模型添加空间条件控制【免费下载链接】diffusers Diffusers: State-of-the-art diffusion models for image, video, and audio generation in PyTorch.项目地址: https://gitcode.com/GitHub_Trending/di/diffusers导读本文基于 Hugging Face Diffusers 仓库的 API 文档与源码系统讲解SanaControlNetModel的模型设计、构造参数、前向传播流程以及它与SanaControlNetPipeline的协作方式。读完本文你将掌握如何加载官方 Sana ControlNet 检查点用边缘图HED、深度图、分割图等条件图对 Sana 文本生成图像模型施加精确的空间控制并能理解零初始化模块在其中的核心作用。一、背景ControlNet 条件控制与 Sana 的结合ControlNet 架构最初由论文Adding Conditional Control to Text-to-Image Diffusion Models作者 Lvmin Zhang、Anyi Rao、Maneesh Agrawala提出其核心思想是在不破坏大规模预训练扩散模型的前提下为其附加一个可训练的条件控制分支从而让生成过程接受边缘图、深度图、分割图、人体关键点等额外空间条件的约束。论文摘要的核心要点如下ControlNet 将已经生产就绪的大型扩散模型锁定复用其经过数十亿图像预训练得到的深层编码层作为强骨干网络去学习多样化的条件控制神经架构通过零卷积zero-initialized convolution layers连接——即从零开始逐步增长参数的零初始化卷积层确保微调过程中不会引入有害噪声论文使用边缘、深度、分割、人体姿态等多种条件控制结合 Stable Diffusion 进行验证支持单一或多重条件、有无提示词等组合实验表明 ControlNet 的训练在小数据集50k和大数据集1m上都很稳健。在本仓库中SanaControlNetModel是把这一思想落地到 Sana 模型一种高效、可缩放的大规模文本生成图像 Transformer 架构上的具体实现。该模型的原始代码库来自 NVlabs/Sana官方 ControlNet 检查点发布在 Efficient-Large-Model 的模型库中。SanaControlNetModel由社区开发者 ishan24 贡献其配套的推理管线为SanaControlNetPipeline源码位于 src/diffusers/pipelines/sana/pipeline_sana_controlnet.py。与经典 ControlNet 的 UNet 结构不同Sana 的骨干是 DiTDiffusion Transformer因此SanaControlNetModel的控制分支也是一条并行的 Transformer 路径它复用SanaTransformerBlock提取各层隐藏状态再用零初始化的线性层把每一层的残差样本输出给主 Transformer由主 Transformer 在每个对应模块处把条件残差叠加回自身隐藏状态。二、SanaControlNetModel 类概览SanaControlNetModel定义于 src/diffusers/models/controlnets/controlnet_sana.py类签名如下class SanaControlNetModel(ModelMixin, AttentionMixin, ConfigMixin, PeftAdapterMixin)它组合了四个基类各自承担不同职责基类职责ModelMixin提供模型保存/加载save_pretrained/from_pretrained、参数统计、跨设备部署等通用能力AttentionMixin提供注意力处理器的挂载与切换能力支持自定义AttentionProcessor如 xFormers、SDPA 等ConfigMixin提供配置管理能力__init__中所有参数经register_to_config自动写入self.configPeftAdapterMixin支持 PEFT如 LoRA适配层的注入与融合类级别还声明了三个与训练/推理优化相关的属性_supports_gradient_checkpointing True _no_split_modules [SanaTransformerBlock, PatchEmbed] _skip_layerwise_casting_patterns [patch_embed, norm]_supports_gradient_checkpointing True支持梯度检查点可在训练时以计算换显存_no_split_modules声明这些模块在模型并行如 device_map 切分时不应被拆分_skip_layerwise_casting_patterns声明patch_embed与norm相关层在逐层低精度转换layerwise casting时跳过用于保证数值稳定性。三、构造函数参数详解SanaControlNetModel.__init__的全部参数如下均为关键字参数且全部写入 config参数类型默认值说明in_channelsint32输入潜在表示的通道数与主 Transformer 的输入通道一致例如 4 倍下采样 VAE 的潜在通道 × 组归一化展开out_channelsint \| None32输出通道数None时回退为in_channelsnum_attention_headsint70自注意力头数attention_head_dimint32每个自注意力头的维度num_layersint7复用SanaTransformerBlock的层数num_cross_attention_headsint \| None20交叉注意力头数cross_attention_head_dimint \| None112每个交叉注意力头的维度cross_attention_dimint \| None2240交叉注意力输入文本嵌入的维度caption_channelsint2304文本编码器Gemma2输出嵌入的维度mlp_ratiofloat2.5MLP 隐藏层相对inner_dim的放大比例dropoutfloat0.0Dropout 概率attention_biasboolFalse注意力层是否使用偏置sample_sizeint32输入潜在图的空间尺寸高宽决定 PatchEmbed 的位置编码网格patch_sizeint1Patch 大小1表示不做 patch 合并norm_elementwise_affineboolFalse归一化层是否使用逐元素仿射参数norm_epsfloat1e-6归一化层的 epsiloninterpolation_scaleint \| NoneNone位置编码插值缩放非None时启用 sincos 位置编码代码中的关键派生关系out_channels out_channels or in_channels inner_dim num_attention_heads * attention_head_dim # 默认 70 × 32 2240inner_dim默认 2240是整个模型的隐藏维度也即各 Transformer 模块、条件嵌入模块的统一宽度。测试用例 tests/pipelines/sana/test_sana_controlnet.py 中给出了一个迷你配置示例便于快速验证结构与前向正确性controlnet SanaControlNetModel( patch_size1, in_channels4, out_channels4, num_layers1, num_attention_heads2, attention_head_dim4, num_cross_attention_heads2, cross_attention_head_dim4, cross_attention_dim8, caption_channels8, sample_size32, )四、模型内部结构五个组成部分从源码看SanaControlNetModel的__init__将网络组织为五个部分1. Patch 嵌入Patch Embeddingself.patch_embed PatchEmbed( heightsample_size, widthsample_size, patch_sizepatch_size, in_channelsin_channels, embed_diminner_dim, interpolation_scaleinterpolation_scale, pos_embed_typesincos if interpolation_scale is not None else None, )PatchEmbed位于 src/diffusers/models/embeddings.py将图像潜在张量从(batch, in_channels, H, W)展平为 token 序列并叠加位置编码。该模块同时被主SanaTransformer2DModel使用保证两侧的 token 空间完全对齐。2. 附加条件嵌入Additional Condition Embeddingsself.time_embed AdaLayerNormSingle(inner_dim) # 时间步自适应层归一化 self.caption_projection PixArtAlphaTextProjection(in_featurescaption_channels, hidden_sizeinner_dim) # 文本投影 self.caption_norm RMSNorm(inner_dim, eps1e-5, elementwise_affineTrue) # 文本嵌入归一化AdaLayerNormSingle把去噪时间步timestep编码并映射为自适应归一化的缩放/平移参数与 Sana 主干的时间条件方式一致PixArtAlphaTextProjection把 Gemma2 文本编码器的输出caption_channels2304维投影到inner_dimRMSNorm对投影后的文本嵌入做归一化。3. Transformer 模块栈self.transformer_blocks nn.ModuleList( [SanaTransformerBlock(...) for _ in range(num_layers)] )这里直接复用 Sana 主干的SanaTransformerBlock定义于 src/diffusers/models/transformers/sana_transformer.py意味着 ControlNet 的条件分支与主生成分支共享相同的 Transformer 结构从而能够对齐每一层的语义特征。4. 零初始化输入模块与逐层控制模块self.input_block zero_module(nn.Linear(inner_dim, inner_dim)) for _ in range(len(self.transformer_blocks)): controlnet_block nn.Linear(inner_dim, inner_dim) controlnet_block zero_module(controlnet_block) self.controlnet_blocks.append(controlnet_block)zero_module是经典 ControlNet 中零卷积思想的直接体现其实现位于 src/diffusers/models/controlnets/controlnet.py将传入模块的所有参数初始化为零。与原始 ControlNet 使用零初始化卷积层不同Sana 版本使用零初始化的线性层input_block接收controlnet_cond经patch_embed编码后的条件 token在进入 Transformer 栈之前就注入条件信息每个controlnet_block对应一个 Transformer 层负责把该层的输出特征投影成控制残差样本。由于这些模块参数从零开始训练初期条件分支不会干扰主模型输出从而保证微调稳定——这正是论文所述zero convolutions 确保不会有害噪声影响微调的设计意图。5. 梯度检查点开关self.gradient_checkpointing False前向传播时会检查该开关开启状态下通过_gradient_checkpointing_func逐层计算 Transformer 块用于降低训练显存占用。五、forward 前向传播流程forward方法的签名其余参数均可选def forward( self, hidden_states: torch.Tensor, # (batch, channel, height, width) 加噪潜在图 encoder_hidden_states: torch.Tensor, # 文本条件嵌入如 prompt 编码结果 timestep: torch.LongTensor, # 去噪时间步 controlnet_cond: torch.Tensor, # ControlNet 条件输入张量 conditioning_scale: float 1.0, # ControlNet 输出缩放因子 encoder_attention_maskNone, # 文本嵌入注意力掩码 attention_maskNone, # 隐藏状态注意力掩码 attention_kwargsNone, # 传给 AttentionProcessor 的额外参数 return_dict: bool True, )前向流程按源码顺序分为以下阶段第 1 步掩码转 bias。若attention_mask/encoder_attention_mask是二维掩码1保留0丢弃先转换为可广播的注意力 bias(1 - mask) * -10000.0并补一个长度为 1 的 query 维度以便广播到[batch, heads, query, key]形状的注意力分数上。第 2 步输入注入。计算 patch 后的空间尺寸然后hidden_states self.patch_embed(hidden_states) hidden_states hidden_states self.input_block(self.patch_embed(controlnet_cond.to(hidden_states.dtype)))条件图先经过patch_embed编码与主潜在图共享相同的 patch 化与位置编码方式再过零初始化的input_block以加法方式注入主隐藏状态。从源码结构看这里隐含一个前提controlnet_cond的通道数与in_channels一致且空间尺寸需与hidden_states匹配条件图在管线中会先经 VAE 编码到潜在空间。第 3 步时间步与文本条件嵌入。timestep, embedded_timestep self.time_embed(timestep, batch_sizebatch_size, hidden_dtypehidden_states.dtype) encoder_hidden_states self.caption_projection(encoder_hidden_states) encoder_hidden_states encoder_hidden_states.view(batch_size, -1, hidden_states.shape[-1]) encoder_hidden_states self.caption_norm(encoder_hidden_states)第 4 步Transformer 栈并收集逐层残差样本。依次执行每个SanaTransformerBlock并把每一层的输出hidden_states追加进block_res_samples元组block_res_samples () for block in self.transformer_blocks: hidden_states block(hidden_states, attention_mask, encoder_hidden_states, encoder_attention_mask, timestep, post_patch_height, post_patch_width) block_res_samples block_res_samples (hidden_states,)第 5 步控制模块投影与缩放。controlnet_block_res_samples () for block_res_sample, controlnet_block in zip(block_res_samples, self.controlnet_blocks): controlnet_block_res_samples controlnet_block_res_samples (controlnet_block(block_res_sample),) controlnet_block_res_samples [sample * conditioning_scale for sample in controlnet_block_res_samples]每个残差样本经零初始化的controlnet_block线性投影后统一乘以conditioning_scale默认 1.0。这就是在管线层面调节条件强度的机制scale 越大条件约束越强。第 6 步返回。return_dictFalse时返回(controlnet_block_res_samples,)元组否则返回SanaControlNetOutput(controlnet_block_samples...)。六、SanaControlNetOutput 输出结构dataclass class SanaControlNetOutput(BaseOutput): controlnet_block_samples: tuple[torch.Tensor]SanaControlNetOutput继承自BaseOutput具备属性访问与字典互转能力只包含一个字段controlnet_block_samples由每个 Transformer 层投影出的控制残差样本组成的元组。主SanaTransformer2DModel会按层接收这些样本并叠加到自身隐藏状态上。在 src/diffusers/models/transformers/sana_transformer.py 中可以找到主 Transformer 的消费逻辑if controlnet_block_samples is not None and 0 index_block len(controlnet_block_samples): hidden_states hidden_states controlnet_block_samples[index_block - 1]即主模型的第i个 Transformer 块会把controlnet_block_samples[i-1]加到自己的输出上——这正是条件残差注入的最终落点也印证了SanaControlNetModel的num_layers必须与主 Transformer 的层数匹配的设计前提。七、端到端使用SanaControlNetPipeline 实战1. 管线组成与加载示例SanaControlNetPipelinesrc/diffusers/pipelines/sana/pipeline_sana_controlnet.py由五个模块组成组件类型作用tokenizerGemmaTokenizer/GemmaTokenizerFast文本分词text_encoderGemma2PreTrainedModel文本编码vaeAutoencoderDC图像↔潜在空间编解码transformerSanaTransformer2DModel主去噪骨干controlnetSanaControlNetModel空间条件控制分支schedulerDPMSolverMultistepScheduler去噪调度器管线文档字符串中的官方示例节选自 pipeline_sana_controlnet.py展示了完整用法其中torch_dtype支持按组件分别指定精度并用device_mapbalanced做多设备均衡放置import torch from diffusers import SanaControlNetPipeline from diffusers.utils import load_image pipe SanaControlNetPipeline.from_pretrained( ishan24/Sana_600M_1024px_ControlNetPlus_diffusers, variantfp16, torch_dtype{default: torch.bfloat16, controlnet: torch.float16, transformer: torch.float16}, device_mapbalanced, ) # 加载一张 HED 边缘图作为条件输入也可用 PIL.Image.open 读取本地条件图 cond_image load_image(path/to/hed_example.png) prompt a cat with a neon sign that says Sana image pipe( prompt, control_imagecond_image, ).images[0] image.save(output.png)管线内部通过SanaLoraLoaderMixin还支持为 ControlNet 与 Transformer 注入 LoRA 适配器。2. 条件图的编码链路SanaControlNetModel的输入controlnet_cond不是原始像素图而是条件图在潜在空间中的表示。管线的处理链路见 pipeline_sana_controlnet.pycontrol_image self.prepare_image( imagecontrol_image, widthwidth, heightheight, batch_sizebatch_size * num_images_per_prompt, num_images_per_promptnum_images_per_prompt, devicedevice, dtypeself.vae.dtype, do_classifier_free_guidanceself.do_classifier_free_guidance, guess_modeFalse, ) control_image self.vae.encode(control_image).latent control_image control_image * self.vae.config.scaling_factorprepare_image用PixArtImageProcessor将条件图缩放到目标尺寸并在启用无分类器引导CFG时沿 batch 维度复制一份torch.cat([image] * 2)条件图经AutoencoderDC编码为潜在表示再乘以 VAE 的scaling_factor归一化得到与主潜在图同尺度、同通道数的controlnet_cond。这也解释了模型参数中in_channels与主 Transformer 输入通道一致的原因若传入的controlnet不是SanaControlNetModel类型管线会直接抛出ValueError提示。3. 去噪循环中的条件注入去噪循环内pipeline_sana_controlnet.py每一步先执行 ControlNet 得到残差样本再交给主 Transformercontrolnet_block_samples self.controlnet( latent_model_input.to(dtypecontrolnet_dtype), encoder_hidden_statesprompt_embeds.to(dtypecontrolnet_dtype), encoder_attention_maskprompt_attention_mask, timesteptimestep, return_dictFalse, attention_kwargsself.attention_kwargs, controlnet_condcontrol_image, conditioning_scalecontrolnet_conditioning_scale, )[0] noise_pred self.transformer( latent_model_input.to(dtypetransformer_dtype), encoder_hidden_statesprompt_embeds.to(dtypetransformer_dtype), encoder_attention_maskprompt_attention_mask, timesteptimestep, return_dictFalse, attention_kwargsself.attention_kwargs, controlnet_block_samplestuple(t.to(dtypetransformer_dtype) for t in controlnet_block_samples), )[0]注意这里 ControlNet 与主 Transformer 分别以自己的 dtype 运行支持torch_dtype字典按组件混合精度残差样本在传递前统一转换为主 Transformer 的精度。CFG 时latent_model_input已拼接为双份控制条件同样按双份处理保证引导计算的对应关系正确。4. 关键调用参数速查__call__方法pipeline_sana_controlnet.py的关键参数及默认值参数默认值说明promptNone提示词字符串或列表与prompt_embeds二选一negative_prompt负向提示词Sana 建议为num_inference_steps20去噪步数越多质量越高但更慢guidance_scale4.5无分类器引导强度1 时启用 CFGcontrol_imageNone条件输入支持 PIL 图、torch.Tensor、np.ndarray及对应列表controlnet_conditioning_scale1.0条件强度多 ControlNet 时可用列表逐网设置height/width1024输出图像尺寸须能被 32 整除use_resolution_binningTrue是否先把尺寸映射到最近的分辨率档位如 512/1024/2048/4096 bin生成后再裁剪回请求尺寸clean_captionFalse是否清洗提示词需要beautifulsoup4与ftfymax_sequence_length300提示词最大序列长度complex_human_instruction内置列表复杂人工指令Complex Human Instruction用于提示词增强传None可关闭output_typepil输出格式可选pil、np、latent、pt此外管线还支持timesteps/sigmas自定义调度、callback_on_step_end步进回调、latents预置噪声、attention_kwargs如 LoRA scale等高级能力。5. 内存优化手段管线声明了模型卸载顺序text_encoder-controlnet-transformer-vae可配合enable_model_cpu_offload()使用。测试 tests/pipelines/sana/test_sana_controlnet.py 中的TestSanaControlNetPipelineMemory专门覆盖了 CPU offload、group offload 与逐层低精度转换layerwise casting的内存优化路径。另外test_vae_tiling验证了 VAE 分块tiling解码在高分辨率128px 以上下不影响生成结果大图显存不足时可用pipe.vae.enable_tiling( tile_sample_min_height96, tile_sample_min_width96, tile_sample_stride_height64, tile_sample_stride_width64, )八、测试与验证仓库为SanaControlNetModel与配套管线提供了完整的测试支撑tests/pipelines/sana/test_sana_controlnet.py 中TestSanaControlNetPipeline.test_inference用随机噪声条件图执行完整的前向生成并断言输出形状(3, 32, 32)test_vae_tiling验证 VAE 分块前后生成结果的最大绝对差小于 0.2TestSanaControlNetPipelineMemory验证各类内存优化路径测试中的 dummy 配置同时给出了SanaControlNetModel与SanaTransformer2DModel的最小可运行参数组合可作为自建小模型的参考模板。九、总结SanaControlNetModel是 ControlNet 思想在 Sana 这一 Transformer 扩散架构上的完整落地它以与主模型同构的SanaTransformerBlock栈提取逐层特征通过零初始化的线性层zero_module实现于 src/diffusers/models/controlnets/controlnet.py实现条件残差注入配合SanaControlNetPipeline即可用边缘图、深度图等条件精确控制 Sana 的图像生成。理解它的构造参数尤其是num_layers需与主 Transformer 匹配与前向流程条件图先经 VAE 编码到潜在空间是将其用于实际生成或自定义训练的关键前提。【免费下载链接】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),仅供参考
分享:

看完干货,该让你的企业上线了

免费需求沟通 · 48 小时内出具建站方案 · 河南本地可上门