如何使用Roberta模型计算词嵌入与句嵌入?
RoBERTa词嵌入与句嵌入计算方法
一、词嵌入的正确性确认
你从outputs[0]提取最后一层隐藏状态的方式是正确的。RoBertaModel的输出中,outputs[0]的形状为(batch_size, sequence_length, hidden_size),其中每个位置的向量对应输入序列中对应token的词嵌入,包括RoBERTa自动添加的特殊token<s>(句首)和</s>(句尾)。
注意:RoBERTa的tokenizer会对原文本做子词切分,比如"yellow"可能被切分为多个子词,每个子词对应一个嵌入向量,如果你需要整词的嵌入,可能需要对子词嵌入做进一步聚合(比如平均)。
另外,你手动做pad的方式可以优化,直接用tokenizer的批量编码方法更规范,还能生成attention_mask,后续计算句嵌入时可以过滤pad部分的影响:
from transformers import RobertaModel, RobertaTokenizer import torch model = RobertaModel.from_pretrained('roberta-base') tokenizer = RobertaTokenizer.from_pretrained('roberta-base') captions = ["example caption", "lorem ipsum", "this bird is yellow has red wings", "hi", "example"] # 用tokenizer批量编码,自动pad并生成attention_mask encoding = tokenizer(captions, padding=True, truncation=True, return_tensors='pt') input_ids = encoding['input_ids'] attention_mask = encoding['attention_mask'] outputs = model(input_ids, attention_mask=attention_mask) word_embeddings = outputs[0].contiguous() # 形状: (5, max_seq_len, 768)
二、句嵌入的常用计算方式
RoBERTa本身没有直接输出句嵌入,常用的三种计算方式如下:
1. 用句首特殊token <s> 的嵌入
RoBERTa预训练时,<s>(对应输入序列的第一个token)被用作句子级任务的聚合标记,直接取它的隐藏状态作为句嵌入是最常用的方式:
# 取每个序列的第一个token的嵌入作为句嵌入 sentence_embeddings_cls = outputs[0][:, 0, :] # 形状: (5, 768)
2. 所有token嵌入的均值(带attention_mask过滤pad)
对有效token(非pad部分)的嵌入做平均,避免pad的0向量影响结果:
# 扩展attention_mask维度,和词嵌入形状匹配 mask = attention_mask.unsqueeze(-1).expand(word_embeddings.size()) # 只保留有效token的嵌入 masked_embeddings = word_embeddings * mask # 计算有效token的数量 sum_mask = mask.sum(1) sum_mask = torch.clamp(sum_mask, min=1e-9) # 避免除以0 # 计算均值 sentence_embeddings_avg = masked_embeddings.sum(1) / sum_mask # 形状: (5, 768)
3. 所有token嵌入的最大池化(带attention_mask过滤pad)
对有效token的嵌入做最大池化,提取最具代表性的特征:
# 将pad部分的嵌入设为极小值,避免影响max计算 masked_embeddings = word_embeddings.masked_fill(~mask.bool(), -1e9) # 对序列维度做max sentence_embeddings_max = masked_embeddings.max(1)[0] # 形状: (5, 768)
总结
- 词嵌入直接取
outputs[0]即可,注意子词和特殊token的对应关系; - 句嵌入优先尝试
<s>的嵌入,若效果不佳可尝试均值或最大池化,记得结合attention_mask过滤pad部分。
内容的提问来源于stack exchange,提问作者user23232264
相关产品推荐
相关产品推荐

