如何拼接类BERT句子表征与词嵌入——基于Keras与HuggingFace实现
这个问题其实在结合预训练Transformer和传统词嵌入的场景里挺常见的——要把BERT输出的子词表征聚合回词级,才能和你的词嵌入层输出对齐拼接。不管是Keras还是PyTorch都能实现,我给你分别讲讲具体的做法:
Keras实现方案
首先你需要在预处理阶段记录每个子词对应的原始词ID,BERT的tokenizer提供了word_ids()方法可以帮你做到这一点。举个预处理的例子:
from transformers import BertTokenizer tokenizer = BertTokenizer.from_pretrained("bert-base-uncased") sample_sentence = "your sample sentence here" encoding = tokenizer.encode_plus( sample_sentence, add_special_tokens=True, max_length=max_transformer_len, padding="max_length", truncation=True, return_offsets_mapping=False, return_word_ids=True # 关键:获取子词对应的原始词ID ) # 把None替换成-100,因为Keras输入不能包含None值 word_ids = [-100 if id is None else id for id in encoding["word_ids"]]
word_ids里的每个元素对应一个子词所属的原始词索引(比如[-100, 0, 0, 1, -100],其中-100对应[CLS]、[SEP]这类特殊token)。
接下来在模型里,你需要把word_ids作为输入层传入,然后用TensorFlow的tf.math.segment_max来对每个原始词对应的子词做最大池化:
import tensorflow as tf from tensorflow.keras import layers from transformers import TFBertModel # 1. 定义模型输入 input_ids = layers.Input(shape=(max_transformer_len,), dtype=tf.int32) token_type_ids = layers.Input(shape=(max_transformer_len,), dtype=tf.int32) attention_mask = layers.Input(shape=(max_transformer_len,), dtype=tf.int32) word_ids_input = layers.Input(shape=(max_transformer_len,), dtype=tf.int32) # 新增的word_ids输入 # 2. 获取BERT子词嵌入 encoder = TFBertModel.from_pretrained("bert-base-uncased") embedding1 = encoder(input_ids, token_type_ids=token_type_ids, attention_mask=attention_mask)[0] # 3. 处理word_ids,把-100替换为max_sentence_len(单独分组处理特殊token) processed_word_ids = tf.where( tf.equal(word_ids_input, -100), tf.constant(max_sentence_len, dtype=tf.int32), word_ids_input ) # 4. 对每个原始词对应的子词做最大池化 pooled_embedding1 = tf.math.segment_max(embedding1, processed_word_ids) # 去掉特殊token对应的池化结果,得到和embedding2长度一致的词级表征 pooled_embedding1 = pooled_embedding1[:, :max_sentence_len, :] # 5. 你的词嵌入层部分(修正原代码的变量名错误) input_wordembedding = layers.Input(shape=(max_sentence_len,), dtype='int32', name='we_input') embedding2 = layers.Embedding( output_dim=wordembedding_VECTOR_SIZE, input_dim=wordembedding_VOCAB_SIZE, input_length=max_sentence_len, weights=[emb_matrix], name='emb1' )(input_wordembedding) # 6. 拼接两个词级表征 z = layers.Concatenate(name='merged')([pooled_embedding1, embedding2]) # 构建完整模型 model = tf.keras.Model( inputs=[input_ids, token_type_ids, attention_mask, word_ids_input, input_wordembedding], outputs=z )
PyTorch实现方案
PyTorch的思路和Keras完全一致,只是API略有不同。同样先在预处理阶段获取word_ids:
from transformers import BertTokenizer tokenizer = BertTokenizer.from_pretrained("bert-base-uncased") sample_sentence = "your sample sentence here" encoding = tokenizer.encode_plus( sample_sentence, add_special_tokens=True, max_length=max_transformer_len, padding="max_length", truncation=True, return_word_ids=True, return_tensors="pt" ) # 把None替换成-100 word_ids = torch.tensor([-100 if id is None else id for id in encoding["word_ids"][0]])
然后在模型中实现池化:
import torch import torch.nn as nn from transformers import BertModel class CustomModel(nn.Module): def __init__(self, wordembedding_VOCAB_SIZE, wordembedding_VECTOR_SIZE, max_sentence_len, emb_matrix): super().__init__() self.encoder = BertModel.from_pretrained("bert-base-uncased") self.embedding = nn.Embedding( num_embeddings=wordembedding_VOCAB_SIZE, embedding_dim=wordembedding_VECTOR_SIZE, _weight=torch.tensor(emb_matrix, dtype=torch.float32) ) self.max_sentence_len = max_sentence_len def forward(self, input_ids, token_type_ids, attention_mask, word_ids, input_wordembedding): # 获取BERT子词嵌入 outputs = self.encoder(input_ids, token_type_ids=token_type_ids, attention_mask=attention_mask) embedding1 = outputs.last_hidden_state # shape: (batch_size, max_transformer_len, hidden_size) # 处理word_ids:把-100替换为max_sentence_len processed_word_ids = torch.where( word_ids == -100, torch.tensor(self.max_sentence_len, device=word_ids.device), word_ids ) # 用scatter_reduce实现按segment的最大池化 batch_size, hidden_size = embedding1.size(0), embedding1.size(-1) pooled = torch.zeros( batch_size, self.max_sentence_len + 1, hidden_size, device=embedding1.device, dtype=embedding1.dtype ) pooled = pooled.scatter_reduce( dim=1, index=processed_word_ids.unsqueeze(-1).repeat(1,1,hidden_size), src=embedding1, reduce="max", include_self=False ) # 去掉特殊token对应的部分,得到词级表征 pooled_embedding1 = pooled[:, :self.max_sentence_len, :] # 获取词嵌入 embedding2 = self.embedding(input_wordembedding) # 拼接 z = torch.cat([pooled_embedding1, embedding2], dim=-1) return z
最后补充一点:预处理时要确保word_ids的长度和input_ids一致,并且每个样本的word_ids中,原始词的索引是连续的0到max_sentence_len-1,这样池化后的结果才能和你的词嵌入层输出的长度完美对齐。
内容的提问来源于stack exchange,提问作者Toni.M
相关产品推荐
相关产品推荐

