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

如何在nn.Sequential中展平nn.Embedding?维度不匹配问题求解

问题分析与解决

你的问题出在原模型的view((1, -1))和你用的nn.Flatten(1, -1)作用维度不一致:

原模型中,输入inputs的形状是[context_size](比如n-gram上下文有2个词,就是长度为2的张量),经过Embedding后得到[context_size, embedding_dim]的张量,view((1, -1))是把这个张量重塑成[1, context_size*embedding_dim]——本质是给单样本手动添加了batch维度(把单个样本包装成batch_size=1的批量输入)。

而你用nn.Flatten(1, -1)时,输入是[context_size, embedding_dim],Flatten(1, -1)会把从第1维到最后一维的所有维度展平,最终得到[context_size*embedding_dim]的1D张量。但PyTorch的Linear层要求输入至少是2D(格式为[batch_size, 特征数]),此时Linear会错误地把1D张量当成[context_size*embedding_dim, 1]来处理,自然和你设置的in_features=context_size*embedding_dim不匹配,导致维度错误。

修正方案

方案一:手动添加batch维度后再展平

在Embedding之后增加nn.Unsqueeze(0)来补全batch维度,再用Flatten展平后续维度,和原模型逻辑完全对齐:

nn.Sequential(
    nn.Embedding(vocab_size, embedding_dim, device=device),
    nn.Unsqueeze(0),  # 将[context_size, embedding_dim]转为[1, context_size, embedding_dim]
    nn.Flatten(1, -1),  # 从第1维展平,得到[1, context_size*embedding_dim]
    nn.Linear(context_size * embedding_dim, 128, device=device),
    nn.ReLU(),
    nn.Linear(128, vocab_size, device=device),
    nn.LogSoftmax(dim=1)
)

方案二:调整输入的维度

如果你能确保输入inputs本身带有batch维度(比如把原来的[context_size]改成[1, context_size]),那么Embedding会输出[1, context_size, embedding_dim],此时你的原代码就可以正常运行——nn.Flatten(1, -1)会直接把[1, context_size, embedding_dim]展平为[1, context_size*embedding_dim],符合Linear层的输入要求。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.23 13:35:08