如何用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
相关产品推荐
相关产品推荐

