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

PyTorch中如何使Embedding层将2*50输入输出为2*3维度结果

PyTorch Embedding层实现指定维度输出方案

原代码核心问题

  • nn.Embedding第一个参数传错:该参数为词表大小(即输入token的最大取值+1),不是输入序列长度。示例中X_train由randint(10)生成,token取值范围为0-9,词表大小应设为10。
  • 输入张量维度处理错误:转换张量时额外套了一层方括号,把原本(2,50)的输入变成了(1,2,50),输出维度不符合预期。
  • 缺少序列聚合步骤:Embedding层会为每个输入token单独生成嵌入向量,输入形状为(样本数, 序列长度)时,原始输出形状为(样本数, 序列长度, 嵌入维度)。要从(2,50)的输入得到(2,3)的输出,必须对长度为50的序列维度做聚合(池化)操作。

修正后可运行代码

import torch
import torch.nn as nn
import numpy as np

# 生成原始训练集
X_train = np.random.randint(10, size=(2, 100))

# 拆分不经过嵌入层的部分,形状(2, 50)
X_train_notembedding = X_train[:, 0:50]
# 拆分需要输入嵌入层的部分,形状(2, 50)
X_train_embedding = X_train[:, 50:100]

# 转换为嵌入层要求的LongTensor类型,保持原有形状不额外增维
X_train_embedding = torch.LongTensor(X_train_embedding)
# 初始化嵌入层:词表大小10,嵌入维度3
embedding = nn.Embedding(num_embeddings=10, embedding_dim=3)
# 得到原始嵌入输出,形状为(2, 50, 3)
embedding_raw = embedding(X_train_embedding)
# 对序列维度做平均池化,压缩序列长度,得到(2, 3)的嵌入输出
embedding_output = embedding_raw.mean(dim=1)

# 将非嵌入部分转换为浮点张量,和嵌入输出类型对齐
X_train_notembedding = torch.FloatTensor(X_train_notembedding)
# 沿特征维度拼接,最终得到形状(2, 53)的特征矩阵
X_train_new = torch.cat([X_train_notembedding, embedding_output], dim=1)

# 维度验证
print(embedding_output.shape)  # 输出: torch.Size([2, 3])
print(X_train_new.shape)       # 输出: torch.Size([2, 53])

关键注意事项

  • 如果你的输入token取值范围不是0-9,比如token最大id是99,就把num_embeddings设为100即可。
  • 序列聚合方式不局限于平均池化,可根据任务需求替换为最大池化、注意力加权聚合等,只要把长度为50的序列维度压缩即可得到每个样本对应的3维嵌入向量。
  • 拼接前要保证两部分张量的dtype一致,非嵌入部分如果是整数类型需要转成浮点型,否则会报类型错误。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.03 06:12:39