File size: 3,273 Bytes
bbb5f1f
b84d865
 
c5fa4cf
b84d865
 
 
c5fa4cf
 
 
 
 
 
 
b84d865
 
 
 
 
c5fa4cf
 
 
 
 
 
 
 
b84d865
c5fa4cf
b84d865
ec336f5
 
 
c5fa4cf
 
 
ec336f5
c5fa4cf
ec336f5
 
c5fa4cf
ec336f5
 
c5fa4cf
 
ec336f5
c5fa4cf
ec336f5
 
 
 
 
 
b84d865
bbb5f1f
c5fa4cf
b84d865
c5fa4cf
 
 
 
 
 
 
b84d865
 
c5fa4cf
 
 
b84d865
c5fa4cf
b84d865
c5fa4cf
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
b84d865
c5fa4cf
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
import os
import gradio as gr
import random as rand
import torch  
from huggingface_hub import InferenceClient
from sentence_transformers import SentenceTransformer


custom_css = """
@import url('https://fonts.googleapis.com/css2?family=Fredoka:wght@300..700&display=swap'); 
body, input, textarea, button, select, label, p, span, div, h1, h2, h3, h4, h5, h6 {
    font-family: 'Fredoka', sans-serif !important;
}
"""

with open("knowledgebase.txt", "r", encoding="utf-8") as file:
    knowledge_base = file.read()

def preprocess_text(text):
    cleaned_text = text.strip()
    chunks = cleaned_text.split("\n")
    cleaned_chunks = []
    for chunk in chunks:
        stripped_chunk = chunk.strip()
        if len(stripped_chunk) > 0:
            cleaned_chunks.append(stripped_chunk)
    return cleaned_chunks

cleaned_chunks = preprocess_text(knowledge_base)

model = SentenceTransformer('all-MiniLM-L6-v2')

def create_embeddings(text_chunks):
    chunk_embeddings = model.encode(text_chunks, convert_to_tensor=True)
    print("Embeddings Shape:", chunk_embeddings.shape)
    return chunk_embeddings

chunk_embeddings = create_embeddings(cleaned_chunks)

def get_top_chunks(query, chunk_embeddings, text_chunks):
    query_embedding = model.encode(query, convert_to_tensor=True)
    query_embedding_normalized = query_embedding / query_embedding.norm()
    chunk_embeddings_normalized = chunk_embeddings / chunk_embeddings.norm(dim=1, keepdim=True)

    similarities = torch.matmul(chunk_embeddings_normalized, query_embedding_normalized)
    top_indices = torch.topk(similarities, k=3).indices
    
    top_chunks = []
    for i in top_indices:
        chunk = text_chunks[i]
        top_chunks.append(chunk)
    return top_chunks


client = InferenceClient("Qwen/Qwen2.5-7B-Instruct", token = os.getenv("hf")) 

def respond(message, history):
    relevant_context = get_top_chunks(message, chunk_embeddings, cleaned_chunks)
    context_str = "\n".join(relevant_context)
    
    system_prompt = f"You are a helpful assistant. Use the following context to answer the user's question accurately:\n\n{context_str}"
    
    messages = [{"role": "system", "content": system_prompt}]
    
    if history:
        messages.extend(history)
        
    messages.append({"role": "user", "content": message})
    
    response = ""
    for msg in client.chat_completion(messages, max_tokens=200, temperature=1, top_p=0.5, stream=True):
        token = msg.choices[0].delta.content
        if token:
            response += token
            yield response


with gr.Blocks(theme='d8ahazard/material_design_rd', css=custom_css) as demo:
    gr.Markdown("# Digital Wellbeing")
    
    chatbot = gr.Chatbot(label="Chat History")
    msg = gr.Textbox(placeholder="Type your question here...")
    
    with gr.Row():
        submit_btn = gr.Button("Submit", variant="primary")
        clear_btn = gr.ClearButton([msg, chatbot], value="Clear Conversation")

    submit_event = submit_btn.click(respond, inputs=[msg, chatbot], outputs=[chatbot])
    msg_submit_event = msg.submit(respond, inputs=[msg, chatbot], outputs=[chatbot])
    
    submit_event.then(lambda: "", None, [msg])
    msg_submit_event.then(lambda: "", None, [msg])

if __name__ == "__main__":
    demo.launch()