ssdataanalysis's picture
Upload folder using huggingface_hub
a34947b verified
Raw
History Blame Contribute Delete
8.7 kB
import os
import json
from concurrent.futures import ThreadPoolExecutor
import gradio as gr
import spaces
from huggingface_hub import InferenceClient
@spaces.GPU
def _ensure_gpu_functions_present_for_zerogpu():
"""ZeroGPU startup probe: required for zero-a10g spaces detection."""
return None
MODEL_MATRIX = [
{
"key": "gemma4-31b",
"label": "Gemma 4 31B-IT",
"model_id": "unsloth/gemma-4-31B-it-unsloth-bnb-4bit",
},
{
"key": "gemma4-26b",
"label": "Gemma 4 26B-IT",
"model_id": "unsloth/gemma-4-26B-A4B-it",
},
{
"key": "qwen3.6-35b",
"label": "Qwen 3.6 35B-IT",
"model_id": "unsloth/Qwen3.6-35B-A3B-GGUF",
},
{
"key": "qwen3.6-27b",
"label": "Qwen 3.6 27B",
"model_id": "unsloth/Qwen3.6-27B-GGUF",
},
]
HF_TOKEN = os.getenv("HF_API_TOKEN") or os.getenv("HUGGINGFACE_HUB_TOKEN") or os.getenv("HF_TOKEN")
HF_TIMEOUT = int(os.getenv("HF_TIMEOUT", "120"))
_SYSTEM_PROMPT = (
"You are a useful assistant. Answer briefly and directly unless the user asks for "
"long-form detail. Keep responses focused and practical."
)
def _init_client() -> InferenceClient:
return InferenceClient(token=HF_TOKEN)
def _extract_chat_text(response):
if response is None:
return "No response from model."
if isinstance(response, dict):
if "choices" in response and response["choices"]:
choice = response["choices"][0]
if isinstance(choice, dict) and "message" in choice and isinstance(choice["message"], dict):
return (choice["message"].get("content") or "").strip()
if isinstance(choice, dict) and "text" in choice:
return str(choice["text"]).strip()
if "generated_text" in response:
return str(response["generated_text"]).strip()
if "text" in response:
return str(response["text"]).strip()
if hasattr(response, "choices"):
choices = getattr(response, "choices")
if choices:
first = choices[0]
if hasattr(first, "message") and getattr(first.message, "content", None) is not None:
return str(first.message.content).strip()
if hasattr(first, "text"):
return str(first.text).strip()
text = str(response)
if text and text != "{}":
return text.strip()
return "Model returned an empty response."
def _build_messages(history, user_message):
messages = [{"role": "system", "content": _SYSTEM_PROMPT}]
for pair in history:
if not pair:
continue
user_turn, assistant_turn = pair
if user_turn:
messages.append({"role": "user", "content": str(user_turn)})
if assistant_turn:
messages.append({"role": "assistant", "content": str(assistant_turn)})
messages.append({"role": "user", "content": user_message})
return messages
def _chat_completion(model_id, messages, max_new_tokens, temperature, top_p):
client = _init_client()
try:
response = client.chat_completion(
model=model_id,
messages=messages,
max_tokens=max_new_tokens,
temperature=temperature,
top_p=top_p,
timeout=HF_TIMEOUT,
)
return _extract_chat_text(response)
except Exception as chat_error:
fallback_prompt = "\n".join([f"{m['role']}: {m['content']}" for m in messages]) + "\nassistant:"
try:
response = client.text_generation(
prompt=fallback_prompt,
max_new_tokens=max_new_tokens,
temperature=temperature,
top_p=top_p,
timeout=HF_TIMEOUT,
)
return str(response).strip()
except Exception as text_error:
return (
f"Could not query model '{model_id}'."
f" Chat error: {chat_error.__class__.__name__}"
f"\nFallback error: {text_error.__class__.__name__}"
)
def _safe_append(history, user_message, assistant_message):
updated = list(history or [])
updated.append((user_message, assistant_message))
return updated
def _query_model(model_id, messages, max_new_tokens, temperature, top_p):
return _chat_completion(model_id, messages, max_new_tokens, temperature, top_p)
def run_compare(
user_message,
hist1,
hist2,
hist3,
hist4,
temperature,
top_p,
max_tokens,
):
if not user_message or not str(user_message).strip():
return hist1, hist2, hist3, hist4, "", "", "", ""
user_message = str(user_message).strip()
m1, m2, m3, m4 = MODEL_MATRIX
model_ids = [m1["model_id"], m2["model_id"], m3["model_id"], m4["model_id"]]
histories = [
list(hist1 or []),
list(hist2 or []),
list(hist3 or []),
list(hist4 or []),
]
messages = [
_build_messages(histories[i], user_message)
for i in range(4)
]
with ThreadPoolExecutor(max_workers=4) as executor:
futures = [
executor.submit(_query_model, model_ids[idx], messages[idx], max_tokens, temperature, top_p)
for idx in range(4)
]
answers = [f.result() for f in futures]
new_hist1 = _safe_append(histories[0], user_message, answers[0])
new_hist2 = _safe_append(histories[1], user_message, answers[1])
new_hist3 = _safe_append(histories[2], user_message, answers[2])
new_hist4 = _safe_append(histories[3], user_message, answers[3])
return (
new_hist1,
new_hist2,
new_hist3,
new_hist4,
m1["label"],
m2["label"],
m3["label"],
m4["label"],
)
def run_clear():
return [], [], [], [], "", "", "", ""
css = """
.gradio-container { max-width: 100%; }
.column {
border: 1px solid #d7dee8;
border-radius: 12px;
padding: 10px;
background: linear-gradient(160deg, #f8fbff, #f2f4ff);
}
"""
with gr.Blocks(css=css) as demo:
gr.Markdown(
"# Side-by-side Unsloth model comparison\n"
"All generations run through Hugging Face hosted inference (zero local GPU)."
)
with gr.Row():
status1 = gr.Textbox(label="Gemma4 31B-IT", value="", interactive=False)
status2 = gr.Textbox(label="Gemma4 26B-IT", value="", interactive=False)
status3 = gr.Textbox(label="Qwen 3.6 35B-IT", value="", interactive=False)
status4 = gr.Textbox(label="Qwen 3.6 27B", value="", interactive=False)
with gr.Row(equal_height=True):
with gr.Column(elem_classes=["column"]):
gr.Markdown("### Gemma4 31B-IT")
chat1 = gr.Chatbot(height=470, label="Gemma4 31B-IT")
with gr.Column(elem_classes=["column"]):
gr.Markdown("### Gemma4 26B-IT")
chat2 = gr.Chatbot(height=470, label="Gemma4 26B-IT")
with gr.Column(elem_classes=["column"]):
gr.Markdown("### Qwen 3.6 35B-IT")
chat3 = gr.Chatbot(height=470, label="Qwen 3.6 35B-IT")
with gr.Column(elem_classes=["column"]):
gr.Markdown("### Qwen 3.6 27B")
chat4 = gr.Chatbot(height=470, label="Qwen 3.6 27B")
with gr.Accordion("Advanced generation settings", open=False):
temperature = gr.Slider(0.1, 1.2, value=0.7, step=0.05, label="Temperature")
top_p = gr.Slider(0.2, 1.0, value=0.9, step=0.05, label="Top-p")
max_tokens = gr.Slider(64, 2048, value=512, step=16, label="Max new tokens")
with gr.Row():
prompt = gr.Textbox(
label="User prompt",
placeholder="Ask the same question to all models...",
lines=2,
scale=3,
)
send = gr.Button("Compare", variant="primary")
clear = gr.Button("Clear")
send.click(
run_compare,
inputs=[
prompt,
chat1,
chat2,
chat3,
chat4,
temperature,
top_p,
max_tokens,
],
outputs=[chat1, chat2, chat3, chat4, status1, status2, status3, status4],
)
prompt.submit(
run_compare,
inputs=[
prompt,
chat1,
chat2,
chat3,
chat4,
temperature,
top_p,
max_tokens,
],
outputs=[chat1, chat2, chat3, chat4, status1, status2, status3, status4],
)
clear.click(
run_clear,
outputs=[chat1, chat2, chat3, chat4, status1, status2, status3, status4],
)
if __name__ == "__main__":
demo.queue()
demo.launch()