Agent学习记录-8

阅读学习时间:约5分钟

笔记

Agent的学习记录--8

来源Github的学习指南Agent Learning Hub的学习:

Stage 2——综合所学笔记4-7

综合实战:带可信来源与证据链的学术科研 Agent:

import os
import re
import sys
import json
import sqlite3
from typing import List, Dict, Any, Optional

from config import get_llm
from langchain_core.tools import tool
from langchain_core.messages import SystemMessage, HumanMessage, ToolMessage, AIMessage

import importlib

我们将前四个模块的代码引入:

mod_01 = importlib.import_module("01_tools_ecosystem")
query_sqlite = mod_01.query_sqlite
execute_python_code = mod_01.execute_python_code
search_academic = mod_01.search_academic

mod_02 = importlib.import_module("02_rag_pipeline")
load_and_chunk_documents = mod_02.load_and_chunk_documents
SimpleVectorRetriever = mod_02.SimpleVectorRetriever

mod_03 = importlib.import_module("03_memory_systems")
SessionMemoryManager = mod_03.SessionMemoryManager
LongTermMemoryStore = mod_03.LongTermMemoryStore
trim_context_window = mod_03.trim_context_window

mod_04 = importlib.import_module("04_robustness_guardrails")
safe_execute_tool = mod_04.safe_execute_tool
LoopCircuitBreaker = mod_04.LoopCircuitBreaker
CitationGroundingGuard = mod_04.CitationGroundingGuard、

# 初始化 RAG 向量索引
print(">>> [系统初始化] 正在加载学术文献库并构建向量索引...")
doc_chunks = load_and_chunk_documents()
rag_retriever = SimpleVectorRetriever(doc_chunks)

# 全局暂存本次会话中实际召回的文档,供最后的引用校验器比对
SESSION_RETRIEVED_DOCS: List[Dict[str, Any]] = []

我们于是将RAG接入Agent工具:

@tool
def retrieve_papers_rag(query: str) -> str:
    global SESSION_RETRIEVED_DOCS
    print(f"  [Tool: RAG] 正在检索学术文献库: '{query}' ...")
    related_doc = rag_retriever.retrieve(query, top_k = 2)
    if not related_doc:
        return "未检索到与查询匹配的学术文献片段。"

    return_doc = []
    for doc, score in related_doc: # related_doc 为元组
        doc_id = len(SESSION_RETRIEVED_DOCS) + 1
        SESSION_RETRIEVED_DOCS.append({
            "id": doc_id,
            "text": doc.page_content,
            "source": doc.metadata.get("source", "unknown")
        })
        return_doc.append(
            f"[Doc {doc_id}] (来源文件: {doc.metadata.get('source', 'unknown')} | 相关度: {score:.3f}):\n{doc.page_content}"
        )

    return "\n\n".join(return_doc)

AGENT_TOOLS = [retrieve_papers_rag, query_sqlite, execute_python_code, search_academic]
TOOL_MAP = {t.name: t for t in AGENT_TOOLS}

现在我们将所有工具组装成Agent闭环:

def run_grounded_agent(
    user_query: str,
    session_id: str = "default_user",
    session_manager: Optional[SessionMemoryManager] = None,
    memory_store: Optional[LongTermMemoryStore] = None,
    max_steps: int = 8
) -> str:
    
    global SESSION_RETRIEVED_DOCS
    SESSION_RETRIEVED_DOCS = []
    
    if session_manager is None:
        session_manager = SessionMemoryManager()
    if memory_store is None:
        memory_store = LongTermMemoryStore()
        
    llm = get_llm(temperature=0.0)
    model_with_tools = llm.bind_tools(AGENT_TOOLS) # 工具接入
    
    # 1. 注入长期记忆画像
    long_term_data = memory_store.read_memory()
    user_profile = long_term_data.get("user_profile", {})
    
    # 2. 构造具备严苛引用要求的 System Prompt
    system_prompt = (
        "你是一名严谨的顶尖学术助理。回答必须遵循【绝对事实与来源证据】原则。\n"
        f"【已知用户偏好】: {json.dumps(user_profile, ensure_ascii=False)}\n"
        "【已知数据库表结构】: papers (id, title, authors, year, citations, venue), experiments (exp_id, model, accuracy, status)\n"
        "【严格引用规则】:\n"
        "1. 任何通过文献检索获取的概念、公式、理论,必须在句末标注 [Doc X];\n"
        "2. 任何来自数据库的数据(论文发表年、被引数等),必须标明 [DB: papers];\n"
        "3. 数据统计与均值计算必须调用 Python 代码工具验算并标注结果;\n"
        "4. 回答末尾必须给出【证据与来源索引】小节;\n"
        "5. 严禁捏造未经工具返回的事实!"
    )
    
    # 3. 获取并组装会话记忆
    history = session_manager.get_history(session_id) #创建对话记录本
    raw_messages = [SystemMessage(content=system_prompt)] + history.messages + [HumanMessage(content=user_query)] 
    messages = trim_context_window(raw_messages, max_messages=6) # 短期上下文裁剪
    circuit_breaker = LoopCircuitBreaker(max_consecutive_duplicates=2) # 引入重复调用与死循环检测器
    
    print(f"\n[User Query]: {user_query}")
    
    for step in range(max_steps):
        response = model_with_tools.invoke(messages) # 传入对话记录
        messages.append(response)

        if not response.tool_calls: # 如果没有tool calling,返回最后的结果
            print("Models have reached the final answer.")
            final_answer = response.content
            break

        for tool_call in response.tool_calls:
            tool_name = tool_call["name"]
            tool_args = tool_call["args"]
            tool_id = tool_call["id"]

            is_block, block_reason = circuit_breaker.check_and_record(tool_name, tool_args) # 检测重复调用
            if is_block:
                messages.append(ToolMessage(content = block_reason, tool_call_id = tool_id))
                continue
            if tool_name in TOOL_MAP:
                result_str = safe_execute_tool(TOOL_MAP[tool_name],tool_args) # 安全执行代码
            else:
                result_str = f"We can't find the tool in the Tool Map."
            messages.append(ToolMessage(content = result_str, tool_call_id = tool_id))
    else:
        final_answer = "Model have reached the maximum iteration。"

    verification = CitationGroundingGuard.verify_citations(final_answer, SESSION_RETRIEVED_DOCS) #验证引用
    print(f"Citation checking report: is_valid = {verification['is_valid']}, cited document ids are: {verification['cited_ids']}")

    if not verification["is_valid"]:
        print(f"There are hallucination references: {verification['warnings']}")
        final_answer += f"\n\n Not verifed or excess the range of retrieve document {verification['hallucinated_doc_ids']}"

    history.add_user_message(user_query)
    history.add_ai_message(final_answer)        

    return final_answer

