基于Transformer的文档关联对话系统跨度检索模型实现求助
问题描述
我正在开发一个基于Transformer的跨度检索模型的文档关联对话系统,训练数据采用包含用户多轮提问(需从多文档获取信息)的数据集。目前在实现模型时卡在了查询与相关文档的编码及二者间注意力分数计算环节,不确定当前计算方式是否正确,希望得到Hugging Face Transformers库的正确实现指导。
以下是我当前的代码:
import transformers query = "What is the capital city of Turkey?" document = "Ankara is the capital city of Turkey." tokenizer = transformers.AutoTokenizer.from_pretrained("bert-base-cased") model = transformers.AutoModel.from_pretrained("bert-base-cased") query_encoded = tokenizer.encode(query, return_tensors="pt") document_encoded = tokenizer.encode(document, return_tensors="pt") query_outputs = model(query_encoded) document_outputs = model(document_encoded) attention_scores = torch.matmul(query_outputs, document_outputs.transpose(0, 1))
问题分析与正确实现
现有代码的核心问题
- 编码不规范:直接用
tokenizer.encode仅返回input_ids,未生成模型必需的attention_mask,无法屏蔽无效token(如PAD)的干扰,且没有保证输入格式的完整性。 - 输出使用错误:
model()返回的是包含多元素的元组,第一个元素才是最后一层的隐藏状态,直接用整个元组做矩阵乘法会导致维度不匹配报错。 - 注意力计算逻辑错误:未遵循Transformer注意力的缩放规则,且维度转置方式不符合跨序列注意力的计算逻辑,也没考虑掩码的作用。
正确实现方案
1. 规范编码输入
使用tokenizer()方法同时生成input_ids和attention_mask,确保模型获得完整输入,多文档场景还支持批量处理。
2. 提取正确的模型输出
从模型返回结果中提取最后一层的隐藏状态(即outputs[0]),其形状为[batch_size, seq_len, hidden_size],是后续计算的基础。
3. 标准跨注意力计算
按照Transformer的注意力机制,计算时需加入缩放因子(1/√d_k,d_k为隐藏层维度),同时可利用注意力掩码屏蔽无效token。
完整示例代码
import torch from transformers import AutoTokenizer, AutoModel query = "What is the capital city of Turkey?" document = "Ankara is the capital city of Turkey." # 初始化tokenizer和模型 tokenizer = AutoTokenizer.from_pretrained("bert-base-cased") model = AutoModel.from_pretrained("bert-base-cased") # 生成模型所需的完整输入(含input_ids、attention_mask) query_inputs = tokenizer(query, return_tensors="pt", padding=True, truncation=True) doc_inputs = tokenizer(document, return_tensors="pt", padding=True, truncation=True) # 获取最后一层隐藏状态,禁用梯度计算以节省资源 with torch.no_grad(): query_hidden = model(**query_inputs)[0] # 形状: [1, query_seq_len, hidden_size] doc_hidden = model(**doc_inputs)[0] # 形状: [1, doc_seq_len, hidden_size] # 去除batch维度,准备计算注意力 query_hidden = query_hidden.squeeze(0) # [query_seq_len, hidden_size] doc_hidden = doc_hidden.squeeze(0) # [doc_seq_len, hidden_size] # 计算带缩放的跨注意力分数 d_k = query_hidden.size(-1) attention_scores = torch.matmul(query_hidden, doc_hidden.T) / torch.sqrt(torch.tensor(d_k, dtype=torch.float32)) # 应用文档的注意力掩码,将PAD token对应的分数置为极小值 doc_mask = doc_inputs["attention_mask"].squeeze(0) attention_scores = attention_scores.masked_fill(doc_mask == 0, -1e9) # 可选:通过softmax得到归一化的注意力权重 attention_weights = torch.nn.functional.softmax(attention_scores, dim=-1)
多文档场景扩展
如果是多文档检索任务,可将多个文档批量编码后,循环计算查询与每个文档的注意力交互;若核心是跨度抽取,直接使用BertForQuestionAnswering等专门的QA模型会更高效,这类模型已经封装了跨度检索的损失计算和推理逻辑。
内容的提问来源于stack exchange,提问作者SyntaxNavigator
相关产品推荐
相关产品推荐

