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

能否以可微分方式将独热向量转换为torch.nn.Embedding输出?

问题解答

完全可以实现,且该过程天然具备可微分性,核心是用矩阵乘法替代Embedding层的索引操作,两者的输出结果完全一致。

原理说明

PyTorch的torch.nn.Embedding本质是一个形状为[vocab_size, embedding_dim]的参数矩阵。当输入整数令牌(形状[batch_size, seq_len])时,它通过索引取值获取对应行的嵌入向量;而独热向量(形状[batch_size, seq_len, vocab_size])可视为编码后的矩阵,将其与Embedding参数矩阵做矩阵乘法,就能得到和索引操作完全相同的结果,且矩阵乘法是可微分操作,反向传播时能正常更新Embedding的参数。

代码示例

以下代码验证两种输入方式的等价性:

import torch
import torch.nn as nn

# 定义基础参数
vocab_size = 100
embedding_dim = 128
batch_size = 32
seq_len = 10

# 初始化Embedding层
embedding = nn.Embedding(vocab_size, embedding_dim)

# 方式1:输入整数令牌
token_ids = torch.randint(0, vocab_size, (batch_size, seq_len))
embedding_from_ids = embedding(token_ids)

# 方式2:输入独热向量
one_hot_tokens = torch.nn.functional.one_hot(token_ids, num_classes=vocab_size).float()
# 矩阵乘法:[batch, seq, vocab] @ [vocab, embed] → [batch, seq, embed]
embedding_from_onehot = torch.matmul(one_hot_tokens, embedding.weight)

# 验证结果一致
print(torch.allclose(embedding_from_ids, embedding_from_onehot))  # 输出True

注意事项

  • 独热向量必须转换为浮点型,否则无法与Embedding的浮点型参数矩阵进行运算。
  • 该方法的计算效率低于直接输入整数令牌(尤其是词汇量较大时),但完全满足可微分需求。
  • 若需封装成模块,可自定义nn.Module复用原Embedding的权重参数,专门处理独热输入。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.25 06:09:23