You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何为BERT/Transformer模型添加手动类别型特征扩展引用分类

为BERT模型添加类别型额外特征的实现方案

针对引用分类任务中引入类别型手动特征的需求,以下是几种可落地的技术方案,均基于HuggingFace Transformers库实现:

一、嵌入拼接法(最通用)

这是最简单直接的方案,核心是将类别特征的嵌入向量与BERT的<[BOS_never_used_51bce0c785ca2f68081bfa7d91973934]> token表征拼接,再送入分类头。

  1. 预处理类别特征:
    • 对低基数类别(如<10种):用独热编码转换为固定长度向量;
    • 对高基数类别(如>10种):用可训练的嵌入层将类别ID映射为稠密向量。
  2. 模型改造步骤:
    • 获取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

二、注意力融合法(增强特征关联)

如果希望模型自动学习文本与类别特征的关联权重,可采用注意力机制融合两者:

  1. 实现思路:
    • 将类别嵌入作为"查询向量",对BERT的序列输出做注意力加权,得到融合了类别信息的文本表征;
    • 或者直接在输入序列末尾添加类别token,让BERT在训练时学习该token与文本的交互。
  2. 代码示例(注意力加权版):
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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.20 20:09:30