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()