基于HuggingFace搭建BERT-OPT编码解码模型的技术求助
问题描述
我想要搭建一个Encoder-Decoder模型,结构如下:
- 使用Bert-base-uncased作为输入编码器
- 以BERT的CLS token输出为输入,通过线性层连接编码器与解码器
- 使用OPT-125M作为解码器,输入为线性层的输出
这么做是为了复现《In-Context Autoencoder》论文的思路并自行测试。我选择用HuggingFace结合PyTorch实现,因为能大幅减少开发量,而且我不了解OPT-125M或BERT的原生实现,同时HuggingFace的优化方便在普通台式机上运行。
现在遇到的问题是:OPT-125M模型似乎必须用tokenizer处理输入,我绕不开这一步。想问有没有办法直接把线性层的输出输入OPT-125M,或者有没有性能相当的替代方案?
以下是我写的框架代码,目前因为OPT输入格式错误报错:
from transformers import BertTokenizer, BertModel, AutoModelForCausalLM tokenizer = BertTokenizer.from_pretrained('bert-base-uncased') model = BertModel.from_pretrained('bert-base-uncased') OPT = AutoModelForCausalLM.from_pretrained("facebook/opt-125m") import torch from torch import nn class Encoder(nn.Module): def __init__(self): super(Encoder, self).__init__() self.tokenizer = BertTokenizer.from_pretrained('bert-base-uncased') self.model = BertModel.from_pretrained('bert-base-uncased') def forward(self, input_text): inputs = self.tokenizer(input_text, return_tensors="pt", padding=True, truncation=True, max_length=512) outputs = self.model(**inputs) return outputs.last_hidden_state[:, 0, :] # CLS token embeddings class LinearTransformation(nn.Module): def __init__(self, input_dim, output_dim): super(LinearTransformation, self).__init__() self.linear = nn.Linear(input_dim, output_dim) def forward(self, x): return self.linear(x) class Decoder(nn.Module): def __init__(self): super(Decoder, self).__init__() self.model = AutoModelForCausalLM.from_pretrained("facebook/opt-125m") def forward(self, x): # Assuming x is prepared correctly for the OPT model output = self.model(input_ids=x) return output class BertOptPipeline(nn.Module): def __init__(self): super(BertOptPipeline, self).__init__() self.encoder = Encoder() self.linear_transformation = LinearTransformation(768, 512) self.decoder = Decoder() def forward(self, input_text): encoded = self.encoder(input_text) transformed = self.linear_transformation(encoded) print(transformed.shape) # Further processing may be needed here to match the decoder's input requirements decoded = self.decoder(transformed) return decoded pipeline = BertOptPipeline() input_text = "thank you for your help" output = pipeline(input_text)
解决方案
你不需要用tokenizer处理输入给OPT,HuggingFace的AutoModelForCausalLM(包括OPT)支持直接传入隐藏层嵌入向量作为模型输入,而非必须传input_ids。
核心修改点
OPT的forward方法可以接受inputs_embeds参数,这个参数就是模型输入层的嵌入向量,正好对应你线性层的输出结果。只需修改Decoder的forward方法,把传入的线性层输出作为inputs_embeds传给OPT模型,而不是当作input_ids。
额外注意几个细节:
- OPT-125M的隐藏层维度为512,你的线性层已经把BERT的768维CLS输出转成512维,维度完全匹配。
- 因为是自回归模型,OPT默认会处理后续token生成,但如果只是获取初始输入对应的输出,直接传入
inputs_embeds并提取logits即可。 - 建议将tokenizer从Encoder类中移出,避免重复初始化,提升运行效率。
修改后的完整代码
from transformers import BertTokenizer, BertModel, AutoModelForCausalLM import torch from torch import nn # 全局初始化tokenizer,避免重复加载 bert_tokenizer = BertTokenizer.from_pretrained('bert-base-uncased') class Encoder(nn.Module): def __init__(self): super(Encoder, self).__init__() self.model = BertModel.from_pretrained('bert-base-uncased') def forward(self, input_text): inputs = bert_tokenizer(input_text, return_tensors="pt", padding=True, truncation=True, max_length=512) outputs = self.model(**inputs) return outputs.last_hidden_state[:, 0, :] # CLS token embeddings (shape: [batch_size, 768]) class LinearTransformation(nn.Module): def __init__(self, input_dim=768, output_dim=512): super(LinearTransformation, self).__init__() self.linear = nn.Linear(input_dim, output_dim) def forward(self, x): return self.linear(x) # shape: [batch_size, 512] class Decoder(nn.Module): def __init__(self): super(Decoder, self).__init__() self.model = AutoModelForCausalLM.from_pretrained("facebook/opt-125m") def forward(self, inputs_embeds): # 直接传入inputs_embeds替代input_ids,添加attention_mask避免无意义的padding影响 batch_size = inputs_embeds.shape[0] attention_mask = torch.ones(batch_size, inputs_embeds.shape[1], device=inputs_embeds.device) outputs = self.model(inputs_embeds=inputs_embeds, attention_mask=attention_mask) return outputs.logits class BertOptPipeline(nn.Module): def __init__(self): super(BertOptPipeline, self).__init__() self.encoder = Encoder() self.linear_transformation = LinearTransformation() self.decoder = Decoder() def forward(self, input_text): encoded = self.encoder(input_text) transformed = self.linear_transformation(encoded) # 为匹配OPT输入格式,增加序列长度维度(从[batch,512]变为[batch,1,512]) transformed = transformed.unsqueeze(1) decoded_logits = self.decoder(transformed) return decoded_logits # 测试运行 pipeline = BertOptPipeline() input_text = "thank you for your help" output = pipeline(input_text) print("Output logits shape:", output.shape) # 输出应为 [1, 1, vocab_size]
自回归生成扩展
如果需要让OPT从CLS编码生成完整句子,可以在Decoder中新增generate方法,同样传入inputs_embeds:
def generate(self, inputs_embeds, max_length=50): batch_size = inputs_embeds.shape[0] attention_mask = torch.ones(batch_size, inputs_embeds.shape[1], device=inputs_embeds.device) outputs = self.model.generate( inputs_embeds=inputs_embeds, attention_mask=attention_mask, max_length=max_length, do_sample=False ) return outputs
生成的token序列可以用OPT的tokenizer解码成文本。
内容的提问来源于stack exchange,提问作者Florian Jäger
相关产品推荐
相关产品推荐

