如何在排除PAD token的前提下高效均值池化BERT嵌入?
高效实现排除PAD token的BERT句嵌入均值池化
当然有高效实现方式,核心是利用attention_mask(Tokenizer返回的掩码张量)批量处理,全程无需遍历单个序列,完全通过张量操作完成:
用掩码标记有效token
Tokenizer做padding时会返回attention_mask,形状为(bs, max_seq_len)——非PAD的token位置值为1,PAD位置为0。我们可以把这个掩码扩展到和最后隐藏状态相同的维度,以此过滤PAD对应的嵌入。批量计算有效隐藏状态总和
将扩展后的掩码与最后隐藏状态逐元素相乘,PAD对应的嵌入会被置为0;之后对max_seq_len维度求和,得到每个样本有效token的隐藏状态总和。计算有效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
相关产品推荐
相关产品推荐

