You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

训练diffusers/UNet2DConditionModel时矩阵形状不匹配问题求助

修复UNet2DConditionModel前向传播的矩阵形状不匹配错误

错误原因

你遇到的mat1与mat2形状无法相乘问题,核心是自定义UNet时未指定与文本编码器输出维度匹配的cross_attention_dim参数。

CLIP-ViT-base-patch32文本编码器的输出特征维度为512,但你初始化UNet2DConditionModel时未设置该参数,导致模型默认使用的交叉注意力层线性参数维度(错误信息中的1280)与输入的文本特征维度(512)不匹配。错误中的288是批次大小×文本序列长度(比如batch_size=16时,16×18=288),因此会随批次变化。

修复方案

不需要补零,直接修正UNet的初始化参数,添加cross_attention_dim=512,确保与CLIP文本编码器的hidden_size一致:

unet = UNet2DConditionModel(
    in_channels=4,
    out_channels=4,
    layers_per_block=2,
    sample_size=64,
    block_out_channels=(128, 256, 512, 512),
    down_block_types=("DownBlock2D", "DownBlock2D", "DownBlock2D", "AttnDownBlock2D"),
    up_block_types=("AttnUpBlock2D", "UpBlock2D", "UpBlock2D", "UpBlock2D"),
    cross_attention_dim=512,  # 新增:匹配CLIP文本编码器的输出维度
).to(device)

额外说明

  • 若后续仍有注意力相关维度错误,可检查attention_head_dim参数,确保其能被cross_attention_dim整除(比如设置为8,512÷8=64,符合注意力头常规配置)。
  • 文本编码器的输出序列长度不影响矩阵相乘,UNet的交叉注意力层会自动处理不同长度的文本序列。

内容的提问来源于stack exchange,提问作者u1ug

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.19 22:07:35