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

如何在排除PAD token的前提下高效均值池化BERT嵌入?

高效实现排除PAD token的BERT句嵌入均值池化

当然有高效实现方式,核心是利用attention_mask(Tokenizer返回的掩码张量)批量处理,全程无需遍历单个序列,完全通过张量操作完成:

  1. 用掩码标记有效token
    Tokenizer做padding时会返回attention_mask,形状为(bs, max_seq_len)——非PAD的token位置值为1,PAD位置为0。我们可以把这个掩码扩展到和最后隐藏状态相同的维度,以此过滤PAD对应的嵌入。

  2. 批量计算有效隐藏状态总和
    将扩展后的掩码与最后隐藏状态逐元素相乘,PAD对应的嵌入会被置为0;之后对max_seq_len维度求和,得到每个样本有效token的隐藏状态总和。

  3. 计算有效token数量并求均值
    对掩码在max_seq_len维度求和,得到每个样本的有效token数(注意加个极小值避免除以0的情况);再用总和除以有效数量,就能得到排除PAD的均值池化句嵌入。

代码示例(PyTorch)

import torch
from transformers import BertTokenizer, BertModel

# 假设已完成模型推理,得到以下变量
hidden_states = model_output.last_hidden_state  # 形状: (bs, max_seq_len, hidden_sz)
attention_mask = tokenizer_result["attention_mask"]  # 形状: (bs, max_seq_len)

# 将掩码扩展为和隐藏状态同维度,实现逐元素过滤
mask = attention_mask.unsqueeze(-1).expand(hidden_states.size())

# 计算有效隐藏状态的总和
sum_valid_hidden = torch.sum(hidden_states * mask, dim=1)

# 计算有效token数量,防止除以0
num_valid_tokens = torch.clamp(mask.sum(dim=1), min=1e-8)

# 得到最终的均值池化句嵌入
sentence_embeddings = sum_valid_hidden / num_valid_tokens

这种方式依赖张量的批量运算,能充分利用框架的硬件优化(比如CUDA加速),比遍历单个序列的效率高得多,尤其适合大批次处理场景。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.24 19:05:04