如何为BERT/Transformer模型添加手动类别型特征扩展引用分类
为BERT模型添加类别型额外特征的实现方案
针对引用分类任务中引入类别型手动特征的需求,以下是几种可落地的技术方案,均基于HuggingFace Transformers库实现:
一、嵌入拼接法(最通用)
这是最简单直接的方案,核心是将类别特征的嵌入向量与BERT的<[BOS_never_used_51bce0c785ca2f68081bfa7d91973934]> token表征拼接,再送入分类头。
- 预处理类别特征:
- 对低基数类别(如<10种):用独热编码转换为固定长度向量;
- 对高基数类别(如>10种):用可训练的嵌入层将类别ID映射为稠密向量。
- 模型改造步骤:
- 获取BERT输出的<[BOS_never_used_51bce0c785ca2f68081bfa7d91973934]> token表征(维度通常为768);
- 将类别特征向量与该表征拼接;
- 把拼接后的向量送入全连接分类层,输出4类引用的概率。
代码示例:
from transformers import BertModel, BertTokenizer import torch.nn as nn import torch class BertWithCatFeatures(nn.Module): def __init__(self, num_cat_classes, cat_emb_dim=32): super().__init__() self.bert = BertModel.from_pretrained('bert-base-uncased') # 类别特征嵌入层(若为独热编码可跳过此层,直接拼接) self.cat_emb = nn.Embedding(num_cat_classes, cat_emb_dim) # 分类头:BERT维度 + 类别嵌入维度 self.classifier = nn.Linear(768 + cat_emb_dim, 4) def forward(self, input_ids, attention_mask, cat_ids): # 获取BERT的<[BOS_never_used_51bce0c785ca2f68081bfa7d91973934]>表征 bert_output = self.bert(input_ids=input_ids, attention_mask=attention_mask) cls_hidden = bert_output.last_hidden_state[:, 0, :] # 类别特征嵌入 cat_emb = self.cat_emb(cat_ids) # 拼接特征 combined_hidden = torch.cat([cls_hidden, cat_emb], dim=1) # 分类输出 logits = self.classifier(combined_hidden) return logits
二、注意力融合法(增强特征关联)
如果希望模型自动学习文本与类别特征的关联权重,可采用注意力机制融合两者:
- 实现思路:
- 将类别嵌入作为"查询向量",对BERT的序列输出做注意力加权,得到融合了类别信息的文本表征;
- 或者直接在输入序列末尾添加类别token,让BERT在训练时学习该token与文本的交互。
- 代码示例(注意力加权版):
from transformers import BertModel import torch.nn as nn import torch class BertWithCatAttention(nn.Module): def __init__(self, num_cat_classes, cat_emb_dim=32): super().__init__() self.bert = BertModel.from_pretrained('bert-base-uncased') self.cat_emb = nn.Embedding(num_cat_classes, cat_emb_dim) # 注意力层:用类别嵌入查询文本序列 self.attention = nn.MultiheadAttention(embed_dim=768, num_heads=1, batch_first=True) self.classifier = nn.Linear(768, 4) self.cat_proj = nn.Linear(cat_emb_dim, 768) def forward(self, input_ids, attention_mask, cat_ids): bert_output = self.bert(input_ids=input_ids, attention_mask=attention_mask) seq_hidden = bert_output.last_hidden_state # [batch_size, seq_len, 768] cat_emb = self.cat_emb(cat_ids).unsqueeze(1) # [batch_size, 1, 32] # 将类别嵌入映射到BERT维度,作为查询 query = self.cat_proj(cat_emb) # 注意力加权文本序列 attn_output, _ = self.attention(query, seq_hidden, seq_hidden, key_padding_mask=~attention_mask.bool()) # 取加权后的聚合表征做分类 logits = self.classifier(attn_output[:, 0, :]) return logits
三、多分支并行法(分离特征处理)
若想保留文本与类别特征的独立性,可采用双分支结构:文本走BERT分支,类别特征走单独的全连接分支,最后融合输出。
代码示例:
from transformers import BertModel import torch.nn as nn import torch class BertWithDualBranch(nn.Module): def __init__(self, num_cat_classes, cat_hidden_dim=64): super().__init__() self.bert = BertModel.from_pretrained('bert-base-uncased') # 类别特征处理分支 self.cat_branch = nn.Sequential( nn.Embedding(num_cat_classes, cat_hidden_dim), nn.ReLU(), nn.Linear(cat_hidden_dim, 128) ) # 融合后分类头 self.classifier = nn.Linear(768 + 128, 4) def forward(self, input_ids, attention_mask, cat_ids): # BERT文本分支 cls_hidden = self.bert(input_ids=input_ids, attention_mask=attention_mask).last_hidden_state[:, 0, :] # 类别特征分支 cat_hidden = self.cat_branch(cat_ids) # 融合分类 combined_hidden = torch.cat([cls_hidden, cat_hidden], dim=1) logits = self.classifier(combined_hidden) return logits
注意事项
- 若使用预训练BERT,可选择冻结BERT底层参数(仅微调上层和类别嵌入层)或全参数微调,根据数据集大小调整;
- 类别特征的嵌入维度需根据类别数量调整:类别越多,嵌入维度可适当增大;
- 训练时需将类别特征与文本输入一同喂入模型,确保数据加载器能同时返回文本输入和类别ID。
内容的提问来源于stack exchange,提问作者Ryan Mclaren
相关产品推荐
相关产品推荐

