CoolFace
Apppublic

rdassignies/chat_bodacc

sourceHugging Facemitupdated 2y agoView on Hugging Face
1likes
graph.py153 linesDownload Raw Back to root
1#!/usr/bin/env python32# -*- coding: utf-8 -*-3"""4Created on Sun Sep 22 15:43:16 20245 6@author: Raphaël d'Assignies (rdassignies@protonmail.ch)7"""8import json9from typing import Literal, Optional, List, Union, Any10from langchain_openai import ChatOpenAI11import pandas as pd12from langchain_core.prompts import ChatPromptTemplate13from langgraph.graph import END, StateGraph, START14from langchain_core.output_parsers import StrOutputParser15from pydantic import BaseModel, Field16from models import NatureJugement17from nodes import (GradeResults, GraphState, generate_query_node, 18                   generate_results_node, query_feedback_node, 19                   evaluate_query_node, evaluate_results_node)20import streamlit as st21 22 23 24 25# Instanciate pipeline26pipeline = StateGraph(GraphState)27 28pipeline.add_node('generate_query', generate_query_node)29pipeline.add_node('generate_results', generate_results_node)30pipeline.add_node('query_feedback', query_feedback_node)31 32# Only query33#pipeline.add_edge(START,'generate_query')34#pipeline.add_edge('generate_query', generate_query_node)35#pipeline.add_edge('generate_query', END)36 37# Full scenario38pipeline.add_edge(START,'generate_query')39pipeline.add_conditional_edges(40    'generate_query', 41    evaluate_query_node, 42    {'error_query' : 'generate_query',43     'ok' : 'generate_results'44     })45 46pipeline.add_conditional_edges(47    'generate_results', 48    evaluate_results_node,49    {50        "yes": END,51        "no": 'query_feedback',52        "max_generation_reached": END53 54    }  55)56 57 58# Création du graph59graph = pipeline.compile()60 61# Load le dataframe62df = pd.read_json('bodacc.json', orient='table')63 64# Initialise le dictionnaire65inputs = {66        'df_head': df.head().to_csv(), 67        'df': df68    }69 70# Créé un dictionnaire des sorties vide71outputs = {}72 73 74# Titre de l'application75st.title("Chat with BODACC !")76 77# Message d'avertissement78warning_message = (f"Cet outil, purement pédagogique, est basé sur des données réelles allant de {df['dateparution'].min()} "79                   f"à {df['dateparution'].max()}, et permet d'interroger le BODACC en langage naturel. Compte tenu de la variabilité des modèles, nous ne pouvons pas garantir la fiabilité des réponses.")80 81st.warning(warning_message)82# Interface utilisateur pour entrer la requête83user_query = st.text_input("Entrez votre requête:", "Trouve moi les restaurants à reprendre en Bretagne dans les 30 derniers jours")84 85 86# Afficher les résultats avec Streamlit87inputs["instructions"] = user_query88 89 90# Afficher un bouton pour démarrer la recherche91if st.button("Lancer la recherche"):92    config = {"configurable": {"thread_id": "2"}}93    94    # Étape 1 : Afficher le message "Je réfléchis..."95    st.write("Je réfléchis...")96 97    # Stream des résultats au fur et à mesure98    with st.spinner('Recherche en cours...'):99        for output in graph.stream(inputs, stream_mode='values', debug=False):100            # Ajouter les résultats au dictionnaire outputs101            for k, v in output.items():102                if k not in outputs:103                    outputs[k] = []104                outputs[k].append(v)105 106            # Ne pas afficher les messages pour les clés non pertinentes (comme error_query)107            if 'query' in output and len(output['query'])>0:108                st.write(f"query : {output['query']}")109                #st.write(outputs.get('query_feedbacks', 'pas de feedback'))110                #st.write(outputs.get('results_feedbacks', 'pas de resultfeedback'))111            if "results" in output and len(output["results"]) > 0:112                records = json.loads(output['results'])113                st.write(f"Résultats intermédiaires trouvés : {len(records)} résultats jusqu'à présent.")114 115    # Après la fin du traitement116    if "results" in outputs and len(outputs["results"]) > 0:117        # Agréger tous les résultats accumulés118        all_results = []119        for res in outputs["results"]:120            json_data = json.loads(res)  # Convertir chaque ensemble de résultats en JSON121            all_results.extend(json_data)  # Accumuler tous les résultats122 123        results_df = pd.DataFrame(all_results)  # Créer un DataFrame avec tous les résultats accumulés124        # Afficher un aperçu des résultats (jusqu'à 5 premiers)125        num_results = len(results_df)126        st.write(f"J'ai trouvé {num_results} résultats.")127        if num_results > 0:128            preview_count = min(5, num_results)  # Gérer le cas où il y a moins de 5 résultats129            st.write(f"Voici un aperçu des {preview_count} premiers résultats :")130            st.write(results_df.head(preview_count))131            132            trunc = outputs.get('truncated', 'pas de traunc')133            134            if trunc[0] == True:135                st.warning("Les résultats de votre recherche ont été tronqués car celle-ci était trop large ! ")136    137        # Convertir tous les résultats en CSV138        csv = results_df.to_csv(index=False)139 140        # Ajouter un bouton pour télécharger tous les résultats141        st.download_button(142            label="Télécharger le résultat complet au format CSV",143            data=csv,144            file_name="results.csv",145            mime="text/csv"146        )147 148    else:149        # Si aucun résultat n'est trouvé150        st.write("Aucun résultat trouvé.")151 152    153