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

如何使用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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.17 21:30:56