import spaces # MUST be imported before torch import torch import gradio as gr from transformers import AutoTokenizer, AutoModelForCausalLM from peft import PeftModel # ============================================================ # Configuration # ============================================================ BASE_MODEL = "Qwen/Qwen2.5-1.5B-Instruct" ADAPTER_MODEL = "rolmaxx/MediGuide-QLoRA" SYSTEM_PROMPT = """You are MediGuide, a medical conversational assistant. Provide clear, helpful, and appropriately cautious responses. Do not claim to replace a qualified healthcare professional. If symptoms may indicate an emergency, advise the user to seek appropriate professional medical care.""" # ============================================================ # Tokenizer & Model – loaded on CPU (no GPU yet) # ============================================================ print("Loading tokenizer...") tokenizer = AutoTokenizer.from_pretrained(BASE_MODEL) if tokenizer.pad_token is None: tokenizer.pad_token = tokenizer.eos_token print("Loading base model on CPU...") model = AutoModelForCausalLM.from_pretrained( BASE_MODEL, torch_dtype=torch.float16, device_map="cpu", # force CPU – GPU will be used later inside the decorated function ) print("Base model loaded on CPU.") print("Loading MediGuide QLoRA adapter on CPU...") model = PeftModel.from_pretrained( model, ADAPTER_MODEL, torch_device="cpu", # load adapter on CPU ) model = model.merge_and_unload() # merge for faster inference model.eval() print("MediGuide loaded successfully (on CPU).") # ============================================================ # Generation – GPU is requested via @spaces.GPU # ============================================================ @spaces.GPU(duration=60) def generate_response(message, history): """ Generate a response from the model. history: list of [user_msg, assistant_msg] tuples from Gradio. """ # 1. Sanitise inputs if not isinstance(message, str): message = str(message) if message is not None else "" messages = [{"role": "system", "content": SYSTEM_PROMPT}] for turn in history: if not isinstance(turn, (list, tuple)) or len(turn) < 2: continue user_msg, assistant_msg = turn[0], turn[1] if user_msg is not None and str(user_msg).strip(): messages.append({"role": "user", "content": str(user_msg)}) if assistant_msg is not None and str(assistant_msg).strip(): messages.append({"role": "assistant", "content": str(assistant_msg)}) if message.strip(): messages.append({"role": "user", "content": message.strip()}) else: return "Please enter a valid question." # 2. Move model to GPU (now available thanks to @spaces.GPU) model.to("cuda") # 3. Apply chat template prompt = tokenizer.apply_chat_template( messages, tokenize=False, add_generation_prompt=True, ) # 4. Tokenize inputs = tokenizer( prompt, return_tensors="pt", truncation=True, max_length=4096, ) inputs = {k: v.to("cuda") for k, v in inputs.items()} # 5. Generate 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, pad_token_id=tokenizer.eos_token_id, ) # 6. Decode only new tokens generated_tokens = outputs[0][inputs["input_ids"].shape[-1]:] response = tokenizer.decode(generated_tokens, skip_special_tokens=True) # 7. Move model back to CPU to free GPU memory for other requests model.to("cpu") torch.cuda.empty_cache() # optional, helps release memory return response.strip() # ============================================================ # Gradio UI # ============================================================ with gr.Blocks(title="MediGuide") as demo: gr.Markdown( """ # 🩺 MediGuide ### QLoRA Fine‑Tuned Medical Conversational Assistant MediGuide is based on **Qwen2.5‑1.5B‑Instruct** and fine‑tuned using **QLoRA** on medical dialogue data. ⚠️ **Disclaimer** This demo is for research and educational purposes only. It is **not a substitute** for professional medical advice, diagnosis, or treatment. """ ) chatbot = gr.Chatbot(label="MediGuide", height=500) message = gr.Textbox( label="Your question", placeholder="e.g. What are common symptoms of the flu?", lines=2, ) with gr.Row(): send_btn = gr.Button("Send", variant="primary") clear_btn = gr.Button("Clear") def respond(message, history): response = generate_response(message, history) history = history + [(message, response)] return "", history message.submit(respond, [message, chatbot], [message, chatbot]) send_btn.click(respond, [message, chatbot], [message, chatbot]) clear_btn.click(lambda: [], outputs=chatbot) # ============================================================ # Launch # ============================================================ demo.queue().launch()