微调Stable Video Diffusion遇矩阵相乘形状不匹配错误的解决咨询
解决Stable Video Diffusion微调时的矩阵形状不匹配错误
问题详情
我在用自有数据集微调stable-video-diffusion-img2vid-xt-1-1模型,执行output = lora_unet代码行时触发矩阵乘法形状不匹配错误,ChatGPT给出的修复方案无效。
当前参数形状
Frames shape: torch.Size([4, 60, 3, 128, 128]) Timesteps shape: torch.Size([4]) Encoder hidden states shape: torch.Size([4, 768]) Added time IDs shape: torch.Size([4, 1])
错误信息
RuntimeError: mat1 and mat2 shapes cannot be multiplied (4x256 and 768x1280)
相关代码
for frames, labels in tqdm(alphabet_loader): frames = frames.to(device) # Create a label mapping and numerical labels label_mapping = {label: idx for idx, label in enumerate(labels)} numeric_labels = [label_mapping[label] for label in labels] labels = torch.tensor(numeric_labels).to(device) # Random timesteps (typically, a value between 0 and 1000 for diffusion models) timesteps = torch.randint(0, 1000, (frames.size(0),), device=device) # Number of classes and embedding dimension num_classes = len(label_mapping) embedding_dim = 256 # Your current embedding size (this can remain 256 for the embedding) label_embedding = nn.Embedding(num_classes, embedding_dim).to(device) # Generate encoder hidden states from label embeddings encoder_hidden_states = label_embedding(labels) # Shape: [batch_size, num_labels] # Project the encoder hidden states to the required dimension (768) projection_layer = nn.Linear(embedding_dim, 768).to(device) # Project from 256 to 768 encoder_hidden_states = projection_layer(encoder_hidden_states) # Now shape is [batch_size, 768] # Added time IDs (optional, adjust as needed) added_time_ids = torch.zeros_like(timesteps, device=device) added_time_ids = added_time_ids.unsqueeze(-1) # Convert [batch_size] to [batch_size, 1] # Permute frames to the expected shape [batch_size, timesteps, channels, height, width] #frames = frames.permute(0, 2, 1, 3, 4) # [B, T, C, H, W] -> [B, T, C, H, W] # Print shapes for debugging print(f"Frames shape: {frames.shape}") print(f"Timesteps shape: {timesteps.shape}") print(f"Encoder hidden states shape: {encoder_hidden_states.shape}") print(f"Added time IDs shape: {added_time_ids.shape}") # Forward pass optimizer.zero_grad() output = lora_unet( frames, timestep=timesteps, encoder_hidden_states=encoder_hidden_states, added_time_ids=added_time_ids, )["sample"] # Compute the loss (ensure labels are in float format for MSELoss) criterion = nn.MSELoss() loss = criterion(output, labels.float()) # Ensure type compatibility loss.backward() optimizer.step() epoch_loss += loss.item()
错误栈追踪
Traceback (most recent call last): File "/Library/Frameworks/Python.framework/Versions/3.12/lib/python3.12/runpy.py", line 198, in _run_module_as_main return _run_code(code, main_globals, None, ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^ File "/Library/Frameworks/Python.framework/Versions/3.12/lib/python3.12/runpy.py", line 88, in _run_code exec(code, run_globals) File "/Users/sunaynatalreja/.vscode/extensions/ms-python.debugpy-2024.14.0-darwin-arm64/bundled/libs/debugpy/adapter/../../debugpy/launcher/../../debugpy/__main__.py", line 71, in cli.main() File "/Users/sunaynatalreja/.vscode/extensions/ms-python.debugpy-2024.14.0-darwin-arm64/bundled/libs/debugpy/adapter/../../debugpy/launcher/../../debugpy/../debugpy/server/cli.py", line 501, in main run() File "/Users/sunaynatalreja/.vscode/extensions/ms-python.debugpy-2024.14.0-darwin-arm64/bundled/libs/debugpy/adapter/../../debugpy/launcher/../../debugpy/../debugpy/server/cli.py", line 351, in run_file runpy.run_path(target, run_name="__main__") File "/Users/sunaynatalreja/.vscode/extensions/ms-python.debugpy-2024.14.0-darwin-arm64/bundled/libs/debugpy/_vendored/pydevd/_pydevd_bundle/pydevd_runpy.py", line 310, in run_path return _run_module_code(code, init_globals, run_name, pkg_name=pkg_name, script_name=fname) ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^ File "/Users/sunaynatalreja/.vscode/extensions/ms-python.debugpy-2024.14.0-darwin-arm64/bundled/libs/debugpy/_vendored/pydevd/_pydevd_bundle/pydevd_runpy.py", line 127, in _run_module_code _run_code(code, mod_globals, init_globals, mod_name, mod_spec, pkg_name, script_name) File "/Users/sunaynatalreja/.vscode/extensions/ms-python.debugpy-2024.14.0-darwin-arm64/bundled/libs/debugpy/_vendored/pydevd/_pydevd_bundle/pydevd_runpy.py", line 118, in _run_code exec(code, run_globals) File "/Users/sunaynatalreja//FrontEnd/.py", line 288, in output = lora_unet( ^^^^^^^^^^ File "/Library/Frameworks/Python.framework/Versions/3.12/lib/python3.12/site-packages/torch/nn/modules/module.py", line 1736, in _wrapped_call_impl return self._call_impl(*args, **kwargs) ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^ File "/Library/Frameworks/Python.framework/Versions/3.12/lib/python3.12/site-packages/torch/nn/modules/module.py", line 1747, in _call_impl return forward_call(*args, **kwargs) ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^ File "/Library/Frameworks/Python.framework/Versions/3.12/lib/python3.12/site-packages/peft/peft_model.py", line 849, in forward return self.get_base_model()(*args, **kwargs) ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^ File "/Library/Frameworks/Python.framework/Versions/3.12/lib/python3.12/site-packages/torch/nn/modules/module.py", line 1736, in _wrapped_call_impl return self._call_impl(*args, **kwargs) ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^ File "/Library/Frameworks/Python.framework/Versions/3.12/lib/python3.12/site-packages/torch/nn/modules/module.py", line 1747, in _call_impl return forward_call(*args, **kwargs) ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^ File "/Library/Frameworks/Python.framework/Versions/3.12/lib/python3.12/site-packages/diffusers/models/unets/unet_spatio_temporal_condition.py", line 429, in forward aug_emb = self.add_embedding(time_embeds) ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^ File "/Library/Frameworks/Python.framework/Versions/3.12/lib/python3.12/site-packages/torch/nn/modules/module.py", line 1736, in _wrapped_call_impl return self._call_impl(*args, **kwargs) ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^ File "/Library/Frameworks/Python.framework/Versions/3.12/lib/python3.12/site-packages/torch/nn/modules/module.py", line 1747, in _call_impl return forward_call(*args, **kwargs) ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^ File "/Library/Frameworks/Python.framework/Versions/3.12/lib/python3.12/site-packages/diffusers/models/embeddings.py", line 1304, in forward sample = self.linear_1(sample) ^^^^^^^^^^^^^^^^^^^^^ File "/Library/Frameworks/Python.framework/Versions/3.12/lib/python3.12/site-packages/torch/nn/modules/module.py", line 1736, in _wrapped_call_impl return self._call_impl(*args, **kwargs) ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^ File "/Library/Frameworks/Python.framework/Versions/3.12/lib/python3.12/site-packages/torch/nn/modules/module.py", line 1747, in _call_impl return forward_call(*args, **kwargs) ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^ File "/Library/Frameworks/Python.framework/Versions/3.12/lib/python3.12/site-packages/torch/nn/modules/linear.py", line 125, in forward return F.linear(input, self.weight, self.bias) ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^ RuntimeError: mat1 and mat2 shapes cannot be multiplied (4x256 and 768x1280)
解决方案
问题根源
从错误栈可知,问题出在模型的add_embedding模块:linear_1层期望输入维度为768,但实际传入的是256维度的张量。这是因为added_time_ids未被正确投影到模型要求的维度,同时代码中存在每次循环重置嵌入层参数的致命错误。
修复步骤
将嵌入/投影层移到训练循环外初始化
不能在每个batch循环中重新定义nn.Embedding和nn.Linear,否则这些层的参数永远无法被优化。将层定义移到训练前:# 训练前初始化(替换成你的类别数获取方式) num_classes = len(alphabet_loader.dataset.classes) embedding_dim = 256 label_embedding = nn.Embedding(num_classes, embedding_dim).to(device) projection_layer = nn.Linear(embedding_dim, 768).to(device) added_time_proj = nn.Linear(1, 768).to(device) # 优化器包含所有需要训练的参数 optimizer = torch.optim.Adam([ {'params': lora_unet.parameters(), 'lr': 1e-5}, {'params': label_embedding.parameters(), 'lr': 1e-4}, {'params': projection_layer.parameters(), 'lr': 1e-4}, {'params': added_time_proj.parameters(), 'lr': 1e-4} ], lr=1e-4)修复
added_time_ids的维度匹配
为added_time_ids添加投影层,将其从1维映射到模型期望的768维:# 循环内处理added_time_ids added_time_ids = torch.zeros_like(timesteps, device=device).unsqueeze(-1) added_time_ids = added_time_proj(added_time_ids) # 形状变为[4,768]若不需要
added_time_ids,可直接删除该参数的传入。调整帧的形状为模型期望格式
Stable Video Diffusion的Unet通常期望帧形状为[batch_size, channels, num_frames, height, width],转换当前帧格式:frames = frames.permute(0, 2, 1, 3, 4) # 从[B,T,C,H,W]转为[B,C,T,H,W]修正损失计算的形状不匹配
原代码中用生成帧和标签直接计算MSE是错误的,生成任务应使用生成帧与真实帧的重构损失:# 替换原损失计算代码 criterion = nn.MSELoss() loss = criterion(output, frames)
修复后的核心代码片段
# 训练前初始化层 num_classes = len(alphabet_loader.dataset.classes) embedding_dim = 256 label_embedding = nn.Embedding(num_classes, embedding_dim).to(device) projection_layer = nn.Linear(embedding_dim, 768).to(device) added_time_proj = nn.Linear(1, 768).to(device) # 优化器配置 optimizer = torch.optim.Adam([ {'params': lora_unet.parameters(), 'lr': 1e-5}, {'params': label_embedding.parameters(), 'lr': 1e-4}, {'params': projection_layer.parameters(), 'lr': 1e-4}, {'params': added_time_proj.parameters(), 'lr': 1e-4} ], lr=1e-4) criterion = nn.MSELoss() # 训练循环 for frames, labels in tqdm(alphabet_loader): frames = frames.to(device) # 转换帧形状为模型期望格式 frames = frames.permute(0, 2, 1, 3, 4) # 处理标签 label_mapping = {label: idx for idx, label in enumerate(labels)} numeric_labels = [label_mapping[label] for label in labels] labels = torch.tensor(numeric_labels).to(device) timesteps = torch.randint(0, 1000, (frames.size(0),), device=device) # 生成条件嵌入 encoder_hidden_states = label_embedding(labels) encoder_hidden_states = projection_layer(encoder_hidden_states) # 处理added_time_ids added_time_ids = torch.zeros_like(timesteps, device=device).unsqueeze(-1) added_time_ids = added_time_proj(added_time_ids) # 前向传播 optimizer.zero_grad() output = lora_unet( frames, timestep=timesteps, encoder_hidden_states=encoder_hidden_states, added_time_ids=added_time_ids, )["sample"] # 计算重构损失 loss = criterion(output, frames) loss.backward() optimizer.step() epoch_loss += loss.item()
内容的提问来源于stack exchange,提问作者Sunayna Talreja
相关产品推荐
相关产品推荐

