Spaces:
Runtime error
Runtime error
File size: 5,776 Bytes
42d091d 344ab39 42d091d 344ab39 42d091d 344ab39 42d091d 5e0e499 42d091d | 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 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 | import gradio as gr
import torch
import urllib.request
import urllib.parse
import xml.etree.ElementTree as ET
from transformers import AutoModelForCausalLM, AutoTokenizer
from peft import PeftModel
import os
# Base model to use
BASE_MODEL_NAME = "google/gemma-2-2b"
# Check if a custom adapter is uploaded or default to base
ADAPTER_MODEL_NAME = os.environ.get("ADAPTER_MODEL_ID", "")
HF_TOKEN = os.environ.get("HF_TOKEN")
print("Initializing tokenizer...")
tokenizer = AutoTokenizer.from_pretrained(BASE_MODEL_NAME, token=HF_TOKEN)
if tokenizer.pad_token is None:
tokenizer.pad_token = tokenizer.eos_token
print("Initializing model...")
# Free spaces run on CPU. We load in float32 and use low_cpu_mem_usage to fit in CPU RAM.
model = AutoModelForCausalLM.from_pretrained(
BASE_MODEL_NAME,
torch_dtype=torch.float32,
low_cpu_mem_usage=True,
token=HF_TOKEN
)
if ADAPTER_MODEL_NAME:
print(f"Loading custom PEFT adapter: {ADAPTER_MODEL_NAME}...")
try:
model = PeftModel.from_pretrained(model, ADAPTER_MODEL_NAME)
except Exception as e:
print(f"Error loading PEFT adapter: {e}. Using base model instead.")
def fetch_arxiv_papers(topic, max_results=3):
try:
query = urllib.parse.quote(topic)
url = f"http://export.arxiv.org/api/query?search_query=all:{query}&max_results={max_results}"
req = urllib.request.Request(
url,
headers={'User-Agent': 'Mozilla/5.0'}
)
with urllib.request.urlopen(req, timeout=10) as response:
xml_data = response.read()
root = ET.fromstring(xml_data)
ns = {'atom': 'http://www.w3.org/2005/Atom'}
papers = []
for entry in root.findall('atom:entry', ns):
title_node = entry.find('atom:title', ns)
summary_node = entry.find('atom:summary', ns)
id_node = entry.find('atom:id', ns)
title = title_node.text.strip().replace('\n', ' ') if title_node is not None else "Unknown Title"
summary = summary_node.text.strip().replace('\n', ' ') if summary_node is not None else ""
id_url = id_node.text.strip() if id_node is not None else ""
authors = []
for author in entry.findall('atom:author', ns):
name_node = author.find('atom:name', ns)
if name_node is not None:
authors.append(name_node.text.strip())
papers.append({
'title': title,
'authors': authors,
'summary': summary,
'url': id_url
})
return papers
except Exception as e:
print(f"Error fetching from arXiv: {e}")
return []
def summarize_topic(topic):
if not topic.strip():
return "Please enter a valid topic.", ""
papers = fetch_arxiv_papers(topic, max_results=3)
if papers:
context_str = ""
for i, paper in enumerate(papers, 1):
context_str += f"Paper {i}: {paper['title']} by {', '.join(paper['authors'])}\nAbstract: {paper['summary']}\n\n"
formatted_prompt = f"Document:\nTopic: {topic}\n\nRelevant Literature:\n{context_str}Based on the above papers, provide key points and an overview of the research on this topic.\n\nSummary:\n"
else:
formatted_prompt = f"Document:\nTopic: {topic}\nProvide key points and overview of research papers associated with this topic.\n\nSummary:\n"
inputs = tokenizer(formatted_prompt, return_tensors="pt")
with torch.no_grad():
outputs = model.generate(
**inputs,
max_new_tokens=256,
do_sample=True,
temperature=0.7,
top_p=0.9,
repetition_penalty=1.1,
eos_token_id=tokenizer.eos_token_id,
pad_token_id=tokenizer.pad_token_id
)
decoded = tokenizer.decode(outputs[0], skip_special_tokens=True)
if "Summary:" in decoded:
response = decoded.split("Summary:")[-1].strip()
else:
response = decoded[len(formatted_prompt):].strip()
# Format References
refs_output = ""
if papers:
for i, paper in enumerate(papers, 1):
authors_str = ", ".join(paper['authors']) if paper['authors'] else "Unknown Authors"
refs_output += f"**[{i}] {paper['title']}**\n"
refs_output += f"*Authors:* {authors_str}\n"
refs_output += f"*URL:* [{paper['url']}]({paper['url']})\n\n"
else:
refs_output = "No external papers could be retrieved."
return response, refs_output
# Create Gradio interface
with gr.Blocks(theme=gr.themes.Soft()) as demo:
gr.Markdown("# 🔬 Scientific Research Paper Summarizer")
gr.Markdown(
"Enter any scientific topic or research question. The assistant will fetch relevant papers from "
"arXiv in real-time, synthesize the core findings using our fine-tuned Gemma model, and provide direct references."
)
with gr.Row():
with gr.Column():
topic_input = gr.Textbox(
label="Research Topic / Keywords",
placeholder="e.g. quantum machine learning, artificial intelligence in healthcare..."
)
submit_btn = gr.Button("Generate Overview", variant="primary")
with gr.Column():
answer_output = gr.Textbox(label="Model Overview & Points", interactive=False)
refs_output = gr.Markdown(label="References Cited")
submit_btn.click(
fn=summarize_topic,
inputs=[topic_input],
outputs=[answer_output, refs_output]
)
if __name__ == "__main__":
demo.launch()
|