ai-chat / app.py
CreitinGameplays's picture
Update app.py
e973aca verified
Raw
History Blame Contribute Delete
8.7 kB
import os
from threading import Thread
from typing import Iterator
import spaces
import gradio as gr
import torch
from transformers import (
AutoModelForCausalLM,
AutoTokenizer,
BitsAndBytesConfig,
TextIteratorStreamer,
)
# ============================================================
# CONFIG
# ============================================================
MODEL_ID = "CreitinGameplays/GLM-4.7-Flash-Fable-5-Distill"
# Keep these conservative for ZeroGPU.
MAX_NEW_TOKENS = 1024
DEFAULT_MAX_NEW_TOKENS = 4096
# Smaller context = much faster prompt processing.
MAX_INPUT_TOKEN_LENGTH = 6112
DEFAULT_SYSTEM_PROMPT = (
"Answer clearly and directly. "
"Use two or three sentences unless the user asks "
"for a longer response."
)
# ============================================================
# TOKENIZER
# ============================================================
print(f"Loading tokenizer: {MODEL_ID}")
tokenizer = AutoTokenizer.from_pretrained(
MODEL_ID,
padding_side="left",
)
if hasattr(tokenizer, "use_default_system_prompt"):
tokenizer.use_default_system_prompt = False
# ============================================================
# MODEL
# ============================================================
print(f"Loading model: {MODEL_ID}")
quantization_config = BitsAndBytesConfig(
load_in_8bit=True,
)
model = AutoModelForCausalLM.from_pretrained(
MODEL_ID,
quantization_config=quantization_config,
torch_dtype=torch.bfloat16,
device_map={"": "cuda"},
low_cpu_mem_usage=True,
)
model.eval()
print("Model loaded.")
print(f"CUDA: {torch.cuda.is_available()}")
# ============================================================
# CONVERSATION
# ============================================================
def build_conversation(
message: str,
chat_history,
system_prompt: str,
):
messages = []
if system_prompt and system_prompt.strip():
messages.append(
{
"role": "system",
"content": system_prompt.strip(),
}
)
if chat_history:
for item in chat_history:
if not isinstance(item, (list, tuple)):
continue
if len(item) != 2:
continue
user_message, assistant_message = item
if user_message:
messages.append(
{
"role": "user",
"content": str(user_message),
}
)
if assistant_message:
messages.append(
{
"role": "assistant",
"content": str(assistant_message),
}
)
messages.append(
{
"role": "user",
"content": message,
}
)
return messages
# ============================================================
# DYNAMIC GPU DURATION
# ============================================================
def get_gpu_duration(
message: str,
chat_history,
system_prompt: str,
max_new_tokens: int,
temperature: float,
top_p: float,
top_k: int,
repetition_penalty: float,
):
"""
Estimate a GPU reservation based on requested output length.
The goal is to avoid reserving an unnecessarily large GPU
window while still giving longer answers enough time.
"""
tokens = int(max_new_tokens)
if tokens <= 128:
return 30
if tokens <= 256:
return 45
if tokens <= 384:
return 60
# Do not request enormous durations by default.
return 60
# ============================================================
# GENERATION
# ============================================================
@spaces.GPU(duration=get_gpu_duration)
def generate(
message: str,
chat_history,
system_prompt: str = DEFAULT_SYSTEM_PROMPT,
max_new_tokens: int = DEFAULT_MAX_NEW_TOKENS,
temperature: float = 0.5,
top_p: float = 0.95,
top_k: int = 40,
repetition_penalty: float = 1.0,
) -> Iterator[str]:
if not message or not message.strip():
yield "Please enter a message."
return
# Hard cap to protect ZeroGPU runtime.
max_new_tokens = min(
int(max_new_tokens),
MAX_NEW_TOKENS,
)
# --------------------------------------------------------
# Build conversation
# --------------------------------------------------------
conversation = build_conversation(
message=message,
chat_history=chat_history,
system_prompt=system_prompt,
)
# --------------------------------------------------------
# Tokenize with native model chat template
# --------------------------------------------------------
input_ids = tokenizer.apply_chat_template(
conversation,
tokenize=True,
add_generation_prompt=True,
return_tensors="pt",
)
# --------------------------------------------------------
# Context limit
# --------------------------------------------------------
if input_ids.shape[1] > MAX_INPUT_TOKEN_LENGTH:
input_ids = input_ids[
:,
-MAX_INPUT_TOKEN_LENGTH:
]
gr.Warning(
f"Conversation trimmed to "
f"{MAX_INPUT_TOKEN_LENGTH} tokens."
)
# --------------------------------------------------------
# CUDA
# --------------------------------------------------------
input_ids = input_ids.to("cuda")
# --------------------------------------------------------
# Stream generation
# --------------------------------------------------------
streamer = TextIteratorStreamer(
tokenizer,
timeout=30.0,
skip_prompt=True,
skip_special_tokens=True,
)
generation_kwargs = {
"input_ids": input_ids,
"streamer": streamer,
"max_new_tokens": max_new_tokens,
"do_sample": True,
"temperature": float(temperature),
"top_p": float(top_p),
"top_k": int(top_k),
"repetition_penalty": float(repetition_penalty),
"num_beams": 1,
"use_cache": True,
}
thread = Thread(
target=model.generate,
kwargs=generation_kwargs,
daemon=True,
)
thread.start()
generated_text = ""
try:
for text in streamer:
generated_text += text
yield generated_text
except Exception as exc:
print(f"Generation error: {exc}")
raise
finally:
thread.join(timeout=1.0)
# ============================================================
# UI
# ============================================================
chat_interface = gr.ChatInterface(
fn=generate,
additional_inputs=[
gr.Textbox(
label="System prompt",
lines=4,
value=DEFAULT_SYSTEM_PROMPT,
),
gr.Slider(
label="Max new tokens",
minimum=16,
maximum=MAX_NEW_TOKENS,
step=16,
value=DEFAULT_MAX_NEW_TOKENS,
),
gr.Slider(
label="Temperature",
minimum=0.1,
maximum=1.5,
step=0.1,
value=0.5,
),
gr.Slider(
label="Top-p",
minimum=0.1,
maximum=1.0,
step=0.05,
value=0.95,
),
gr.Slider(
label="Top-k",
minimum=0,
maximum=100,
step=1,
value=40,
),
gr.Slider(
label="Repetition penalty",
minimum=1.0,
maximum=1.5,
step=0.05,
value=1.0,
),
],
examples=[
["Hello there!"],
[
"Explain Python in three sentences."
],
[
"Explain the plot of Cinderella briefly."
],
[
"Write a short paragraph about open-source AI."
],
],
)
# ============================================================
# APP
# ============================================================
with gr.Blocks(
title="GLM-4.7-Flash Fable Chat",
) as demo:
gr.Markdown(
"""
# GLM-4.7-Flash Fable Chat
Fast ZeroGPU chat interface.
"""
)
chat_interface.render()
# ============================================================
# LAUNCH
# ============================================================
if __name__ == "__main__":
demo.queue(
max_size=10,
).launch(
show_error=True,
)