如何用nn.Embedding与torch.sum等价替换sum模式的nn.EmbeddingBag
sum模式EmbeddingBag的等价替换实现
核心前提
PyTorch官方定义中,mode="sum"的nn.EmbeddingBag的计算逻辑,本质就是对输入索引做嵌入查找后按样本维度求和,完全可以用nn.Embedding + torch.sum实现等价替换,两种实现的权重、输出、梯度计算逻辑完全一致。
1. 基础定义替换
原有EmbeddingBag定义:
import torch import torch.nn as nn # n为词表大小,m为嵌入维度 EE = nn.EmbeddingBag(n, m, mode="sum", sparse=True)
等价Embedding定义,可直接复用原有权重:
EE_eq = nn.Embedding(n, m, sparse=True) # 直接迁移原有权重,保证参数完全一致 EE_eq.weight = EE.weight
如果原EmbeddingBag设置了padding_idx、max_norm等参数,替换时需要在nn.Embedding中传入完全相同的参数,保证逻辑对齐
2. 前向计算替换
分两种常见输入场景:
场景1:定长输入(每个样本对应固定数量的索引,无需传offsets)
原EmbeddingBag调用方式:
# input_indices shape为 [batch_size, 每个样本的索引数k] input_indices = torch.randint(0, n, (32, 5)) output_ori = EE(input_indices)
等价实现:
# 先做嵌入查找,输出shape [batch_size, k, m] embeds = EE_eq(input_indices) # 沿样本内索引维度求和,输出shape [batch_size, m],与原输出完全一致 output_eq = torch.sum(embeds, dim=1)
场景2:变长输入(每个样本索引数量不固定,需传offsets参数)
原EmbeddingBag调用方式:
# 所有样本的索引拼接为1维张量 input_indices = torch.tensor([1, 2, 4, 5, 6, 7, 8, 9, 10]) # offsets标记每个样本的起始索引位置 offsets = torch.tensor([0, 3, 5]) output_ori = EE(input_indices, offsets=offsets)
等价实现:
batch_size = len(offsets) # 构造每个索引对应的样本id sample_ids = torch.zeros_like(input_indices) for i in range(1, batch_size): sample_ids[offsets[i-1]:offsets[i]] = i - 1 sample_ids[offsets[-1]:] = batch_size - 1 # 嵌入查找后按样本id求和 embeds = EE_eq(input_indices) output_eq = torch.zeros(batch_size, m, device=embeds.device) output_eq.index_add_(dim=0, index=sample_ids, source=embeds)
等价验证
可通过torch.allclose(output_ori, output_eq)验证两种实现的输出完全一致,参数更新逻辑也完全对齐。
内容的提问来源于stack exchange,提问作者user2584234
相关产品推荐
相关产品推荐

