基于Llama2实现回归任务的输出维度异常问题求助
解决Llama2适配回归任务的输出维度问题
你的问题核心是:当前模型使用了所有token的最后隐藏层状态输入到线性层,导致输出维度为(batch_size, seq_len, 1)(比如输入"Hello world!"分词为4个token,输出就是1,4,1),但回归任务需要的是单值输出(batch_size, 1)。要解决这个问题,需要把序列维度的隐藏态压缩成一个向量,以下是几种可行方案:
方案1:使用BOS token的隐藏态
Llama分词器会自动在输入开头添加<s>(BOS,序列起始)token,对应input_ids的第0个索引位置。直接提取这个token的隐藏态作为线性层的输入即可:
import torch import torch.nn as nn from transformers import LlamaModel, LlamaTokenizer class TransformerModel(nn.Module): def __init__(self, model_name:str, additional_layer_size:int = 1): super(TransformerModel, self).__init__() self.transformer = LlamaModel.from_pretrained(model_name, torch_dtype=torch.float32, cache_dir="hugginface_cache/models") self.tokenizer = LlamaTokenizer.from_pretrained(model_name, cache_dir="hugginface_cache/tokenizer") # 修复Llama默认无pad token的问题 self.tokenizer.pad_token = self.tokenizer.eos_token self.additional_layer = nn.Linear(self.transformer.config.hidden_size, additional_layer_size) def forward(self, input_text): # 加入padding和truncation,支持批量输入 inputs = self.tokenizer( input_text, return_tensors="pt", padding=True, truncation=True, max_length=512 # 根据任务需求设置最大序列长度 ).to("cuda") outputs = self.transformer(**inputs) # 提取BOS token(第一个token)的隐藏态 bos_hidden = outputs.last_hidden_state[:, 0, :] # shape: (batch_size, hidden_size) # 输出单值回归结果 return self.additional_layer(bos_hidden)
方案2:使用EOS token的隐藏态
如果希望聚焦输入序列的末尾语义,可以通过attention_mask定位真实的序列末尾(EOS,序列结束)token:
def forward(self, input_text): inputs = self.tokenizer( input_text, return_tensors="pt", padding=True, truncation=True, max_length=512 ).to("cuda") last_hidden_state = self.transformer(**inputs).last_hidden_state attention_mask = inputs.attention_mask # 计算每个序列的最后一个有效token索引 seq_lengths = attention_mask.sum(dim=1) - 1 # 索引从0开始,所以减1 # 提取对应位置的隐藏态 eos_hidden = last_hidden_state[torch.arange(last_hidden_state.size(0)), seq_lengths, :] return self.additional_layer(eos_hidden)
方案3:均值池化(利用全序列信息)
如果想融合整个输入序列的语义信息,可以对所有非padding的token隐藏态做均值池化:
def forward(self, input_text): inputs = self.tokenizer( input_text, return_tensors="pt", padding=True, truncation=True, max_length=512 ).to("cuda") last_hidden_state = self.transformer(**inputs).last_hidden_state attention_mask = inputs.attention_mask # 用mask屏蔽padding token的影响 mask = attention_mask.unsqueeze(-1).expand(last_hidden_state.size()) sum_hidden = torch.sum(last_hidden_state * mask, dim=1) # 计算均值,避免除以0 mean_hidden = sum_hidden / torch.clamp(mask.sum(dim=1), min=1e-9) return self.additional_layer(mean_hidden)
额外注意事项
- Llama默认没有pad token,必须手动设置
self.tokenizer.pad_token = self.tokenizer.eos_token,否则批量输入会报错。 - 回归任务的损失函数建议使用
nn.MSELoss(),与单值输出格式匹配。 - 三种方案各有优劣:BOS方案计算最快,均值池化能利用全序列信息,EOS方案适合聚焦序列末尾语义,可根据任务特性选择。
内容的提问来源于stack exchange,提问作者Lukas Fehring
相关产品推荐
相关产品推荐

