loydef / app.py
Hercule66's picture
Update app.py
d874c19 verified
Raw
History Blame Contribute Delete
11.6 kB
"""
Système RAG OHADA avec API Hugging Face
Optimisé pour Space CPU gratuit - Réponses en 2-5 secondes
"""
import gradio as gr
import pandas as pd
import numpy as np
import faiss
import json
import os
from datetime import datetime
from sentence_transformers import SentenceTransformer
from huggingface_hub import InferenceClient
import warnings
warnings.filterwarnings('ignore')
# ============================================
# CONFIGURATION
# ============================================
class Config:
CSV_PATH = "ohada_documents.csv"
LOGO_PATH = "loydef.jpeg"
FEEDBACK_FILE = "user_feedback.json"
EMBEDDING_MODEL = "sentence-transformers/paraphrase-multilingual-mpnet-base-v2"
API_MODEL = "Qwen/Qwen2.5-3B-Instruct" # Pas de restriction d'accès
TOP_K = 3
PRIMARY_COLOR = "#001044"
ACCENT_COLOR = "#e03a3c"
config = Config()
# ============================================
# CHARGEMENT (Uniquement embedding + FAISS)
# ============================================
print("🚀 Initialisation du système RAG OHADA...")
df = pd.read_csv(config.CSV_PATH)
print(f"✅ {len(df)} segments OHADA chargés")
embedding_model = SentenceTransformer(config.EMBEDDING_MODEL)
print("✅ Modèle d'embedding chargé")
texts = df['text'].tolist()
embeddings = embedding_model.encode(texts, show_progress_bar=True, batch_size=32)
print(f"✅ Embeddings créés : {embeddings.shape}")
dimension = embeddings.shape[1]
index = faiss.IndexFlatL2(dimension)
index.add(embeddings.astype('float32'))
print(f"✅ Index FAISS créé avec {index.ntotal} vecteurs")
# Client API
hf_token = os.environ.get("HF_TOKEN")
client = InferenceClient(token=hf_token) if hf_token else InferenceClient()
print("✅ Client API Hugging Face initialisé")
# ============================================
# GESTION DES FEEDBACKS
# ============================================
def save_feedback(question, answer, sources, feedback_type):
feedback_data = {
"timestamp": datetime.now().isoformat(),
"question": question,
"answer": answer,
"sources": [
{
"source_file": src['source'],
"segment_id": src['segment_id'],
"distance": float(src['distance'])
} for src in sources
],
"feedback": feedback_type
}
if os.path.exists(config.FEEDBACK_FILE):
with open(config.FEEDBACK_FILE, 'r', encoding='utf-8') as f:
feedbacks = json.load(f)
else:
feedbacks = []
feedbacks.append(feedback_data)
with open(config.FEEDBACK_FILE, 'w', encoding='utf-8') as f:
json.dump(feedbacks, f, ensure_ascii=False, indent=2)
return f"✅ Merci pour votre retour ({feedback_type}) !"
# ============================================
# FONCTIONS RAG
# ============================================
def search_documents(query, top_k=config.TOP_K):
query_embedding = embedding_model.encode([query])
distances, indices = index.search(query_embedding.astype('float32'), top_k)
results = []
for i, idx in enumerate(indices[0]):
results.append({
'text': df.iloc[idx]['text'],
'source': df.iloc[idx]['source_file'],
'segment_id': df.iloc[idx]['segment_id'],
'distance': distances[0][i],
'word_count': df.iloc[idx]['word_count']
})
return results
def generate_answer(question, context_docs):
"""Génère une réponse via l'API Hugging Face (rapide)"""
context = "\n\n".join([
f"Document {i+1} (Source: {doc['source']}):\n{doc['text'][:400]}"
for i, doc in enumerate(context_docs)
])
messages = [
{
"role": "system",
"content": "Tu es un assistant juridique expert en droit OHADA. Réponds de manière précise et professionnelle en te basant UNIQUEMENT sur les documents fournis. Si l'information n'est pas dans les documents, dis-le clairement."
},
{
"role": "user",
"content": f"""DOCUMENTS DE RÉFÉRENCE :
{context}
QUESTION : {question}
Fournis une réponse structurée avec :
1. La réponse directe
2. Les références aux sources pertinentes"""
}
]
try:
response = client.chat_completion(
model=config.API_MODEL,
messages=messages,
max_tokens=500,
temperature=0.7,
top_p=0.9,
)
return response.choices[0].message.content
except Exception as e:
error_msg = str(e)
if "429" in error_msg:
return "⚠️ Limite de requêtes atteinte. Veuillez réessayer dans quelques instants."
elif "authentication" in error_msg.lower():
return "⚠️ Erreur d'authentification. Vérifiez le token HF_TOKEN dans les secrets du Space."
else:
return f"❌ Erreur lors de la génération : {error_msg}"
# ============================================
# FONCTION PRINCIPALE
# ============================================
def rag_pipeline(question, top_k):
if not question.strip():
return "⚠️ Veuillez poser une question.", "", ""
# Recherche
docs = search_documents(question, top_k=int(top_k))
# Génération via API (rapide !)
answer = generate_answer(question, docs)
# Formatage des sources
sources_text = ""
for i, doc in enumerate(docs):
sources_text += f"""
📄 **Source {i+1}** : {doc['source']}
- Segment ID : {doc['segment_id']}
- Pertinence : {(1 - doc['distance']/10)*100:.1f}%
- Mots : {doc['word_count']}
*Extrait :* {doc['text'][:300]}...
---
"""
global last_query
last_query = {
'question': question,
'answer': answer,
'sources': docs
}
return answer, sources_text, "✅ Réponse générée"
def handle_like():
if 'last_query' in globals():
return save_feedback(
last_query['question'],
last_query['answer'],
last_query['sources'],
"like"
)
return "⚠️ Aucune réponse à évaluer"
def handle_dislike():
if 'last_query' in globals():
return save_feedback(
last_query['question'],
last_query['answer'],
last_query['sources'],
"dislike"
)
return "⚠️ Aucune réponse à évaluer"
# ============================================
# CSS PERSONNALISÉ
# ============================================
custom_css = f"""
.gradio-container {{
font-family: 'Inter', sans-serif;
max-width: 1400px !important;
}}
.header-container {{
background: linear-gradient(135deg, {config.PRIMARY_COLOR} 0%, #002266 100%);
padding: 30px;
border-radius: 15px;
margin-bottom: 30px;
text-align: center;
box-shadow: 0 4px 20px rgba(0,16,68,0.3);
}}
.header-title {{
color: white;
font-size: 2.5em;
font-weight: 700;
margin: 10px 0;
}}
.header-subtitle {{
color: #ffffff;
font-size: 1.2em;
opacity: 0.9;
}}
.primary-btn {{
background: linear-gradient(135deg, {config.ACCENT_COLOR} 0%, #ff4444 100%) !important;
color: white !important;
border: none !important;
padding: 12px 30px !important;
font-weight: 600 !important;
border-radius: 8px !important;
transition: all 0.3s ease !important;
}}
.response-box {{
background: #f8fafc !important;
border-left: 4px solid {config.ACCENT_COLOR} !important;
padding: 20px !important;
border-radius: 10px !important;
font-size: 1.05em !important;
line-height: 1.7 !important;
}}
.feedback-btn {{
padding: 10px 20px !important;
border-radius: 8px !important;
font-weight: 500 !important;
}}
.like-btn {{ background: #10b981 !important; color: white !important; }}
.dislike-btn {{ background: #ef4444 !important; color: white !important; }}
"""
# ============================================
# INTERFACE GRADIO
# ============================================
with gr.Blocks(css=custom_css, theme=gr.themes.Soft()) as demo:
with gr.Row():
with gr.Column(scale=1):
if os.path.exists(config.LOGO_PATH):
gr.Image(config.LOGO_PATH, height=100, show_label=False, container=False)
with gr.Column(scale=4):
gr.HTML(f"""
<div class="header-container">
<h1 class="header-title">🏛️ Assistant Juridique OHADA</h1>
<p class="header-subtitle">Propulsé par l'API Hugging Face - Réponses en 2-5 secondes ⚡</p>
<p style="color: rgba(255,255,255,0.8); font-size: 0.9em; margin-top: 10px;">
{len(df)} segments indexés | Modèle : {config.API_MODEL}
</p>
</div>
""")
with gr.Row():
with gr.Column(scale=2):
question_input = gr.Textbox(
label="❓ Votre question juridique",
placeholder="Exemple : Quelles sont les conditions de constitution d'une société anonyme ?",
lines=3
)
with gr.Accordion("⚙️ Paramètres", open=False):
top_k_slider = gr.Slider(
minimum=1,
maximum=5,
value=3,
step=1,
label="📚 Nombre de documents à consulter"
)
search_btn = gr.Button("🔍 Rechercher", variant="primary", elem_classes=["primary-btn"], size="lg")
gr.Examples(
examples=[
["Quelles sont les conditions de constitution d'une société anonyme selon l'OHADA ?"],
["Quelle est la durée maximale d'une société selon l'acte uniforme ?"],
["Quelles sont les obligations comptables des entreprises selon l'OHADA ?"],
["Comment se fait la nomination des commissaires aux comptes ?"],
],
inputs=question_input,
label="💡 Questions exemples"
)
with gr.Column(scale=3):
status_output = gr.Textbox(label="📊 Statut", interactive=False)
answer_output = gr.Textbox(label="✨ Réponse", lines=10, interactive=False, elem_classes=["response-box"])
with gr.Row():
like_btn = gr.Button("👍 Like", elem_classes=["feedback-btn", "like-btn"])
dislike_btn = gr.Button("👎 Dislike", elem_classes=["feedback-btn", "dislike-btn"])
feedback_output = gr.Textbox(label="💬 Feedback", interactive=False)
with gr.Row():
sources_output = gr.Markdown(label="📚 Sources consultées")
gr.HTML(f"""
<div style="text-align: center; margin-top: 30px; padding: 20px;
background: {config.PRIMARY_COLOR}; color: white; border-radius: 10px;">
<p style="margin: 0;">
🔒 Données officielles OHADA | 🤖 API HF | ⚡ CPU-optimized
</p>
</div>
""")
search_btn.click(
fn=rag_pipeline,
inputs=[question_input, top_k_slider],
outputs=[answer_output, sources_output, status_output]
)
like_btn.click(fn=handle_like, outputs=feedback_output)
dislike_btn.click(fn=handle_dislike, outputs=feedback_output)
if __name__ == "__main__":
demo.launch(
server_name="0.0.0.0",
server_port=7860,
share=False
)