Amey9766's picture
Update app.py
8602cb1 verified
Raw
History Blame Contribute Delete
5.13 kB
import os
import re
import threading
from typing import List, Dict, Optional
import gradio as gr
import torch
from transformers import AutoTokenizer, AutoModelForCausalLM, TextIteratorStreamer
MODEL_ID = "Amey9766/qwen-0.6b-hospitality-housekeeping"
tokenizer = None
model = None
device = None
def load_model(hf_access_token: Optional[str] = None):
global tokenizer, model, device
if tokenizer is not None and model is not None:
return
device = "cuda" if torch.cuda.is_available() else "cpu"
dtype = torch.float16 if device == "cuda" else torch.float32
use_token = hf_access_token or os.getenv("HF_TOKEN")
tokenizer = AutoTokenizer.from_pretrained(
MODEL_ID,
token=use_token,
trust_remote_code=True,
use_fast=True,
)
model = AutoModelForCausalLM.from_pretrained(
MODEL_ID,
token=use_token,
torch_dtype=dtype,
device_map="auto" if device == "cuda" else None,
trust_remote_code=True,
)
if device == "cpu":
model.to(device)
print("✅ Loaded model:", getattr(model.config, "_name_or_path", "unknown"))
def build_plain_prompt(system_message: str, history: List[Dict[str, str]], user_message: str) -> str:
"""
Universal prompt builder that does NOT require tokenizer.chat_template.
This works with any CausalLM.
"""
hard_rules = (
"You are a professional hotel housekeeping assistant.\n"
"STRICT RULES:\n"
"1) Respond in English only.\n"
"2) Answer ONLY the user's last question.\n"
"3) Do NOT generate follow-up questions.\n"
"4) Do NOT mention rules, instructions, or your reasoning.\n"
"5) Provide only the final answer.\n"
)
sys = (system_message or "").strip()
prompt = f"{hard_rules}\nSYSTEM NOTE: {sys}\n\n"
# Convert Gradio "messages" history into a readable transcript
for m in history:
role = (m.get("role") or "user").lower()
content = (m.get("content") or "").strip()
if not content:
continue
if role == "user":
prompt += f"User: {content}\n"
else:
prompt += f"Assistant: {content}\n"
prompt += f"User: {user_message.strip()}\nAssistant:"
return prompt
def clean_output(text: str) -> str:
"""
Removes common fine-tune artifacts without being too aggressive.
"""
# Remove leading parenthetical meta like: "(Answering in English...)"
text = re.sub(r"^\s*\(.*?\)\s*", "", text, flags=re.DOTALL)
# If the model starts adding "Question:" sections, cut everything after it
cut_markers = [
"\nQuestion:",
"\nNow, let's",
"\nNow let's",
"\nBased on the rules",
"\nAccording to the rules",
]
for marker in cut_markers:
idx = text.lower().find(marker.lower())
if idx != -1:
text = text[:idx].strip()
break
return text.strip()
def respond(
message: str,
history: List[Dict[str, str]],
system_message: str,
max_tokens: int,
temperature: float,
top_p: float,
hf_token: gr.OAuthToken,
):
load_model(hf_token.token if hf_token else None)
prompt = build_plain_prompt(system_message, history, message)
inputs = tokenizer(prompt, return_tensors="pt")
inputs = {k: v.to(device) for k, v in inputs.items()}
streamer = TextIteratorStreamer(
tokenizer,
skip_prompt=True,
skip_special_tokens=True,
)
# Lower randomness to reduce “roleplay / training artifact” behavior
gen_kwargs = dict(
**inputs,
max_new_tokens=int(max_tokens),
do_sample=True,
temperature=float(temperature),
top_p=float(top_p),
repetition_penalty=1.15,
streamer=streamer,
eos_token_id=tokenizer.eos_token_id,
pad_token_id=tokenizer.eos_token_id,
)
thread = threading.Thread(target=model.generate, kwargs=gen_kwargs)
thread.start()
out = ""
for chunk in streamer:
out += chunk
# Stream the cleaned output live
cleaned = clean_output(out)
# Hard stop if it starts self-questioning
if any(x in out.lower() for x in ["\nquestion:", "now, let's generate", "based on the rules"]):
yield cleaned
break
yield cleaned
chatbot = gr.ChatInterface(
respond,
type="messages",
title="Hospitality Housekeeping Assistant",
description=f"Running model: `{MODEL_ID}`",
additional_inputs=[
gr.Textbox(
value="Give SOP-style housekeeping answers. Use bullet points when helpful.",
label="System message",
),
gr.Slider(1, 1024, value=256, step=1, label="Max new tokens"),
gr.Slider(0.1, 1.0, value=0.3, step=0.05, label="Temperature"),
gr.Slider(0.5, 1.0, value=0.85, step=0.05, label="Top-p"),
],
)
with gr.Blocks() as demo:
with gr.Sidebar():
gr.LoginButton()
gr.Markdown(f"**Model:** `{MODEL_ID}`")
chatbot.render()
if __name__ == "__main__":
demo.launch()