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

如何用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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.02 15:57:00