Spaces:
Chz0
/
Running on Zero

File size: 6,274 Bytes
e29bfe0
b733698
 
15c28bf
b733698
 
 
15c28bf
 
b733698
4c61b56
b733698
 
 
4c61b56
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1e5f9f4
15c28bf
b733698
 
 
 
 
1e5f9f4
b733698
15c28bf
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1e5f9f4
15c28bf
98fb82a
aec92b4
b733698
15c28bf
 
 
 
 
1e5f9f4
b733698
 
 
 
 
15c28bf
 
b733698
1e5f9f4
b733698
 
15c28bf
b733698
 
 
 
 
 
 
 
 
 
15c28bf
b733698
1e5f9f4
 
 
 
 
 
 
 
 
 
 
 
 
 
 
d80c1ef
 
aec92b4
b733698
 
15c28bf
b733698
 
 
 
15c28bf
 
 
 
 
 
 
 
1e5f9f4
15c28bf
 
1e5f9f4
15c28bf
 
1e5f9f4
15c28bf
d80c1ef
1e5f9f4
15c28bf
b733698
15c28bf
b733698
 
15c28bf
 
 
 
 
 
 
 
 
1e5f9f4
 
b733698
 
 
8353a84
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
156
import spaces
import os
import torch
import random
import numpy as np
import pandas as pd
import gradio as gr
from pathlib import Path

from transformers import pipeline, set_seed
import transformers.utils.import_utils as transformers_import_utils
from sentence_transformers import SentenceTransformer
from datasets import load_dataset

# Transformers 5 may query the CUDA stream while preparing CPU attention
# masks (used internally when we call embedding_model.encode(...) on CPU).
# In the GPU-less ZeroGPU parent process, that stray CUDA runtime call
# leaves CUDA in a state that then breaks the next forked @spaces.GPU
# worker's attempt to initialize its dynamically assigned GPU, causing
# "RuntimeError: No CUDA GPUs are available" even though the code never
# explicitly touches CUDA outside the decorated function.
# Fix: skip only that query in the ZeroGPU parent; real GPU workers keep
# the original Transformers behavior once CUDA_VISIBLE_DEVICES is set.
# See: https://huggingface.co/spaces/zero-gpu-explorers/README/discussions/181
if hasattr(transformers_import_utils, "is_cuda_stream_capturing"):
    _original_is_cuda_stream_capturing = transformers_import_utils.is_cuda_stream_capturing

    def _zerogpu_safe_is_cuda_stream_capturing(*args, **kwargs):
        if os.getenv("SPACES_ZERO_GPU") == "1" and not os.getenv("CUDA_VISIBLE_DEVICES"):
            return False
        return _original_is_cuda_stream_capturing(*args, **kwargs)

    transformers_import_utils.is_cuda_stream_capturing = _zerogpu_safe_is_cuda_stream_capturing

# 1. Load Datasets
print("Loading datasets...")
dataset = load_dataset("cbasu/Med-EASi", split="train")
dataset2 = load_dataset("cbasu/Med-EASi", split="validation")
dataset3 = load_dataset("cbasu/Med-EASi", split="test")
dataset4 = load_dataset("csv", data_files="custom_terms.csv", split="train")

# 2. Format dataset rows into strings
texts = []
for row in dataset:
    content = f"Medical Context: {row['Expert']}\nSimplified Explanation: {row['Simple']}\nTerms: {row['expert_terms']}, {row['layman_terms']}"
    texts.append(content)
for row in dataset2:
    content = f"Medical Context: {row['Expert']}\nSimplified Explanation: {row['Simple']}\nTerms: {row['expert_terms']}, {row['layman_terms']}"
    texts.append(content)
for row in dataset3:
    content = f"Medical Context: {row['Expert']}\nSimplified Explanation: {row['Simple']}\nTerms: {row['expert_terms']}, {row['layman_terms']}"
    texts.append(content)
