Kimi-K3 / app.py
devoppro's picture
Update app.py
8a6e1e0 verified
Raw History Blame Contribute Delete
10.5 kB
import spaces
import os
import gc
import time
import traceback
# ---------------------------------------------------------
# Persistent storage
# ---------------------------------------------------------
DATA_DIR = "/data"
HF_HOME = os.path.join(DATA_DIR, "hf-cache")
TRANSFORMERS_CACHE = os.path.join(HF_HOME, "transformers")
HF_HUB_CACHE = os.path.join(HF_HOME, "hub")
AIRLLM_CACHE = os.path.join(DATA_DIR, "airllm-cache")
for directory in [
DATA_DIR,
HF_HOME,
TRANSFORMERS_CACHE,
HF_HUB_CACHE,
AIRLLM_CACHE,
]:
os.makedirs(directory, exist_ok=True)
# Set these BEFORE importing transformers / huggingface libraries
os.environ["HF_HOME"] = HF_HOME
os.environ["HF_HUB_CACHE"] = HF_HUB_CACHE
os.environ["TRANSFORMERS_CACHE"] = TRANSFORMERS_CACHE
# AirLLM cache location
os.environ["AIRLLM_CACHE_DIR"] = AIRLLM_CACHE
# Prevent excessive CPU thread usage
os.environ.setdefault("TOKENIZERS_PARALLELISM", "false")
# ---------------------------------------------------------
# Imports
# ---------------------------------------------------------
import gradio as gr
import torch
from airllm import AutoModel
# ---------------------------------------------------------
# Configuration
# ---------------------------------------------------------
MODEL_ID = "devoppro/Kimi-K3"
MAX_NEW_TOKENS = 128
MAX_LENGTH = 2048
# ---------------------------------------------------------
# Global model
# ---------------------------------------------------------
model = None
tokenizer = None
model_status = "Not initialized"
# ---------------------------------------------------------
# Model loading
# ---------------------------------------------------------
def initialize_model():
global model
global tokenizer
global model_status
if model is not None:
return model
print("=" * 70)
print("Kimi K3 + AirLLM")
print("=" * 70)
print(f"Model: {MODEL_ID}")
print(f"HF_HOME: {HF_HOME}")
print(f"HF_HUB_CACHE: {HF_HUB_CACHE}")
print(f"AirLLM cache: {AIRLLM_CACHE}")
print()
print("Checking persistent storage...")
try:
total, used, free = torch.cuda.mem_get_info()
print(f"CUDA memory available: {free / 1024**3:.2f} GB")
except Exception:
print("CUDA memory information unavailable outside GPU worker.")
print()
print("Preparing Kimi K3...")
print("If this is the first startup, Hugging Face will download the")
print("large model checkpoint to the persistent /data volume.")
print()
start = time.time()
try:
model = AutoModel.from_pretrained(
MODEL_ID,
compression="4B",
trust_remote_code=True,
)
tokenizer = model.tokenizer
elapsed = time.time() - start
model_status = (
f"Model ready — initialization took {elapsed:.1f} seconds"
)
print("=" * 70)
print(model_status)
print("=" * 70)
return model
except Exception as e:
model_status = f"Initialization failed: {e}"
print()
print("MODEL INITIALIZATION ERROR")
print(traceback.format_exc())
raise
# ---------------------------------------------------------
# Initialize at module level
# ---------------------------------------------------------
#
# ZeroGPU expects model registration/loading at module scope.
#
# The actual ZeroGPU backend handles the CUDA registration
# and disk offload mechanism.
# ---------------------------------------------------------
try:
initialize_model()
except Exception as e:
print()
print("WARNING:")
print("Kimi K3 was not initialized during startup.")
print(str(e))
# ---------------------------------------------------------
# Generation
# ---------------------------------------------------------
@spaces.GPU(
duration=180,
size="xlarge",
)
def generate(
message,
history,
max_new_tokens,
temperature,
):
global model
global tokenizer
if model is None:
initialize_model()
if not message or not message.strip():
return "", history
# -----------------------------------------------------
# Build conversation
# -----------------------------------------------------
conversation = []
if history:
for item in history:
if isinstance(item, dict):
role = item.get("role")
content = item.get("content", "")
if role in ["user", "assistant"]:
conversation.append(
{
"role": role,
"content": content,
}
)
elif isinstance(item, (list, tuple)):
if len(item) == 2:
user_message, assistant_message = item
if user_message:
conversation.append(
{
"role": "user",
"content": user_message,
}
)
if assistant_message:
conversation.append(
{
"role": "assistant",
"content": assistant_message,
}
)
conversation.append(
{
"role": "user",
"content": message,
}
)
# -----------------------------------------------------
# Tokenize
# -----------------------------------------------------
try:
prompt = tokenizer.apply_chat_template(
conversation,
tokenize=False,
add_generation_prompt=True,
)
except Exception:
# Fallback for processors/tokenizers without
# chat-template support.
prompt = ""
for item in conversation:
prompt += (
f"{item['role'].upper()}: "
f"{item['content']}\n"
)
prompt += "ASSISTANT:"
print()
print("Generating response...")
print(f"Prompt length: {len(prompt)} characters")
inputs = tokenizer(
[prompt],
return_tensors="pt",
padding=True,
truncation=True,
max_length=MAX_LENGTH,
)
input_ids = inputs["input_ids"].cuda()
attention_mask = None
if "attention_mask" in inputs:
attention_mask = inputs["attention_mask"].cuda()
# -----------------------------------------------------
# Generation
# -----------------------------------------------------
generation_kwargs = {
"max_new_tokens": int(max_new_tokens),
"do_sample": True,
"temperature": float(temperature),
}
if attention_mask is not None:
generation_output = model.generate(
input_ids,
attention_mask=attention_mask,
**generation_kwargs,
)
else:
generation_output = model.generate(
input_ids,
**generation_kwargs,
)
# -----------------------------------------------------
# Decode
# -----------------------------------------------------
generated_tokens = generation_output[0][
input_ids.shape[-1]:
]
answer = tokenizer.decode(
generated_tokens,
skip_special_tokens=True,
)
answer = answer.strip()
# -----------------------------------------------------
# Cleanup
# -----------------------------------------------------
del input_ids
if attention_mask is not None:
del attention_mask
del generation_output
gc.collect()
try:
torch.cuda.empty_cache()
except Exception:
pass
new_history = list(history or [])
new_history.append(
{
"role": "user",
"content": message,
}
)
new_history.append(
{
"role": "assistant",
"content": answer,
}
)
return "", new_history
# ---------------------------------------------------------
# UI
# ---------------------------------------------------------
with gr.Blocks() as demo:
gr.Markdown(
"""
# 🤖 Kimi K3 + AirLLM
### Kimi K3 running through AirLLM
This demo uses the persistent `/data` storage volume for the
large model checkpoint.
**First startup may take a long time.**
"""
)
status = gr.Markdown(
"### Model status\n"
"The model is being prepared..."
)
chatbot = gr.Chatbot(
label="Kimi K3",
type="messages",
height=600,
)
with gr.Row():
message = gr.Textbox(
label="Message",
placeholder="Ask Kimi K3 anything...",
scale=5,
)
send = gr.Button(
"Send",
variant="primary",
scale=1,
)
with gr.Row():
max_tokens = gr.Slider(
minimum=16,
maximum=512,
value=128,
step=16,
label="Max new tokens",
)
temperature = gr.Slider(
minimum=0.1,
maximum=1.5,
value=0.7,
step=0.1,
label="Temperature",
)
clear = gr.Button("Clear conversation")
# -----------------------------------------------------
# Events
# -----------------------------------------------------
send.click(
fn=generate,
inputs=[
message,
chatbot,
max_tokens,
temperature,
],
outputs=[
message,
chatbot,
],
)
message.submit(
fn=generate,
inputs=[
message,
chatbot,
max_tokens,
temperature,
],
outputs=[
message,
chatbot,
],
)
clear.click(
fn=lambda: [],
inputs=None,
outputs=chatbot,
)
# Update status after startup
status.value = (
f"### Model status\n"
f"`{model_status}`"
)
# ---------------------------------------------------------
# Launch
# ---------------------------------------------------------
demo.launch(
server_name="0.0.0.0",
server_port=7860,
ssr_mode=True,
)