ginnigarg's picture
tried fixing tensor error
5b83b4e verified
Raw
History Blame Contribute Delete
3.86 kB
import streamlit as st
from transformers import pipeline
from sentence_transformers import SentenceTransformer
import faiss
import numpy as np
import json
import re
import pkg_resources
from symspellpy import SymSpell
# -----------------------------
# 1. Load Portfolio Data
# -----------------------------
with open("portfolio.json", "r") as f:
portfolio_data = json.load(f)
# Extract texts only (documents to embed)
documents = [entry["text"] for entry in portfolio_data]
# -----------------------------
# 2. Initialize SymSpell
# -----------------------------
sym_spell = SymSpell(max_dictionary_edit_distance=2, prefix_length=7)
# Load dictionary
dictionary_path = pkg_resources.resource_filename(
"symspellpy", "frequency_dictionary_en_82_765.txt"
)
sym_spell.load_dictionary(dictionary_path, term_index=0, count_index=1)
# If you have a custom dict (else comment out)
# sym_spell.load_dictionary("custom_dictionary.txt", term_index=0, count_index=1)
# -----------------------------
# 3. QA Model (Generative)
# -----------------------------
from transformers import AutoTokenizer, AutoModelForSeq2SeqLM, pipeline
model_name = "google/flan-t5-large"
tokenizer = AutoTokenizer.from_pretrained(model_name)
model = AutoModelForSeq2SeqLM.from_pretrained(
model_name,
device_map=None, # don’t use meta-tensors
low_cpu_mem_usage=False # force eager weight loading
)
qa_pipeline = pipeline("text2text-generation", model=model, tokenizer=tokenizer, device=-1)
# qa_pipeline = pipeline("text2text-generation", model="google/flan-t5-large", device=-1)
# -----------------------------
# 4. Embed Portfolio Documents
# -----------------------------
@st.cache_resource
def load_data():
model = SentenceTransformer("all-MiniLM-L6-v2")
doc_embeddings = model.encode(documents, convert_to_numpy=True)
index = faiss.IndexHNSWFlat(doc_embeddings.shape[1], 32)
index.add(np.array(doc_embeddings).astype("float32"))
return model, index
model, hnsw_index = load_data()
# -----------------------------
# 5. Preprocessing Functions
# -----------------------------
def preprocess_text(text):
tokens = re.findall(r'\w+|\d+\w*|\S+', text)
return tokens
def correct_spelling(text):
tokens = preprocess_text(text)
corrected_tokens = []
for token in tokens:
if token.isdigit() or re.match(r'\d+\w*', token) or re.match(r'[.,]', token):
corrected_tokens.append(token)
else:
suggestions = sym_spell.lookup(token, max_edit_distance=2, verbosity=1)
if suggestions:
corrected_tokens.append(suggestions[0].term)
else:
corrected_tokens.append(token)
return " ".join(corrected_tokens)
# -----------------------------
# 6. RAG Function
# -----------------------------
def rag_qa(question, k=3):
question = correct_spelling(question)
question_embedding = model.encode([question], convert_to_numpy=True)
distances, retrieved_indices = hnsw_index.search(
np.array(question_embedding).astype("float32"), k
)
retrieved_contexts = [documents[idx] for idx in retrieved_indices[0]]
# Build context
context = "\n".join(retrieved_contexts)
prompt = (
f"Context:\n{context}\n\n"
f"Q: {question}\n"
f"A: If the answer is unclear from the context, respond with 'I don't know'."
)
response = qa_pipeline(prompt, max_length=500)
return response[0]['generated_text']
# -----------------------------
# 7. Streamlit UI
# -----------------------------
st.title("🧠 Ask Anything About Ginni!")
question = st.text_input("Ask your question:")
if st.button("Get Answer"):
if question.strip():
answer = rag_qa(question)
st.success(f"**Answer:** {answer}")
else:
st.warning("Please enter a valid question.")