WEAIbot / app.py
Nelly-43's picture
Update app.py
728ed59 verified
Raw
History Blame Contribute Delete
5.39 kB
import os
import glob
import warnings
import gradio as gr
from langchain_openai import ChatOpenAI
from langchain_community.vectorstores import Chroma
from langchain_huggingface import HuggingFaceEmbeddings
from langchain.agents import create_agent
from langchain.agents.middleware import dynamic_prompt, ModelRequest
from huggingface_hub import snapshot_download
warnings.filterwarnings('ignore')
os.environ["WANDB_DISABLED"] = "true"
DATA_DIR = "./chroma"
ENDPOINT_URL = "https://router.huggingface.co/v1"
MODEL_NAME = "meta-llama/Llama-3.1-70B-Instruct"
chunk_size=1500
chunk_overlap=30
separator="\n"
max_tokens=1000
splitter_type='recursive'
docs_path = f"{DATA_DIR}/refs/"
refs_path = f"{DATA_DIR}/links"
snapshot_download(repo_id="CGIAR/weai-refs",
repo_type="dataset",
token=os.getenv('HF_TOKEN'),
local_dir=DATA_DIR
)
llm_client = ChatOpenAI(base_url=ENDPOINT_URL,
api_key=os.getenv('HF_TOKEN'),
model=MODEL_NAME,
temperature=0,
max_retries=2,
extra_headers={"X-HF-Bill-To": "cgiar"}
)
embeddings = HuggingFaceEmbeddings(
model_name="sentence-transformers/all-mpnet-base-v2",
encode_kwargs={"normalize_embeddings": True},
model_kwargs = {"device": "cpu"}
)
# Create a Chroma vector database
docs_vector_db = Chroma(persist_directory=docs_path, embedding_function=embeddings)
refs_vector_db = Chroma(persist_directory=refs_path, embedding_function=embeddings)
@dynamic_prompt
def ref_context(request: ModelRequest) -> str:
"""Inject context into state messages."""
last_query = request.state["messages"][-1].text
ref_content = refs_vector_db.similarity_search(last_query, k=10)
system_message = (
"""
Use the given context to add citations to the attached research findings and:
- include source URLs as citations
- Format citations as markdown links: [Source Title](URL)
- Group sources in a "Sources:" section at the end of your response
Make any necessary edits to the findings and only use links and citations from the 'Link/URL' field in the context.
### OUTDATED LINKS
these organisations and their corresponding websites are no longer active:
- USAID
- Feed the Future
"""
f"### Context\n\n{ref_content}"
)
return system_message
@dynamic_prompt
def doc_context(request: ModelRequest) -> str:
"""Inject context into state messages."""
last_query = request.state["messages"][-1].text
doc_content = docs_vector_db.similarity_search(last_query, k=10)
system_message = (
"""You are a research agent specialized in the Women's Empowerment in Agriculture Index (WEAI).
Use the following context to answer questions.
Be as detailed as possible, but don't make up any information that's not from the context and where possible reference related studies and resources
from the context you have.
"""
f"\n\n{doc_content}"
)
return system_message
def weai_support(query: str):
findings = response_agent.invoke({"messages": [{"role": "user", "content": query}]})
response = findings['messages'][-1].content
return findings, citation_agent.invoke({"messages": [{"role": "user", "content": response}]})
response_agent = (create_agent(llm_client, tools=[], middleware=[doc_context]))
citation_agent = (create_agent(llm_client, tools=[], middleware=[ref_context]))
with gr.Blocks() as demo:
with gr.Sidebar():
gr.LoginButton()
gr.Markdown("# WEAI-bot")
chatbot = gr.Chatbot(type='messages',
allow_tags=True)
msg = gr.Textbox()
clear = gr.ClearButton([msg, chatbot])
def handle_undo(history, undo_data: gr.UndoData):
return history[:undo_data.index], history[undo_data.index]['content'][0]["text"]
def handle_retry(history, retry_data: gr.RetryData):
new_history = history[:retry_data.index]
previous_prompt = history[retry_data.index]['content'][0]["text"]
yield from support_agent_fn(previous_prompt, new_history)
def support_agent_fn(message, history):
findings, response = weai_support(message)
response = response['messages'][-1].content#.split('<|start_header_id|>assistant<|end_header_id|>')[-1].strip()
history.append({"role": "user", "content": message})
history.append({"role": "assistant", "content": response})
return "", history
def handle_like(data: gr.LikeData):
if data.liked:
print("You upvoted this response: ", data.value)
else:
print("You downvoted this response: ", data.value)
def handle_edit(history, edit_data: gr.EditData):
new_history = history[:edit_data.index]
new_history[-1]['content'] = [{"text": edit_data.value, "type": "text"}]
return new_history
msg.submit(support_agent_fn, [msg, chatbot], [msg, chatbot])
chatbot.undo(handle_undo, chatbot, [chatbot, msg])
chatbot.retry(handle_retry, chatbot, chatbot)
chatbot.like(handle_like, None, None)
chatbot.edit(handle_edit, chatbot, chatbot)
if __name__ == "__main__":
demo.launch()