LLAMA 2词嵌入3维转2维后值重复问题求助
解决Llama2(AutoModelForCausalLM)生成的3维张量转2维时的值重复问题
问题根源
你之前的代码存在两个核心问题:
- 错误取了
outputs[0]:AutoModelForCausalLM的outputs[0]是所有token的预测logits,形状为[batch_size, seq_len, vocab_size],这不是句子级嵌入,而是每个token的下一个token预测概率的log值。 - 盲目取第一个token的logits:Llama2输入的第一个token默认是
<s>(bos token),不同样本的bos token对应的logits本身就高度相似,直接取x[0]自然会出现值重复的情况。
正确解决方案:获取句子级嵌入
CausalLM没有专门的last_hidden_state属性,但当你设置output_hidden_states=True时,outputs.hidden_states会返回所有层的隐藏状态(从embedding层到最后一层),我们可以基于最后一层的隐藏状态,通过池化得到2维的句子嵌入([batch_size, hidden_size]),适配逻辑回归、SVM等模型的输入要求。
方法1:取最后一个有效token的隐藏状态(推荐)
Llama2是因果语言模型,最后一个有效token的隐藏状态最能代表整句语义:
import torch with torch.no_grad(): outputs = model( features['input_ids'].to(device), features['attention_mask'].to(device), output_hidden_states=True ) # 获取最后一层的隐藏状态(形状:[batch_size, seq_len, hidden_size]) last_hidden = outputs.hidden_states[-1] attention_mask = features['attention_mask'].to(device) # 计算每个样本最后一个有效token的索引 seq_lengths = torch.sum(attention_mask, dim=1) - 1 # 有效长度减1得到最后一个非pad token的位置 # 扩展维度以匹配hidden state的形状,方便索引 seq_lengths = seq_lengths.unsqueeze(-1).unsqueeze(-1).expand(-1, -1, last_hidden.size(-1)) # 取出最后一个有效token的隐藏状态,压缩为2维 cls_train = torch.gather(last_hidden, dim=1, index=seq_lengths).squeeze(1) # 形状:[batch_size, hidden_size] aux = cls_train.to("cpu") Y = database['label']
方法2:均值池化(对所有非pad token取平均)
如果句子较长,均值池化也能稳定提取全局语义:
import torch with torch.no_grad(): outputs = model( features['input_ids'].to(device), features['attention_mask'].to(device), output_hidden_states=True ) last_hidden = outputs.hidden_states[-1] attention_mask = features['attention_mask'].unsqueeze(-1).to(device) # 扩展维度为[batch_size, seq_len, 1] # 用attention_mask过滤pad token的隐藏状态 masked_hidden = last_hidden * attention_mask # 对每个样本的有效token隐藏状态取平均 cls_train = torch.sum(masked_hidden, dim=1) / torch.sum(attention_mask, dim=1) # 形状:[batch_size, hidden_size] aux = cls_train.to("cpu") Y = database['label']
关于pad_token的说明
你添加tokenizer.add_special_tokens({'pad_token': '[PAD]'})是正确的操作,因为Llama2原生没有pad token,这一步让tokenizer能正确处理文本padding。只要你的attention_mask是正确生成的(pad位置为0,有效token位置为1),就不会影响上述池化逻辑,因为我们会通过attention_mask过滤掉pad token的影响。
内容的提问来源于stack exchange,提问作者lucasa.lisboa
相关产品推荐
相关产品推荐

