Commit 2e30afa1 authored by Delvallez Delvallez's avatar Delvallez Delvallez

system.* de RAG4HN fonctionne en local

parent 8829efd9
{
"cells": [
{
"cell_type": "code",
"execution_count": 1,
"id": "afdf459d",
"metadata": {},
"outputs": [
{
"name": "stderr",
"output_type": "stream",
"text": [
"/mnt/memsupplementaire/ProjetIndividuel/Travaux/stagexairag/challenge/gits/Annote_Retrieval-Augmented-Generation-for-historical-newspapers-main/raghn/lib/python3.12/site-packages/tqdm/auto.py:21: TqdmWarning: IProgress not found. Please update jupyter and ipywidgets. See https://ipywidgets.readthedocs.io/en/stable/user_install.html\n",
" from .autonotebook import tqdm as notebook_tqdm\n"
]
}
],
"source": [
"from langchain_ollama import ChatOllama\n",
"from langchain_core.prompts import PromptTemplate\n",
"from langchain_tavily import TavilySearch\n",
"from langchain_cohere import CohereRerank\n",
"from langchain_core.output_parsers import StrOutputParser\n",
"from typing_extensions import TypedDict\n",
"from typing import List\n",
"from langchain_core.documents import Document\n",
"from langchain_chroma import Chroma\n",
"from langchain_huggingface import HuggingFaceEmbeddings\n",
"import numpy as np\n",
"# from rank_bm25 import BM25Okapi\n",
"from flair.data import Sentence\n",
"from flair.models import SequenceTagger\n",
"from sklearn.feature_extraction.text import TfidfVectorizer\n",
"from sklearn.metrics.pairwise import cosine_similarity\n",
"from langgraph.graph import StateGraph\n",
"\n",
"from pprint import pprint\n",
"\n",
"from dotenv import load_dotenv\n",
"load_dotenv()\n",
"\n",
"import os\n"
]
},
{
"cell_type": "code",
"execution_count": 2,
"id": "2096d49c-d3dc-4329-ada7-aff56d210198",
"metadata": {},
"outputs": [],
"source": [
"local_llm = 'llama3'"
]
},
{
"cell_type": "code",
"execution_count": 3,
"id": "267c63e1-4c2f-439d-8d95-4c6aa01f41cf",
"metadata": {
"metadata": {}
},
"outputs": [
{
"name": "stdout",
"output_type": "stream",
"text": [
"2026-04-28 17:01:43,925 SequenceTagger predicts: Dictionary with 17 tags: O, S-PER, B-PER, E-PER, I-PER, S-LOC, B-LOC, E-LOC, I-LOC, S-ORG, B-ORG, E-ORG, I-ORG, S-HumanProd, B-HumanProd, E-HumanProd, I-HumanProd\n"
]
}
],
"source": [
"# Création des retriver pour les titre/summary et pour les documents\n",
"model_name = \"intfloat/e5-small\"\n",
"model_kwargs = {'device': 'cpu'}\n",
"encode_kwargs = {'normalize_embeddings': True}\n",
"hf = HuggingFaceEmbeddings(\n",
" model_name=model_name,\n",
" model_kwargs=model_kwargs,\n",
" encode_kwargs=encode_kwargs,\n",
")\n",
"vectordb = Chroma(persist_directory=\"corpus_db\", embedding_function = hf)\n",
"titledb = Chroma(persist_directory=\"title_db\", embedding_function=hf)\n",
"top_retrieve = 20\n",
"\n",
"# Définition documents\n",
"retriever = vectordb.as_retriever(search_type=\"mmr\",\n",
" search_kwargs={'k': top_retrieve, 'lambda_mult': 0.25}\n",
")\n",
"title_retriever = titledb.as_retriever() # retriever titre/summary\n",
"\n",
"\n",
"# Outils pour le reranking\n",
"compressor = CohereRerank(model = 'rerank-multilingual-v3.0', top_n = top_retrieve, cohere_api_key=os.getenv(\"COHERE_API_KEY\"))\n",
"\n",
"vectorizer = TfidfVectorizer()\n",
"tagger = SequenceTagger.load(\"hmbert/flair-hipe-2022-newseye-fr\")"
]
},
{
"cell_type": "code",
"execution_count": null,
"id": "566a4ded-00db-4a2f-8105-7a0e7290de46",
"metadata": {},
"outputs": [],
"source": [
"\n"
]
},
{
"cell_type": "code",
"execution_count": 4,
"id": "7416fb59",
"metadata": {},
"outputs": [],
"source": [
"# Implem du reranker\n",
"def rerank(docs, question):\n",
" rerank_docs = compressor.compress_documents(docs, question)\n",
" texts = [doc.page_content for doc in rerank_docs]\n",
" print(rerank_docs[0].metadata)\n",
" ners = [doc.metadata['title'] for doc in rerank_docs]\n",
" sentence = Sentence(question)\n",
" tagger.predict(sentence)\n",
" sen_dict = sentence.to_dict(tag_type='ner')\n",
" aner = \" \".join([ner['labels'][0]['value'] for ner in sen_dict['entities']] + ['O'])\n",
"\n",
" all_ner = ners + [aner]\n",
" tfidf_matrix = vectorizer.fit_transform(all_ner)\n",
" query_vector = tfidf_matrix[-1]\n",
" doc_vectors = tfidf_matrix[:-1]\n",
" ner_scores = cosine_similarity(query_vector, doc_vectors).flatten()\n",
" co_scores = np.array([float(doc.metadata['relevance_score']) for doc in rerank_docs])\n",
"\n",
" scores = 0.8 * co_scores + 0.2 * ner_scores\n",
" max_idx = np.argsort(-scores)\n",
" final_docs = []\n",
" for idx in max_idx[:3]:\n",
" if scores[idx] > 0.5:\n",
" final_docs.append(texts[idx])\n",
" return final_docs"
]
},
{
"cell_type": "code",
"execution_count": 5,
"id": "1d531a81-6d4d-405e-975a-01ef1c9679fa",
"metadata": {},
"outputs": [],
"source": [
"# Définition du LLM et du prompt à compléter\n",
"prompt = PromptTemplate(\n",
" template=\"\"\"<|begin_of_text|><|start_header_id|>system<|end_header_id|> You are an assistant for question-answering tasks. \n",
" Use the following pieces of retrieved context to answer the question. If you don't know the answer, just say that you don't know. \n",
" Use three sentences maximum and keep the answer concise <|eot_id|><|start_header_id|>user<|end_header_id|>\n",
" Question: {question} \n",
" Context: {context} \n",
" Answer: <|eot_id|><|start_header_id|>assistant<|end_header_id|>\"\"\",\n",
" input_variables=[\"question\", \"document\"],\n",
")\n",
"\n",
"llm = ChatOllama(model=local_llm, temperature=0.3)\n",
"\n",
"rag_chain = prompt | llm | StrOutputParser()"
]
},
{
"cell_type": "code",
"execution_count": null,
"id": "59be4871-afb1-45e8-9e34-0229d3fd6493",
"metadata": {},
"outputs": [],
"source": [
"# Assemblage des éléments\n",
"class GraphState(TypedDict):\n",
" \"\"\"\n",
" Represents the state of our graph.\n",
"\n",
" Attributes:\n",
" question: question\n",
" generation: LLM generation\n",
" web_search: whether to add search\n",
" documents: list of documents \n",
" \"\"\"\n",
" question : str\n",
" generation : str \n",
" title: List[str]\n",
" documents : List[str]\n",
"\n",
"\n",
"def title_retrieve(state):\n",
" \"\"\"\n",
" Retrieve titles from vectorstore\n",
"\n",
" Args:\n",
" state (dict): The current graph state\n",
"\n",
" Returns:\n",
" state (dict): New key added to state, documents, that contains retrieved documents\n",
" \"\"\"\n",
" print(\"---TITLE RETRIEVE---\")\n",
" question = state['question']\n",
"\n",
" titles = title_retriever.invoke(question)\n",
" title = [t.page_content for t in titles]\n",
" print(f\"{len(title)} titles retrieved\")\n",
" return {'title': title, 'question': question}\n",
" # return {'question': question}\n",
" \n",
"\n",
"def retrieve(state):\n",
" \"\"\"\n",
" Retrieve documents from vectorstore\n",
"\n",
" Args:\n",
" state (dict): The current graph state\n",
"\n",
" Returns:\n",
" state (dict): New key added to state, documents, that contains retrieved documents\n",
" \"\"\"\n",
" print(\"---RETRIEVE---\")\n",
" question = state[\"question\"]\n",
" title = state['title']\n",
" docs = retriever.invoke(question)\n",
" \n",
" print(f\"T{len(docs)} titles retieved\")\n",
" print(\"---RERANK---\")\n",
" refined_docs = rerank(docs, question)\n",
" return {\"documents\": refined_docs, \"question\": question}\n",
"\n",
"def generate(state):\n",
" \"\"\"\n",
" Generate answer using RAG on retrieved documents\n",
"\n",
" Args:\n",
" state (dict): The current graph state\n",
"\n",
" Returns:\n",
" state (dict): New key added to state, generation, that contains LLM generation\n",
" \"\"\"\n",
" print(\"---GENERATE---\")\n",
" question = state[\"question\"]\n",
" documents = state[\"documents\"]\n",
" \n",
" generation = rag_chain.invoke({\"context\": documents, \"question\": question})\n",
" return {\"documents\": documents, \"question\": question, \"generation\": generation}\n",
" "
]
},
{
"cell_type": "code",
"execution_count": 7,
"id": "07fa3d08-6a86-4705-a28b-e2721070bc5e",
"metadata": {},
"outputs": [
{
"data": {
"text/plain": [
"<langgraph.graph.state.StateGraph at 0x72d0f6c1e570>"
]
},
"execution_count": 7,
"metadata": {},
"output_type": "execute_result"
}
],
"source": [
"workflow = StateGraph(GraphState)\n",
"\n",
"workflow.add_node(\"title_retrieve\", title_retrieve)\n",
"workflow.add_node(\"retrieve\", retrieve)\n",
"workflow.add_node(\"generate\", generate)"
]
},
{
"cell_type": "code",
"execution_count": 8,
"id": "d9a4b9e4-3ba8-47d6-958c-e5a7112ac6f4",
"metadata": {},
"outputs": [
{
"data": {
"text/plain": [
"<langgraph.graph.state.StateGraph at 0x72d0f6c1e570>"
]
},
"execution_count": 8,
"metadata": {},
"output_type": "execute_result"
}
],
"source": [
"# On ajoute les lien entre les différentes étapes/noeuds\n",
"workflow.set_entry_point(\"title_retrieve\")\n",
"workflow.add_edge(\"title_retrieve\", \"retrieve\")\n",
"workflow.add_edge(\"retrieve\", \"generate\")"
]
},
{
"cell_type": "code",
"execution_count": 9,
"id": "13043b0f-17c7-49d3-9ea7-8f2c0f0c8691",
"metadata": {
"scrolled": true
},
"outputs": [
{
"name": "stdout",
"output_type": "stream",
"text": [
"{'question': 'Qui est Antoine Meillet'}\n",
"---TITLE RETRIEVE---\n",
"{len(title)} titles retieved\n",
"'Finished running: title_retrieve:'\n",
"---RETRIEVE---\n",
"T20 titles retieved\n",
"---RERANK---\n",
"{'title': 'Antoine Meillet', 'relevance_score': 0.99995124}\n",
"'Finished running: retrieve:'\n",
"---GENERATE---\n",
"'Finished running: generate:'\n",
"('Antoine Meillet was a French linguist and philologist. He is considered the '\n",
" \"founder of sociolinguistics and was critical of Ferdinand de Saussure's \"\n",
" '\"Cours de linguistique générale\". He supervised the work of Milman Parry, '\n",
" \"who studied oral tradition and epic poetry in the Balkans with Meillet's \"\n",
" 'guidance.')\n"
]
}
],
"source": [
"app = workflow.compile()\n",
"\n",
"inputs = {\"question\": \"Qui est Antoine Meillet\"}\n",
"print(inputs)\n",
"for output in app.stream(inputs):\n",
" for key, value in output.items():\n",
" pprint(f\"Finished running: {key}:\")\n",
"pprint(value[\"generation\"])"
]
},
{
"cell_type": "code",
"execution_count": null,
"id": "9d247891-0d76-47c2-873e-6c4d519ecd5d",
"metadata": {},
"outputs": [],
"source": []
}
],
"metadata": {
"kernelspec": {
"display_name": "raghn (3.12.3)",
"language": "python",
"name": "python3"
},
"language_info": {
"codemirror_mode": {
"name": "ipython",
"version": 3
},
"file_extension": ".py",
"mimetype": "text/x-python",
"name": "python",
"nbconvert_exporter": "python",
"pygments_lexer": "ipython3",
"version": "3.12.3"
}
},
"nbformat": 4,
"nbformat_minor": 5
}
#!/usr/bin/env python
# coding: utf-8
# In[1]:
from langchain_ollama import ChatOllama
from langchain_core.prompts import PromptTemplate
from langchain_tavily import TavilySearch
from langchain_cohere import CohereRerank
from langchain_core.output_parsers import StrOutputParser
from typing_extensions import TypedDict
from typing import List
from langchain_core.documents import Document
from langchain_chroma import Chroma
from langchain_huggingface import HuggingFaceEmbeddings
import numpy as np
# from rank_bm25 import BM25Okapi
from flair.data import Sentence
from flair.models import SequenceTagger
from sklearn.feature_extraction.text import TfidfVectorizer
from sklearn.metrics.pairwise import cosine_similarity
from langgraph.graph import StateGraph
from pprint import pprint
from dotenv import load_dotenv
load_dotenv()
import os
# In[2]:
local_llm = 'llama3'
# In[3]:
# Création des retriver pour les titre/summary et pour les documents
model_name = "intfloat/e5-small"
model_kwargs = {'device': 'cpu'}
encode_kwargs = {'normalize_embeddings': True}
hf = HuggingFaceEmbeddings(
model_name=model_name,
model_kwargs=model_kwargs,
encode_kwargs=encode_kwargs,
)
vectordb = Chroma(persist_directory="corpus_db", embedding_function = hf)
titledb = Chroma(persist_directory="title_db", embedding_function=hf)
top_retrieve = 20
# Définition documents
retriever = vectordb.as_retriever(search_type="mmr",
search_kwargs={'k': top_retrieve, 'lambda_mult': 0.25}
)
title_retriever = titledb.as_retriever() # retriever titre/summary
# Outils pour le reranking
compressor = CohereRerank(model = 'rerank-multilingual-v3.0', top_n = top_retrieve, cohere_api_key=os.getenv("COHERE_API_KEY"))
vectorizer = TfidfVectorizer()
tagger = SequenceTagger.load("hmbert/flair-hipe-2022-newseye-fr")
# In[ ]:
# In[4]:
# Implem du reranker
def rerank(docs, question):
rerank_docs = compressor.compress_documents(docs, question)
texts = [doc.page_content for doc in rerank_docs]
print(rerank_docs[0].metadata)
ners = [doc.metadata['title'] for doc in rerank_docs]
sentence = Sentence(question)
tagger.predict(sentence)
sen_dict = sentence.to_dict(tag_type='ner')
aner = " ".join([ner['labels'][0]['value'] for ner in sen_dict['entities']] + ['O'])
all_ner = ners + [aner]
tfidf_matrix = vectorizer.fit_transform(all_ner)
query_vector = tfidf_matrix[-1]
doc_vectors = tfidf_matrix[:-1]
ner_scores = cosine_similarity(query_vector, doc_vectors).flatten()
co_scores = np.array([float(doc.metadata['relevance_score']) for doc in rerank_docs])
scores = 0.8 * co_scores + 0.2 * ner_scores
max_idx = np.argsort(-scores)
final_docs = []
for idx in max_idx[:3]:
if scores[idx] > 0.5:
final_docs.append(texts[idx])
return final_docs
# In[5]:
# Définition du LLM et du prompt à compléter
prompt = PromptTemplate(
template="""<|begin_of_text|><|start_header_id|>system<|end_header_id|> You are an assistant for question-answering tasks.
Use the following pieces of retrieved context to answer the question. If you don't know the answer, just say that you don't know.
Use three sentences maximum and keep the answer concise <|eot_id|><|start_header_id|>user<|end_header_id|>
Question: {question}
Context: {context}
Answer: <|eot_id|><|start_header_id|>assistant<|end_header_id|>""",
input_variables=["question", "document"],
)
llm = ChatOllama(model=local_llm, temperature=0.3)
rag_chain = prompt | llm | StrOutputParser()
# In[ ]:
# Assemblage des éléments
class GraphState(TypedDict):
"""
Represents the state of our graph.
Attributes:
question: question
generation: LLM generation
web_search: whether to add search
documents: list of documents
"""
question : str
generation : str
title: List[str]
documents : List[str]
def title_retrieve(state):
"""
Retrieve titles from vectorstore
Args:
state (dict): The current graph state
Returns:
state (dict): New key added to state, documents, that contains retrieved documents
"""
print("---TITLE RETRIEVE---")
question = state['question']
titles = title_retriever.invoke(question)
title = [t.page_content for t in titles]
print(f"{len(title)} titles retrieved")
return {'title': title, 'question': question}
# return {'question': question}
def retrieve(state):
"""
Retrieve documents from vectorstore
Args:
state (dict): The current graph state
Returns:
state (dict): New key added to state, documents, that contains retrieved documents
"""
print("---RETRIEVE---")
question = state["question"]
title = state['title']
docs = retriever.invoke(question)
print(f"T{len(docs)} titles retieved")
print("---RERANK---")
refined_docs = rerank(docs, question)
return {"documents": refined_docs, "question": question}
def generate(state):
"""
Generate answer using RAG on retrieved documents
Args:
state (dict): The current graph state
Returns:
state (dict): New key added to state, generation, that contains LLM generation
"""
print("---GENERATE---")
question = state["question"]
documents = state["documents"]
generation = rag_chain.invoke({"context": documents, "question": question})
return {"documents": documents, "question": question, "generation": generation}
# In[7]:
workflow = StateGraph(GraphState)
workflow.add_node("title_retrieve", title_retrieve)
workflow.add_node("retrieve", retrieve)
workflow.add_node("generate", generate)
# In[8]:
# On ajoute les lien entre les différentes étapes/noeuds
workflow.set_entry_point("title_retrieve")
workflow.add_edge("title_retrieve", "retrieve")
workflow.add_edge("retrieve", "generate")
# In[9]:
app = workflow.compile()
inputs = {"question": "Qui est Antoine Meillet"}
print(inputs)
for output in app.stream(inputs):
for key, value in output.items():
pprint(f"Finished running: {key}:")
pprint(value["generation"])
# In[ ]:
Markdown is supported
0% or
You are about to add 0 people to the discussion. Proceed with caution.
Finish editing this message first!
Please register or to comment