LangGraph 模拟简单的多模型编排

0 阅读1分钟

image.png
假设这三个节点是不同的模型,处理不同的问题,在用户输入后使用模型判断下一步要调用哪个节点回答问题

  • node1: 讲短笑话
  • node2: 李白风格写诗
  • node3: 其他问题

路由节点返回的就是下一步要走的节点,最后并实现打字效果,使用stream优化响应等待时间

from langchain_core.messages import AIMessageChunk
from langchain.messages import AIMessage,HumanMessage,SystemMessage
from langgraph.graph import StateGraph,START,END
from langgraph.graph.message import add_messages,MessagesState
from langchain.chat_models import init_chat_model
from dotenv import load_dotenv
from rich import print as rprint
from IPython.display import display
from typing import Literal
load_dotenv()
llm=init_chat_model(
  model="deepseek-chat"
)
class overAllState(MessagesState):
  userValue: str
  nodeName: str
def node1(state:overAllState):
  res1=llm.invoke([HumanMessage(state["userValue"])])
  return {
    "messages": [res1]
  }
def node2(state:overAllState):
  res2=llm.invoke([HumanMessage(state["userValue"])])
  return {
    "messages": [res2]
  }
def node3(state:overAllState):
  res3=llm.invoke([HumanMessage(state["userValue"])])
  return {
    "messages": [res3]
  }
def route(state:overAllState) -> Literal["node1","node2","node3"]:
  userValue=state["userValue"]
  prompt=f"""
    用户的提示词:{userValue},
    判断分类,**只输出下面其中一个单词,不要多余文字**
    node1: 讲短笑话
    node2: 李白风格写诗
    node3: 其他问题
    只返回:node1 / node2 / node3
  """
  res=llm.invoke(prompt)
  route_name = res.content.strip()
  return route_name

builder=StateGraph(overAllState)
builder.add_node(node1)
builder.add_node(node2)
builder.add_node(node3)
builder.add_conditional_edges(
  source=START,
  path=route,
  path_map={
    "node1": "node1",
    "node2": "node2",
    "node3": "node3"
  }
)
builder.add_edge("node1", END)
builder.add_edge("node2", END)
builder.add_edge("node3", END)
graph=builder.compile()
display(graph)
# res=graph.invoke({"userValue":"给我讲个笑话吧"})
# rprint(res["messages"][-1].content)


for chunk in graph.stream({"userValue":"爱因斯坦终年多少岁"},stream_mode="messages"):
    mes, meta = chunk
    target_nodes = {"node1","node2","node3"}
    if meta.get("langgraph_node") in target_nodes:
        if isinstance(mes, AIMessageChunk) and mes.content:
            print(mes.content, end="", flush=True)