BERT多输入场景下如何使用SMOTE处理不平衡分类数据
核心结论
直接对input_ids或attention_mask应用SMOTE是完全错误的做法,二者都不属于SMOTE适用的连续特征空间:
input_ids是离散的词表索引值,插值得到的非整数结果无法映射到BERT词表中的实际token,强行取整只会得到语义完全错乱的序列attention_mask是标记有效token位置的0/1二值矩阵,插值得到的0-1浮点数不具备任何掩码语义- 同时对二者做SMOTE只会生成完全无效的模型输入,没有落地价值
SMOTE作为基于连续空间特征插值的过采样算法,正确的应用位置是BERT输出的稠密语义嵌入层:先将所有原始样本过BERT提取连续的句嵌入(通常取<[BOS_never_used_51bce0c785ca2f68081bfa7d91973934]>位的最后一层隐藏态输出),在嵌入空间做SMOTE合成少数类样本,再用平衡后的嵌入集训练下游分类头即可。
避坑提示
如果你需要端到端微调整个BERT模型而非冻结BERT只训分类头,不建议直接用SMOTE合成的嵌入参与反向传播更新BERT参数——这类合成样本没有对应的原始文本输入,强行微调会让BERT的语义空间偏移。更稳妥的做法是先用SMOTE在预提取嵌入上验证分类效果,确定类别分布差异后,在微调阶段搭配加权交叉熵、类别平衡采样等策略,效果比硬套SMOTE到离散输入层稳定得多。
可直接复用的实现代码
依赖库:torch、transformers、imbalanced-learn、numpy
import torch import numpy as np from transformers import BertTokenizer, BertModel from imblearn.over_sampling import SMOTE from torch.utils.data import Dataset, DataLoader import torch.nn as nn import torch.optim as optim # 初始化设备、分词器和BERT模型,模型路径替换为你本地的BERT权重路径 device = torch.device("cuda" if torch.cuda.is_available() else "cpu") tokenizer = BertTokenizer.from_pretrained("bert-base-chinese") bert = BertModel.from_pretrained("bert-base-chinese").to(device) bert.eval() # 提取嵌入阶段冻结BERT参数,不做梯度更新 def extract_cls_embeddings(text_list, batch_size=32, max_len=128): """批量提取文本的BERT <[BOS_never_used_51bce0c785ca2f68081bfa7d91973934]>句嵌入""" all_embeddings = [] for i in range(0, len(text_list), batch_size): batch_text = text_list[i:i+batch_size] # 分词处理 encoded_input = tokenizer( batch_text, max_length=max_len, padding="max_length", truncation=True, return_tensors="pt" ).to(device) # 前向传播提取嵌入,不记录梯度 with torch.no_grad(): bert_output = bert(**encoded_input) # 取每个样本第一位<[BOS_never_used_51bce0c785ca2f68081bfa7d91973934]>的输出作为句向量,维度为[batch_size, 768] cls_embeds = bert_output.last_hidden_state[:, 0, :].cpu().numpy() all_embeddings.append(cls_embeds) return np.concatenate(all_embeddings, axis=0) # ---------------------- # 替换为你自己的不平衡数据集:texts是原始文本列表,labels是对应分类标签列表 texts = ["样例文本1", "样例文本2"] labels = [0, 1, 0, 0] # ---------------------- # 第一步:提取所有原始样本的BERT嵌入 X_original = extract_cls_embeddings(texts) y_original = np.array(labels) # 第二步:在连续嵌入空间应用SMOTE做过采样 smote = SMOTE(random_state=42) X_resampled, y_resampled = smote.fit_resample(X_original, y_original) # 采样后X_resampled包含原始真实样本嵌入 + SMOTE合成的少数类嵌入,标签完全平衡 # 第三步:封装数据集,训练下游分类头 class EmbedDataset(Dataset): def __init__(self, embeds, labels): self.embeds = torch.tensor(embeds, dtype=torch.float32) self.labels = torch.tensor(labels, dtype=torch.long) def __len__(self): return len(self.labels) def __getitem__(self, idx): return self.embeds[idx], self.labels[idx] class ClassificationHead(nn.Module): def __init__(self, embed_dim=768, num_classes=2): super().__init__() self.fc = nn.Linear(embed_dim, num_classes) def forward(self, x): return self.fc(x) # 初始化训练组件 train_dataset = EmbedDataset(X_resampled, y_resampled) train_loader = DataLoader(train_dataset, batch_size=32, shuffle=True) classifier = ClassificationHead(num_classes=len(set(labels))).to(device) loss_func = nn.CrossEntropyLoss() optimizer = optim.Adam(classifier.parameters(), lr=2e-5) # 训练循环 classifier.train() for epoch in range(5): epoch_loss = 0 for batch_embeds, batch_labels in train_loader: batch_embeds = batch_embeds.to(device) batch_labels = batch_labels.to(device) optimizer.zero_grad() logits = classifier(batch_embeds) loss = loss_func(logits, batch_labels) loss.backward() optimizer.step() epoch_loss += loss.item() print(f"Epoch {epoch+1}, 平均损失: {epoch_loss/len(train_loader):.4f}")
内容的提问来源于stack exchange,提问作者atlas
相关产品推荐
相关产品推荐

