rsobieski commited on
Commit
713c540
·
verified ·
1 Parent(s): 0c774c2

Update agent.py

Browse files
Files changed (1) hide show
  1. agent.py +20 -17
agent.py CHANGED
@@ -2,7 +2,7 @@ import os
2
  import pandas as pd
3
  from langchain_core.messages import HumanMessage, AIMessage
4
  from langgraph.graph import StateGraph, MessagesState
5
- from langchain_community.llms import HuggingFaceHub
6
  from tools import TOOLS
7
 
8
  # --- Read local QA data for retriever ---
@@ -14,18 +14,18 @@ qa_dict = {
14
  }
15
 
16
  def build_graph():
17
- # Initialize Mistral model using HuggingFaceHub
18
- llm = HuggingFaceHub(
19
- repo_id="mistralai/Mistral-7B-Instruct-v0.3",
20
- model_kwargs={
21
- "max_new_tokens": 512,
22
- "temperature": 0.1,
23
- "do_sample": False
24
- },
25
  huggingfacehub_api_token=os.environ["HF_TOKEN"]
26
  )
27
 
28
- # Retriever node (unchanged)
29
  def retriever_node(state: MessagesState):
30
  query = state["messages"][-1].content.strip()
31
  if query in qa_dict:
@@ -34,7 +34,7 @@ def build_graph():
34
  print("🔍 No match. Sending to LLM.")
35
  return {"messages": state["messages"]}
36
 
37
- # Assistant node (LLM) - updated for HuggingFaceHub
38
  def assistant_node(state: MessagesState):
39
  query = state["messages"][-1].content.strip()
40
 
@@ -49,19 +49,22 @@ def build_graph():
49
  )
50
 
51
  # Format prompt for Mistral model
52
- prompt = f"<s>[INST] {system_prompt}\n\n{query} [/INST]"
 
 
 
53
 
54
  # Generate response
55
- response = llm.invoke(prompt).strip()
56
 
57
  # Clean up response
58
- for tag in ("Final answer:", "Answer:", "assistant:"):
59
  if response.lower().startswith(tag.lower()):
60
  response = response[len(tag):].strip()
61
 
62
  return {"messages": [AIMessage(content=response.strip())]}
63
 
64
- # Tool node (unchanged)
65
  def tool_node(state: MessagesState):
66
  try:
67
  tool_signal = state.get("tool_call", "")
@@ -82,7 +85,7 @@ def build_graph():
82
  print(f"⚠️ Tool error: {e}")
83
  return {"messages": [AIMessage(content="Unknown")]} # fail-safe
84
 
85
- # Build LangGraph (unchanged)
86
  builder = StateGraph(MessagesState)
87
  builder.add_node("retriever", retriever_node)
88
  builder.add_node("assistant", assistant_node)
@@ -96,7 +99,7 @@ def build_graph():
96
 
97
  return builder.compile()
98
 
99
- # Agent class (unchanged)
100
  class BasicAgent:
101
  def __init__(self):
102
  print("✅ BasicAgent initialized with retriever + LLM + tools")
 
2
  import pandas as pd
3
  from langchain_core.messages import HumanMessage, AIMessage
4
  from langgraph.graph import StateGraph, MessagesState
5
+ from langchain_huggingface import HuggingFaceEndpoint
6
  from tools import TOOLS
7
 
8
  # --- Read local QA data for retriever ---
 
14
  }
15
 
16
  def build_graph():
17
+ # Initialize Mistral model with proper configuration
18
+ llm = HuggingFaceEndpoint(
19
+ endpoint_url="https://api-inference.huggingface.co/models/mistralai/Mistral-7B-Instruct-v0.3",
20
+ task="text-generation",
21
+ max_new_tokens=512,
22
+ temperature=0.1,
23
+ top_k=50,
24
+ top_p=0.95,
25
  huggingfacehub_api_token=os.environ["HF_TOKEN"]
26
  )
27
 
28
+ # Retriever node
29
  def retriever_node(state: MessagesState):
30
  query = state["messages"][-1].content.strip()
31
  if query in qa_dict:
 
34
  print("🔍 No match. Sending to LLM.")
35
  return {"messages": state["messages"]}
36
 
37
+ # Assistant node (LLM)
38
  def assistant_node(state: MessagesState):
39
  query = state["messages"][-1].content.strip()
40
 
 
49
  )
50
 
51
  # Format prompt for Mistral model
52
+ messages = [
53
+ {"role": "system", "content": system_prompt},
54
+ {"role": "user", "content": query}
55
+ ]
56
 
57
  # Generate response
58
+ response = llm.invoke(messages).strip()
59
 
60
  # Clean up response
61
+ for tag in ("Final answer:", "Answer:", "assistant:", "assistant:"):
62
  if response.lower().startswith(tag.lower()):
63
  response = response[len(tag):].strip()
64
 
65
  return {"messages": [AIMessage(content=response.strip())]}
66
 
67
+ # Tool node
68
  def tool_node(state: MessagesState):
69
  try:
70
  tool_signal = state.get("tool_call", "")
 
85
  print(f"⚠️ Tool error: {e}")
86
  return {"messages": [AIMessage(content="Unknown")]} # fail-safe
87
 
88
+ # Build LangGraph
89
  builder = StateGraph(MessagesState)
90
  builder.add_node("retriever", retriever_node)
91
  builder.add_node("assistant", assistant_node)
 
99
 
100
  return builder.compile()
101
 
102
+ # Agent class
103
  class BasicAgent:
104
  def __init__(self):
105
  print("✅ BasicAgent initialized with retriever + LLM + tools")