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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.28 01:03:24