如何使用PyTorch MultiheadAttention基于自定义数据集实现分类任务?
基于PyTorch MultiheadAttention实现二分类任务方案
核心思路
MultiheadAttention本身用于建模序列元素间的依赖关系,要完成分类任务,需在注意力模块后搭配分类头:先通过注意力处理输入序列,再通过池化操作提取全局特征,最后送入全连接层输出分类结果。
具体实现步骤
1. 数据封装与加载
将数据集转为PyTorch支持的Dataset和DataLoader格式。注意MultiheadAttention默认输入形状为(seq_len, batch_size, feature_dim),而你的输入是(batch_size, seq_len, feature_dim),后续需在模型中做维度转换。
示例代码:
import torch from torch.utils.data import Dataset, DataLoader class CustomDataset(Dataset): def __init__(self, x, y): self.x = torch.tensor(x, dtype=torch.float32) self.y = torch.tensor(y, dtype=torch.long) def __len__(self): return len(self.x) def __getitem__(self, idx): return self.x[idx], self.y[idx] # 假设x_train、y_train是你的原始数据(如numpy数组) train_dataset = CustomDataset(x_train, y_train) train_loader = DataLoader(train_dataset, batch_size=32, shuffle=True) val_dataset = CustomDataset(x_val, y_val) val_loader = DataLoader(val_dataset, batch_size=32)
2. 构建分类模型
模型包含三个核心部分:维度转换适配注意力模块、MultiheadAttention层、池化+分类头。
示例代码:
class AttentionClassifier(torch.nn.Module): def __init__(self, feature_dim=300, num_heads=6, seq_len=102, num_classes=2): super().__init__() # feature_dim必须能被num_heads整除,300/6=50符合要求 self.multihead_attn = torch.nn.MultiheadAttention(embed_dim=feature_dim, num_heads=num_heads) # 分类头:池化后接全连接层 self.classifier = torch.nn.Sequential( torch.nn.Linear(feature_dim, 128), torch.nn.ReLU(), torch.nn.Linear(128, num_classes) ) def forward(self, x): # 转换维度为MultiheadAttention要求的(seq_len, batch_size, feature_dim) x = x.permute(1, 0, 2) # 自注意力计算:query/key/value均使用输入序列本身 attn_output, _ = self.multihead_attn(query=x, key=x, value=x) # 转换回(batch_size, seq_len, feature_dim) attn_output = attn_output.permute(1, 0, 2) # 全局均值池化:对序列维度取平均,得到全局特征 pooled = torch.mean(attn_output, dim=1) # 输出分类logits logits = self.classifier(pooled) return logits
3. 训练与验证流程
设置损失函数、优化器,执行标准训练循环:
示例代码:
device = torch.device("cuda" if torch.cuda.is_available() else "cpu") model = AttentionClassifier().to(device) criterion = torch.nn.CrossEntropyLoss() optimizer = torch.optim.Adam(model.parameters(), lr=1e-4) # 训练阶段 model.train() for epoch in range(10): total_loss = 0.0 for batch_x, batch_y in train_loader: batch_x = batch_x.to(device) batch_y = batch_y.to(device) optimizer.zero_grad() logits = model(batch_x) loss = criterion(logits, batch_y) loss.backward() optimizer.step() total_loss += loss.item() * batch_x.size(0) avg_loss = total_loss / len(train_loader.dataset) print(f"Epoch {epoch+1}, Average Loss: {avg_loss:.4f}") # 验证阶段 model.eval() correct = 0 total = 0 with torch.no_grad(): for batch_x, batch_y in val_loader: batch_x = batch_x.to(device) batch_y = batch_y.to(device) logits = model(batch_x) _, preds = torch.max(logits, 1) total += batch_y.size(0) correct += (preds == batch_y).sum().item() accuracy = correct / total print(f"Validation Accuracy: {accuracy:.4f}")
关键注意事项
- 维度转换:
MultiheadAttention的输入维度顺序要求严格,必须做好permute操作,否则会触发维度不匹配错误。 - head数选择:
embed_dim必须能被num_heads整除,比如300可选择3、5、6、10等数值。 - 池化方式:除均值池化外,还可使用最大池化,或在输入序列前添加特殊
[CLS]token,取该token的输出作为全局特征(类似BERT的做法)。 - 超参数调优:学习率、batch size、head数量、全连接层维度等需根据任务效果调整。
内容的提问来源于stack exchange,提问作者kowser66
相关产品推荐
相关产品推荐