我们这样子就组装好一个Agent了,现在我们测试:

def test_grounded_agent():
    lt = LongTermMemoryStore()
    lt.update_profile("preferred_language", "Python")
    lt.update_profile("research_interest", "Transformer 注意力机制")
    sm = SessionMemoryManager()
    test_query = (
        "请帮我完成两个任务:\n"
        "1. 从本地文献库中检索出缩放点积注意力 (Scaled Dot-Product Attention) 的公式,并解释为什么要除以根号 dk;\n"
        "2. 从数据库中查询 2017 年发表的该论文对应的被引数是多少,并标注明确的来源!"
    )
    final_output = run_grounded_agent(test_query, session_manager=sm, memory_store=lt)
    print(final_output)
    
if __name__ == "__main__":
    test_grounded_agent()

结果输出:

## 任务1:缩放点积注意力公式及解释

### 公式
缩放点积注意力(Scaled Dot-Product Attention)的计算公式为 [Doc 1]:

\[
\text{Attention}(Q, K, V) = \text{softmax}\left(\frac{QK^T}{\sqrt{d_k}}\right)V
\]

其中,\(Q\)(查询)、\(K\)(键)的维度均为 \(d_k\),\(V\) 为值。

### 为什么要除以 \(\sqrt{d_k}\)?
除以 \(\sqrt{d_k}\) 的主要目的是**防止在 \(d_k\) 较大时点积结果过大,从而将 softmax 函数推入梯度极小的饱和区** [Doc 1]。

**理论分析**(通过 Python 代码验算):
假设 \(Q\) 和 \(K\) 的每个元素独立同分布,均值为 0,方差为 1,则:
- 点积 \(Q \cdot K\) 的期望值为 0
- 点积的方差 \(\text{Var}(Q \cdot K) = d_k\)
- 点积的标准差 \(\text{Std}(Q \cdot K) = \sqrt{d_k}\)

不同 \(d_k\) 下的方差与标准差如下表所示:

| \(d_k\) | 点积方差 | 点积标准差 | 缩放后标准差 |
|--------|----------|------------|--------------|
| 10     | 10.0     | 3.162      | 1.000        |
| 50     | 50.0     | 7.071      | 1.000        |
| 100    | 100.0    | 10.000     | 1.000        |
| 500    | 500.0    | 22.361     | 1.000        |
| 1000   | 1000.0   | 31.623     | 1.000        |

*注:上表数据由 Python 代码模拟计算得出,展示了缩放可将点积标准差归一化为 1。*

**影响**:
1. 当 \(d_k\) 较大时,未缩放的点积值可能非常大(例如 \(d_k=1000\) 时标准差约为 31.6)。
2. 这些大值输入 softmax 会导致输出接近 one-hot 分布(即某个位置的概率接近 1,其余接近 0)。
3. 在 one-hot 区域,softmax 的梯度几乎为零,造成**梯度消失**,使模型难以训练。
4. 除以 \(\sqrt{d_k}\) 后,点积的标准差被归一化为 1,softmax 的输入保持在合理范围,确保梯度能够正常流动,训练过程更加稳定。

## 任务2:查询2017年该论文的被引数

通过查询本地科研数据库(research.db),找到 2017 年发表的论文《Attention Is All You Need》的详细信息如下 [DB: papers]:

- **标题**: Attention Is All You Need
- **作者**: Vaswani et al.
- **年份**: 2017
- **会议**: NeurIPS
- **被引数**: **115,000** 次

该论文是 Transformer 架构的原始论文,其中首次提出了缩放点积注意力机制。

## 【证据与来源索引】

1. **公式与解释**:来自本地文献库 `attention_mechanism.txt` 文件,标记为 **[Doc 1]**2. **理论分析与计算**:通过 Python 代码工具验算得出,展示了不同 \(d_k\) 下的方差变化,支持缩放必要性的解释。
3. **论文被引数据**:来自本地数据库 `papers` 表,通过 SQL 查询获得,标记为 **[DB: papers]**