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

如何用PyTorch实现带注意力机制的BiLSTM问答抽取模型?

用PyTorch实现BiLSTM+注意力的问答起始/结束词预测模型

整体思路

我们需要构建双输入(问题+文档)的模型,核心流程如下:

  • 对问题和文档分别做词嵌入,得到可输入模型的词向量序列
  • 用BiLSTM分别编码问题与文档,生成包含上下文信息的语义表示
  • 引入注意力机制,让文档的每个位置都聚焦问题中的相关信息,增强文档编码的针对性
  • 基于融合注意力后的文档编码,对每个词做三分类预测:非起始/结束词、答案起始词、答案结束词

具体实现步骤及代码

1. 数据预处理(示例)

先将问题、文档转换为模型可处理的张量格式,包括分词、转词表索引、长度补齐等操作:

import torch
import torch.nn as nn
import torch.optim as optim
from torch.utils.data import Dataset, DataLoader

# 模拟词表(实际场景可通过Tokenizer从数据集构建)
vocab = {'<PAD>':0, '<UNK>':1, 'who':2, 'discovered':3, 'neptune':4, 'the':5, 'planet':6, '?':7,
         'With':8, 'a':9, 'prediction':10, 'by':11, 'Urbain':12, 'Le':13, 'Verrier':14, ',':15,
         'telescopic':16, 'observations':17, 'confirming':18, 'existence':19, 'of':20, 'major':21,
         'were':22, 'made':23, 'on':24, 'night':25, 'September':26, '23–24':27, '1846':28}

# 自定义数据集类
class QADataset(Dataset):
    def __init__(self, data, vocab, max_len_q=20, max_len_d=50):
        self.data = data
        self.vocab = vocab
        self.max_len_q = max_len_q
        self.max_len_d = max_len_d
    
    def __len__(self):
        return len(self.data)
    
    def __getitem__(self, idx):
        q, d, start_idx, end_idx = self.data[idx]
        
        # 问题转索引并补齐长度
        q_idx = [self.vocab.get(w, self.vocab['<UNK>']) for w in q.split()]
        q_idx = q_idx[:self.max_len_q] + [self.vocab['<PAD>']]*(self.max_len_q - len(q_idx))
        
        # 文档转索引并补齐长度
        d_idx = [self.vocab.get(w, self.vocab['<UNK>']) for w in d.split()]
        d_idx = d_idx[:self.max_len_d] + [self.vocab['<PAD>']]*(self.max_len_d - len(d_idx))
        
        # 构建标签:0=非起始/结束,1=起始,2=结束;无答案时标签全0
        start_label = torch.zeros(self.max_len_d)
        end_label = torch.zeros(self.max_len_d)
        if start_idx != -1:
            start_label[start_idx] = 1
        if end_idx != -1:
            end_label[end_idx] = 1
        
        return torch.tensor(q_idx), torch.tensor(d_idx), start_label, end_label

# 模拟训练数据(问题,文档,答案起始词索引,答案结束词索引;无答案时索引设为-1)
sample_data = [
    ("who discovered neptune the planet?", 
     "With a prediction by Urbain Le Verrier , telescopic observations confirming the existence of a major planet were made on the night of September 23–24 1846",
     4, 6)
]
dataset = QADataset(sample_data, vocab)
dataloader = DataLoader(dataset, batch_size=1)

2. 模型定义

