A53040 | 程序题:检索增强生成这是一个基于检索增强生成(Retrieval-Augmented Generation, RAG)的问答系统实 现。该系统包含两个主要组件: 1. 检索器(Retriever):使用嵌入模型将文档转换为向量,根据查询语义检索最相关的文档 2. 生成器(Generator):基于检索到的相关文档和原始问题,生成准确的答案 请根据代码逻辑和注释提示,补全Python代码中的空位。…
来源2025年
时间限制1s
内存限制256MB
通过 / 提交0/0
题目描述
程序题:检索增强生成
这是一个基于检索增强生成(Retrieval-Augmented Generation, RAG)的问答系统实 现。该系统包含两个主要组件:
1. 检索器(Retriever):使用嵌入模型将文档转换为向量,根据查询语义检索最相关的文档
2. 生成器(Generator):基于检索到的相关文档和原始问题,生成准确的答案 请根据代码逻辑和注释提示,补全Python代码中的空位。
import torch
import torch.nn.functional as F
from transformers import AutoTokenizer, AutoModel,
AutoModelForCausalLM
class Retriever:
def __init__(self, embedder, tokenizer, corpus):
self.embedder = embedder
self.tokenizer = tokenizer
self.corpus = corpus
self.embeddings = self.build_index(corpus)
def _text_to_vector(self, text_batch):
enc = ____[21]____
enc = {k: v.to(self.embedder.device) for k, v inenc.items()}
with torch.no_grad():
outputs = ____[22]____
last_hidden = outputs.last_hidden_state
attn_mask = enc["attention_mask"]
sent_vecs = ____[23]____
return sent_vecs
def build_index(self, corpus):
vecs_tensor = self._text_to_vector(corpus)
vecs_normalized = F.normalize(vecs_tensor, p=2, dim=1)
return vecs_normalized.cpu()
def search(self, query_vec, topk=3):
if query_vec.dim() == 1:
query_vec = query_vec.unsqueeze(0)
query_normalized = F.normalize(query_vec, p=2, dim=1)
similarities = torch.matmul(query_normalized,self.embeddings.T)
similarities = similarities.squeeze(0)
topk_scores, topk_indices = torch.topk(similarities,k=min(topk, len(self.corpus)))
results = [
{"text": self.corpus[idx], "score":float(topk_scores[i])}
for i, idx in enumerate(topk_indices)
]
return results
def rag_answer(query, retriever, generator, gen_tokenizer, topk=3,
max_new_tokens=128):
query_vec_tensor = retriever._text_to_vector([query])
query_vec = query_vec_tensor.squeeze(0).cpu()
docs = ____[24]____
context = "\n".join([d["text"] for d in docs]) if docs else ""
prompt = f"Use the context to answer.\nContext:\n{context}\n\nQ: {query}\nA:"
inputs = gen_tokenizer(prompt, return_tensors="pt")
inputs = {k: v.to(generator.device) for k, v in inputs.items()}
pad_id = gen_tokenizer.pad_token_id if gen_tokenizer.pad_token_id is not None else gen_tokenizer.eos_token_id
with torch.no_grad():
output_ids = generator.generate( **inputs,
do_sample=True, temperature=0.7, top_p=0.9,
max_new_tokens=max_new_tokens,
pad_token_id=pad_id, eos_token_id=gen_tokenizer.eos_token_id,
)
new_tokens = ____[25]____
answer = gen_tokenizer.decode(new_tokens,
skip_special_tokens=True).strip()
return answer
if __name__ == "__main__":
embedder_name = "../Qwen/Qwen3-Embedding-0.6B"
retriever_tokenizer =
AutoTokenizer.from_pretrained(embedder_name)
embedder = AutoModel.from_pretrained(embedder_name)
generator_name = "../Qwen/Qwen3-0.6B"
gen_tokenizer = AutoTokenizer.from_pretrained(generator_name)
generator = AutoModelForCausalLM.from_pretrained(generator_name)
corpus = [
"Paris is the capital of France.",
"Tokyo is the capital of Japan.",
"Beijing is the capital of China.",
]
retriever = Retriever(embedder, retriever_tokenizer, corpus)
query = "What is the capital of France?"
ans = rag_answer(query, retriever, generator, gen_tokenizer,topk=2)
print("Q:", query)
print("A:", ans)暂无题解
C++ 编辑器
输入
输出
可保存默认模板;新题优先使用已保存模板。
当前快捷键仅展示,暂不支持修改。
- 撤销
Ctrl / ⌘ + Z - 重做
Ctrl / ⌘ + Y - 查找
Ctrl / ⌘ + F - 全选
Ctrl / ⌘ + A - 复制
Ctrl / ⌘ + C - 剪切
Ctrl / ⌘ + X - 粘贴
Ctrl / ⌘ + V - 自动排版
工具栏排版按钮 - 草稿保存
编辑时自动保存到本机
历史
提交记录
状态说明时间源码
AI
作答助手
你好,我是作答助手。可以问思路、复杂度、样例含义或代码报错原因;不会直接给出完整 AC 代码。
确定要清空代码吗?