CoolFace
Apppublic

jaothan/DockerGenAI_Streamlit

sourceHugging Faceupdated 2y agoView on Hugging Face
0likes
loader.py151 linesDownload Raw Back to root
1import os
2import requests
3from dotenv import load_dotenv
4from langchain_community.graphs import Neo4jGraph
5import streamlit as st
6from streamlit.logger import get_logger
7from chains import load_embedding_model
8from utils import create_constraints, create_vector_index
9from PIL import Image
10
11load_dotenv(".env")
12
13url = os.getenv("NEO4J_URI")
14username = os.getenv("NEO4J_USERNAME")
15password = os.getenv("NEO4J_PASSWORD")
16ollama_base_url = os.getenv("OLLAMA_BASE_URL")
17embedding_model_name = os.getenv("EMBEDDING_MODEL")
18# Remapping for Langchain Neo4j integration
19os.environ["NEO4J_URL"] = url
20
21logger = get_logger(__name__)
22
23so_api_base_url = "https://api.stackexchange.com/2.3/search/advanced"
24
25embeddings, dimension = load_embedding_model(
26    embedding_model_name, config={"ollama_base_url": ollama_base_url}, logger=logger
27)
28
29# if Neo4j is local, you can go to http://localhost:7474/ to browse the database
30neo4j_graph = Neo4jGraph(
31    url=url, username=username, password=password, refresh_schema=False
32)
33
34create_constraints(neo4j_graph)
35create_vector_index(neo4j_graph)
36
37
38def load_so_data(tag: str = "neo4j", page: int = 1) -> None:
39    parameters = (
40        f"?pagesize=100&page={page}&order=desc&sort=creation&answers=1&tagged={tag}"
41        "&site=stackoverflow&filter=!*236eb_eL9rai)MOSNZ-6D3Q6ZKb0buI*IVotWaTb"
42    )
43    data = requests.get(so_api_base_url + parameters).json()
44    insert_so_data(data)
45
46
47def load_high_score_so_data() -> None:
48    parameters = (
49        f"?fromdate=1664150400&order=desc&sort=votes&site=stackoverflow&"
50        "filter=!.DK56VBPooplF.)bWW5iOX32Fh1lcCkw1b_Y6Zkb7YD8.ZMhrR5.FRRsR6Z1uK8*Z5wPaONvyII"
51    )
52    data = requests.get(so_api_base_url + parameters).json()
53    insert_so_data(data)
54
55
56def insert_so_data(data: dict) -> None:
57    # Calculate embedding values for questions and answers
58    for q in data["items"]:
59        question_text = q["title"] + "\n" + q["body_markdown"]
60        q["embedding"] = embeddings.embed_query(question_text)
61        for a in q["answers"]:
62            a["embedding"] = embeddings.embed_query(
63                question_text + "\n" + a["body_markdown"]
64            )
65
66    # Cypher, the query language of Neo4j, is used to import the data
67    # https://neo4j.com/docs/getting-started/cypher-intro/
68    # https://neo4j.com/docs/cypher-cheat-sheet/5/auradb-enterprise/
69    import_query = """
70    UNWIND $data AS q
71    MERGE (question:Question {id:q.question_id}) 
72    ON CREATE SET question.title = q.title, question.link = q.link, question.score = q.score,
73        question.favorite_count = q.favorite_count, question.creation_date = datetime({epochSeconds: q.creation_date}),
74        question.body = q.body_markdown, question.embedding = q.embedding
75    FOREACH (tagName IN q.tags | 
76        MERGE (tag:Tag {name:tagName}) 
77        MERGE (question)-[:TAGGED]->(tag)
78    )
79    FOREACH (a IN q.answers |
80        MERGE (question)<-[:ANSWERS]-(answer:Answer {id:a.answer_id})
81        SET answer.is_accepted = a.is_accepted,
82            answer.score = a.score,
83            answer.creation_date = datetime({epochSeconds:a.creation_date}),
84            answer.body = a.body_markdown,
85            answer.embedding = a.embedding
86        MERGE (answerer:User {id:coalesce(a.owner.user_id, "deleted")}) 
87        ON CREATE SET answerer.display_name = a.owner.display_name,
88                      answerer.reputation= a.owner.reputation
89        MERGE (answer)<-[:PROVIDED]-(answerer)
90    )
91    WITH * WHERE NOT q.owner.user_id IS NULL
92    MERGE (owner:User {id:q.owner.user_id})
93    ON CREATE SET owner.display_name = q.owner.display_name,
94                  owner.reputation = q.owner.reputation
95    MERGE (owner)-[:ASKED]->(question)
96    """
97    neo4j_graph.query(import_query, {"data": data["items"]})
98
99
100# Streamlit
101def get_tag() -> str:
102    input_text = st.text_input(
103        "Which tag questions do you want to import?", value="neo4j"
104    )
105    return input_text
106
107
108def get_pages():
109    col1, col2 = st.columns(2)
110    with col1:
111        num_pages = st.number_input(
112            "Number of pages (100 questions per page)", step=1, min_value=1
113        )
114    with col2:
115        start_page = st.number_input("Start page", step=1, min_value=1)
116    st.caption("Only questions with answers will be imported.")
117    return (int(num_pages), int(start_page))
118
119
120def render_page():
121    datamodel_image = Image.open("./images/datamodel.png")
122    st.header("StackOverflow Loader")
123    st.subheader("Choose StackOverflow tags to load into Neo4j")
124    st.caption("Go to http://localhost:7474/ to explore the graph.")
125
126    user_input = get_tag()
127    num_pages, start_page = get_pages()
128
129    if st.button("Import", type="primary"):
130        with st.spinner("Loading... This might take a minute or two."):
131            try:
132                for page in range(1, num_pages + 1):
133                    load_so_data(user_input, start_page + (page - 1))
134                st.success("Import successful", icon="โœ…")
135                st.caption("Data model")
136                st.image(datamodel_image)
137                st.caption("Go to http://localhost:7474/ to interact with the database")
138            except Exception as e:
139                st.error(f"Error: {e}", icon="๐Ÿšจ")
140    with st.expander("Highly ranked questions rather than tags?"):
141        if st.button("Import highly ranked questions"):
142            with st.spinner("Loading... This might take a minute or two."):
143                try:
144                    load_high_score_so_data()
145                    st.success("Import successful", icon="โœ…")
146                except Exception as e:
147                    st.error(f"Error: {e}", icon="๐Ÿšจ")
148
149
150render_page()
151