使用UNet2DConditionModel遇矩阵形状不匹配RuntimeError求助
解决CLIPTextModel与UNet2DConditionModel训练时的形状不匹配错误
错误原因
这个RuntimeError: mat1 and mat2 shapes cannot be multiplied (4928x768 and 1280x128)是因为UNet内部处理文本条件的投影层,输入维度和CLIPTextModel输出的文本嵌入维度不兼容:
- CLIP输出的文本张量维度是
[64,77,768](最后一维为768) - UNet的文本投影层期望输入维度为1280,导致矩阵乘法时形状不匹配。
解决方案
1. 匹配模型版本与权重
确保CLIPTextModel和UNet2DConditionModel来自同一套预训练权重,比如用Stable Diffusion v1.x的CLIP配同版本的UNet,SDXL的CLIP配SDXL的UNet(SDXL的CLIP输出维度是1280,适配对应UNet)。示例代码:
from transformers import CLIPTextModel from diffusers import UNet2DConditionModel # 加载Stable Diffusion v1.5配套模型 text_encoder = CLIPTextModel.from_pretrained("runwayml/stable-diffusion-v1-5/text_encoder") unet = UNet2DConditionModel.from_pretrained("runwayml/stable-diffusion-v1-5/unet")
2. 手动修改UNet的文本投影层
如果必须使用当前的UNet,可以重新定义它的文本投影层,将输入维度改为CLIP输出的768:
import torch.nn as nn # 获取原投影层的输出维度 output_dim = unet.text_projection.out_features # 替换投影层,输入维度适配CLIP的768 unet.text_projection = nn.Linear(768, output_dim)
3. 检查文本嵌入的处理流程
确认你没有对CLIP输出的文本嵌入做错误的维度变换(比如不必要的flatten、池化),保证传入UNet的文本张量格式为[batch_size, seq_len, 768],UNet会自动处理序列维度的均值池化操作。
验证
修改后重新运行训练代码,确认noise_pred = model(...)行不再报错,检查图像张量torch.Size([64,4,64,64])和文本张量torch.Size([64,77,768])能正常完成前向传播。
内容的提问来源于stack exchange,提问作者Artsem Holub
相关产品推荐
相关产品推荐