for row in dataset4:
    content = f"Medical Context: {row['Expert']}\nSimplified Explanation: {row['Simple']}\nTerms: {row['expert_terms']}, {row['layman_terms']}"
    texts.append(content)

print(f"Total processed documents: {len(texts)}")

# 3. Generate Embeddings & Index (this is retrieval-only, keep it on CPU)
print("Indexing dataset embeddings...")
embedding_model = SentenceTransformer('all-MiniLM-L6-v2', device='cpu')
embeddings = embedding_model.encode(texts)

df = pd.DataFrame({
    "Document": texts,
    "Embedding": list(embeddings)
})

# 4. Retrieval Helper Function
def retrieve_with_pandas(query, top_k=3):
    query_embedding = embedding_model.encode([query])[0]
    df['Similarity'] = df['Embedding'].apply(
        lambda x: np.dot(query_embedding, x) / (np.linalg.norm(query_embedding) * np.linalg.norm(x))
    )
    results = df.sort_values(by="Similarity", ascending=False).head(top_k)
    return results["Document"].tolist()

# 5. Prompt Builder Helper Function
def build_prompt(context_docs: list, user_query: str) -> str:
    context = "\n\n".join(context_docs)
    prompt = f"""<s>[INST] You are a helpful medical assistant.
Use the following retrieved medical context to answer the user's question clearly and concisely.
Explain it in very simple terms as if the user has the vocabulary of an elementary school student.
If the context does not contain enough information to answer, state that you don't know based on the provided text.
Answer in 2 to 3 sentences maximum.

Context:
{context}

Question:
{user_query} [/INST]"""
    return prompt

# 6. Load the LLM at MODULE LEVEL and place it on 'cuda' here.
#    Per HF's ZeroGPU docs, a CUDA emulation mode is active outside @spaces.GPU
#    functions specifically so this placement works at startup. Lazy-loading
#    onto cuda from inside a @spaces.GPU function is unsupported/discouraged.
print("Loading BioMistral LLM...")
llm = pipeline(
    'text-generation',
    model='BioMistral/BioMistral-7B',
    torch_dtype=torch.bfloat16,
    device_map="cuda",
)
if llm.model.config.pad_token_id is None:
    llm.model.config.pad_token_id = llm.model.config.eos_token_id

# 7. Answer Generation Helper Function - ONLY the real inference call needs @spaces.GPU
@spaces.GPU(duration=60)
def generate_answer(prompt: str, max_new_tokens: int = 80) -> str:
    set_seed(42)
    output = llm(
        prompt,
        max_new_tokens=max_new_tokens,
        num_return_sequences=1,
        do_sample=False,
    )[0]["generated_text"]

    # Strip out the prompt text to leave only the model's new response
    answer = output[len(prompt):].strip()
    return answer

# 8. RAG Handler Function
def respond_to_query(user_query: str) -> str:
    if not user_query.strip():
        return "Please enter a valid medical question or term."

    # Step A: Retrieve context documents based on user input
    context_docs = retrieve_with_pandas(user_query, top_k=3)

    # Step B: Build prompt with retrieved context
    prompt = build_prompt(context_docs, user_query)

    # Step C: Generate model answer
    answer = generate_answer(prompt, max_new_tokens=80)

    return answer

# 9. Gradio Web Interface
demo = gr.Interface(
    fn=respond_to_query,
    inputs=gr.Textbox(lines=2, placeholder="e.g., What is necrosis? or What is cholera?"),
    outputs=gr.Textbox(label="BioMistral RAG Response", lines=4),
    title="🩺 Medical Terminology RAG Assistant",
    description="Ask a medical question to retrieve context from Med-EASi + custom terms and generate simplified explanations.",
    examples=[
        ["What causes necrosis?"],
        ["What is cholera?"],
        ["What is asbestosis?"],
        ["How can syphilis be treated?"]
    ],
    cache_examples=False,
)

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