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

微调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未被正确投影到模型要求的维度,同时代码中存在每次循环重置嵌入层参数的致命错误。

修复步骤

  1. 将嵌入/投影层移到训练循环外初始化
    不能在每个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)
    
  2. 修复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,可直接删除该参数的传入。

  3. 调整帧的形状为模型期望格式
    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]
    
  4. 修正损失计算的形状不匹配
    原代码中用生成帧和标签直接计算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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.15 07:09:50