class BiLSTMAttentionQA(nn.Module):
    def __init__(self, vocab_size, embed_dim, hidden_dim, num_classes=3):
        super().__init__()
        # 词嵌入层,指定padding_idx避免padding向量参与梯度更新
        self.embedding = nn.Embedding(vocab_size, embed_dim, padding_idx=vocab['<PAD>'])
        # BiLSTM编码层:双向输出,所以hidden_dim设为目标维度的一半
        self.q_bilstm = nn.LSTM(embed_dim, hidden_dim//2, bidirectional=True, batch_first=True)
        self.d_bilstm = nn.LSTM(embed_dim, hidden_dim//2, bidirectional=True, batch_first=True)
        # 注意力权重计算层
        self.attention = nn.Linear(hidden_dim, hidden_dim)
        # 起始/结束词预测分类器
        self.start_classifier = nn.Linear(hidden_dim*2, num_classes)
        self.end_classifier = nn.Linear(hidden_dim*2, num_classes)
        
    def forward(self, q_input, d_input):
        # 1. 生成词嵌入
        q_embed = self.embedding(q_input)  # shape: [batch_size, q_len, embed_dim]
        d_embed = self.embedding(d_input)  # shape: [batch_size, d_len, embed_dim]
        
        # 2. BiLSTM编码上下文
        q_output, _ = self.q_bilstm(q_embed)  # shape: [batch_size, q_len, hidden_dim]
        d_output, _ = self.d_bilstm(d_embed)  # shape: [batch_size, d_len, hidden_dim]
        
        # 3. 注意力机制:计算文档词对问题词的注意力权重
        attn_scores = torch.matmul(self.attention(d_output), q_output.transpose(1,2))  # shape: [batch_size, d_len, q_len]
        attn_weights = torch.softmax(attn_scores, dim=2)  # 对问题维度做softmax,得到注意力权重
        attn_context = torch.matmul(attn_weights, q_output)  # 加权求和得到问题的注意力表示
        
        # 4. 融合文档编码与注意力上下文
        fused_output = torch.cat([d_output, attn_context], dim=2)  # shape: [batch_size, d_len, hidden_dim*2]
        
        # 5. 预测起始和结束词
        start_logits = self.start_classifier(fused_output)  # shape: [batch_size, d_len, num_classes]
        end_logits = self.end_classifier(fused_output)  # shape: [batch_size, d_len, num_classes]
        
        return start_logits, end_logits

3. 训练与预测

# 初始化模型、损失函数、优化器
vocab_size = len(vocab)
embed_dim = 128
hidden_dim = 256
model = BiLSTMAttentionQA(vocab_size, embed_dim, hidden_dim)
# 交叉熵损失,忽略padding位置的无效标签
criterion = nn.CrossEntropyLoss()
optimizer = optim.Adam(model.parameters(), lr=1e-3)

# 训练流程示例
model.train()
for epoch in range(10):
    total_loss = 0.0
    for q_batch, d_batch, start_label, end_label in dataloader:
        optimizer.zero_grad()
        start_logits, end_logits = model(q_batch, d_batch)
        
        # 调整张量形状适配交叉熵损失要求
        start_loss = criterion(start_logits.view(-1, 3), start_label.view(-1).long())
        end_loss = criterion(end_logits.view(-1, 3), end_label.view(-1).long())
        total_batch_loss = start_loss + end_loss
        
        total_batch_loss.backward()
        optimizer.step()
        total_loss += total_batch_loss.item()
    print(f"Epoch {epoch+1}, Loss: {total_loss/len(dataloader):.4f}")

# 预测流程示例
model.eval()
with torch.no_grad():
    q_batch, d_batch, _, _ = next(iter(dataloader))
    start_logits, end_logits = model(q_batch, d_batch)
    
    # 获取每个词的预测类别
    start_pred = torch.argmax(start_logits, dim=2).squeeze().numpy()
    end_pred = torch.argmax(end_logits, dim=2).squeeze().numpy()
    
    # 输出预测结果
    doc_words = sample_data[0][1].split()
    for idx, word in enumerate(doc_words):
        if idx >= len(start_pred):
            break
        if start_pred[idx] == 1:
            pred_desc = "start word of answer"
        elif end_pred[idx] == 2:
            pred_desc = "end word of answer"
        else:
            pred_desc = "not start word or end of answer"
        print(f"the output layer predict word “{word}” is {pred_desc}")

关键优化点

  • 可替换随机初始化的词嵌入层为预训练词嵌入(如GloVe、Word2Vec),提升模型语义理解能力
  • 注意力机制可扩展为多头注意力,增强模型对不同语义信息的捕捉能力
  • 训练时可加入学习率调度器,动态调整学习率优化训练效果
  • 针对无答案的样本,可单独设计损失权重或分类分支,提升模型对无答案场景的识别能力

内容的提问来源于stack exchange,提问作者user9799714

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.22 04:22:51