RAG_Paper / app.py
praneth02's picture
Update app.py
cdf96cb verified
Raw
History Blame Contribute Delete
13 kB
import re
import time
import torch
from sentence_transformers import SentenceTransformer
import chromadb
import google.generativeai as genai
import gradio as gr
# Load embedding model with GPU support (optimized for QA tasks)
embedder = SentenceTransformer('multi-qa-mpnet-base-dot-v1', device='cuda' if torch.cuda.is_available() else 'cpu')
# Configure Gemini API
genai.configure(api_key="AIzaSyDXG4o4UnII5VFD1u5TaWgleG2kCfJ0Ofw")
# Initialize Gemini model
gemini_instance = genai.GenerativeModel('gemini-2.0-flash')
# Define document paths and labels (using .md files)
doc_paths = [
'1002215.md',
'cancers-15-00321.md',
'ijo-57-06-1245.md'
]
doc_labels = [
"Early-stage triple negative breast cancer: the therapeutic role of immunotherapy and the prognostic value of pathological complete response",
"Immunotherapy for Triple-Negative Breast Cancer: Combination Strategies to Improve Outcome",
"Triple‑negative breast cancer therapy: Current and future perspectives (Review)"
]
def extract_and_chunk_docs(doc_paths, doc_labels):
segmented_docs = []
doc_info = []
for doc_path, doc_label in zip(doc_paths, doc_labels):
try:
with open(doc_path, 'r', encoding='utf-8') as md_file:
full_text = md_file.read()
if not full_text.strip():
print(f"No content extracted from {doc_label}")
segmented_docs.append([])
doc_info.append({"label": doc_label, "gemini_structure": "No content extracted"})
continue
# Identify pages based on '-----'
pages = full_text.split('-----')
full_text_lines = []
for page_num, page_content in enumerate(pages, 1):
lines = page_content.split('\n')
for line in lines:
line = line.strip()
if line:
full_text_lines.append({'text': line, 'page': page_num})
text_for_gemini = "\n".join([entry['text'] for entry in full_text_lines])
prompt = f"""You are an expert in analyzing research papers. Given the following text from a Markdown file, identify all potential section titles (e.g., Abstract, Introduction, Methods, Results, Discussion) and subsections (e.g., '2.1 Data Analysis'). Include headings that might define or characterize triple-negative breast cancer (TNBC), such as 'Definition,' 'Characteristics,' or similar, and explicitly include 'References' as a section title if present. Only headings starting with '## **' are sections. Return only the titles (without '## **'), one per line, without explanation.
Text:
{text_for_gemini}
"""
max_attempts = 5
for attempt in range(max_attempts):
try:
response = gemini_instance.generate_content(prompt)
titles = response.text.strip().split('\n')
break
except Exception as e:
if "429" in str(e):
delay = 2 ** attempt
print(f"Rate limit hit for {doc_label}. Waiting {delay} seconds...")
time.sleep(delay)
else:
print(f"Error identifying titles for {doc_label}: {e}")
segmented_docs.append([])
doc_info.append({"label": doc_label, "gemini_structure": f"Error: {e}"})
break
else:
print(f"Error for {doc_label}: Max retries exceeded.")
segmented_docs.append([])
doc_info.append({"label": doc_label, "gemini_structure": "Max retries exceeded"})
continue
titles = list(dict.fromkeys([title.strip() for title in titles if title.strip()]))
print(f"Gemini Identified Titles for {doc_label}: {titles}")
chunks = []
current_chunk = ""
current_title = "Unknown"
for line_info in full_text_lines:
line = line_info['text']
page_num = line_info['page'] # Use page number from full_text_lines
cleaned_line = re.sub(r'(?i)copyright.*|all\s*rights\s*reserved', '', line)
cleaned_line = re.sub(r'\s+', ' ', cleaned_line).strip()
if not cleaned_line:
continue
if cleaned_line.startswith('## **'):
normalized_line = re.sub(r'^##\s*\*\*|\*\*', '', cleaned_line).strip()
matched_title = next((title for title in titles if normalized_line.lower() == title.lower()), None)
if matched_title:
# Save the previous chunk if it exists and is not "References"
if current_chunk and current_title.lower() != "references":
chunks.append({
'text': current_chunk.strip(),
'page': page_num, # Assign the current page
'section': current_title
})
current_title = matched_title
current_chunk = ""
print(f"Detected title on page {page_num}: '{current_title}'")
continue
# Only add to chunk if current section is not "References"
if current_title.lower() != "references":
current_chunk += " " + cleaned_line
if len(current_chunk.split()) > 100:
chunks.append({
'text': current_chunk.strip(),
'page': page_num, # Assign the current page
'section': current_title
})
current_chunk = ""
# Save the final chunk if it exists and is not "References"
if current_chunk and current_title.lower() != "references":
chunks.append({
'text': current_chunk.strip(),
'page': page_num, # Assign the final page
'section': current_title
})
segmented_docs.append(chunks)
doc_info.append({"label": doc_label, "gemini_structure": "Chunked by titles starting with '## **', excluding 'References'"})
print(f"\n=== {doc_label} ===")
print("Identified Section Titles:", titles)
print("\nChunks (excluding References):")
for i, chunk in enumerate(chunks):
print(f"Chunk {i + 1}: [Page: {chunk['page']}, Section: '{chunk['section']}'] {chunk['text'][:100]}...")
except Exception as e:
print(f"Error processing {doc_label}: {e}")
segmented_docs.append([])
doc_info.append({"label": doc_label, "gemini_structure": f"Error: {str(e)}"})
return segmented_docs, doc_info
def compute_segment_embeddings(segmented_docs):
segment_embeddings = []
for doc_segments in segmented_docs:
if doc_segments:
embeddings = embedder.encode(
[seg['text'] for seg in doc_segments],
convert_to_tensor=False,
show_progress_bar=True,
batch_size=32
)
segment_embeddings.append(embeddings)
else:
segment_embeddings.append([])
return segment_embeddings
def save_to_vector_store(segmented_docs, segment_embeddings, doc_labels):
db_instance = chromadb.Client()
try:
embeddings_store = db_instance.create_collection("research_docs_mpnet_run_b")
except:
embeddings_store = db_instance.get_collection("research_docs_mpnet_run_b")
for i, (doc_segments, doc_embeds) in enumerate(zip(segmented_docs, segment_embeddings)):
if doc_segments and doc_embeds.size > 0:
for j, (segment, embed) in enumerate(zip(doc_segments, doc_embeds)):
embeddings_store.add(
embeddings=[embed.tolist()],
documents=[segment['text']],
metadatas=[{
"label": doc_labels[i],
"page": segment['page'], # Include page in metadata
"section": segment['section']
}],
ids=[f"{doc_labels[i]}seg{j}"]
)
return embeddings_store
def process_query(query, embeddings_store, doc_labels):
query_embed = embedder.encode([query], convert_to_tensor=False)[0].tolist()
query_results = embeddings_store.query(query_embeddings=[query_embed], n_results=3) # Limit to top 3
retrieved_contexts = []
ref_citations = []
for doc, meta in zip(query_results["documents"][0], query_results["metadatas"][0]):
label = meta["label"]
page = meta["page"]
section = meta["section"]
retrieved_contexts.append(doc)
ref_citations.append(f"[Ref: {label}, Page: {page}, Section: '{section}']")
combined_context = "\n".join(retrieved_contexts) if retrieved_contexts else "No relevant context found."
citation_str = " | ".join(ref_citations) if ref_citations else "N/A"
answer_prompt = f"""You are an AI assistant for research papers. Use only the provided context to answer the query concisely (1-2 sentences max). If the context lacks a clear answer, state so briefly.
Context:
{combined_context}
Query: {query}
Answer:"""
max_attempts = 5
for attempt in range(max_attempts):
try:
response = gemini_instance.generate_content(answer_prompt)
response_text = response.text.strip()
break
except Exception as e:
if "429" in str(e):
delay = 2 ** attempt
print(f"Rate limit for query '{query}'. Waiting {delay} seconds...")
time.sleep(delay)
else:
response_text = f"Error generating answer: {e}"
break
else:
response_text = "Error: Max retries exceeded."
final_response = f"{response_text}\n\n*References*: {citation_str}"
return final_response
def chatbot_response(message, history):
response = process_query(message, embeddings_store, doc_labels)
return history + [{"role": "user", "content": message}, {"role": "assistant", "content": response}]
# Custom theme (unchanged)
custom_theme = gr.themes.Soft(
primary_hue="blue",
secondary_hue="gray",
neutral_hue="slate",
text_size="lg",
spacing_size="md",
radius_size="lg",
).set(
body_background_fill="#f0f4f8",
body_text_color="#1e293b",
input_background_fill="#ffffff",
input_border_color="#cbd5e1",
input_shadow="0 2px 4px rgba(0,0,0,0.1)",
button_primary_background_fill="#3b82f6",
button_primary_text_color="#ffffff",
button_primary_background_fill_hover="#2563eb",
block_title_text_color="#1e40af",
block_border_color="#e2e8f0",
block_background_fill="#ffffff",
)
# Initialize system
doc_segments, doc_metadata = extract_and_chunk_docs(doc_paths, doc_labels)
segment_embeds = compute_segment_embeddings(doc_segments)
embeddings_store = save_to_vector_store(doc_segments, segment_embeds, doc_labels)
# Custom CSS (simplified, no debug styling)
css = """
.header { text-align: center; margin-bottom: 20px; }
.gradio-container { max-width: 900px; margin: auto; }
.chatbot .prose { max-width: 100%; }
.chatbot .bubble-wrap:nth-child(even) { background-color: #dbeafe; color: #1e40af; }
.chatbot .bubble-wrap:nth-child(odd) { background-color: #f1f5f9; color: #1e293b; }
.chatbot .bubble { border-radius: 10px; padding: 10px; }
"""
# Gradio interface
with gr.Blocks(theme=custom_theme, css=css, title="Research Paper Chatbot") as interface:
gr.Markdown("# Research Paper Chatbot\nAsk questions about the research papers and get concise, referenced answers.", elem_classes="header")
chatbot = gr.Chatbot(label="Conversation", height=500, type="messages", avatar_images=(None, "https://png.pngtree.com/png-vector/20201224/ourmid/pngtree-future-intelligent-technology-robot-ai-png-image_2588803.jpg"))
with gr.Row():
with gr.Column(scale=8):
msg = gr.Textbox(placeholder="Type your question here", show_label=False, container=False)
with gr.Column(scale=2):
submit_btn = gr.Button("Send", variant="primary")
msg.submit(chatbot_response, [msg, chatbot], chatbot).then(lambda: "", None, msg)
submit_btn.click(chatbot_response, [msg, chatbot], chatbot).then(lambda: "", None, msg)
if __name__ == "__main__":
interface.launch()