| from __future__ import annotations |
|
|
| import os |
| import sys |
| import threading |
| from typing import Generator |
|
|
| import gradio as gr |
| import torch |
| from transformers import AutoModelForCausalLM, AutoTokenizer, TextIteratorStreamer |
|
|
| MODEL_ID = os.getenv("MODEL_ID", "IFM/K2-Horizon-0.9B") |
| HF_TOKEN = os.getenv("HF_TOKEN") |
|
|
| print(f"[startup] Loading tokenizer for {MODEL_ID}...", flush=True) |
| tokenizer = AutoTokenizer.from_pretrained( |
| MODEL_ID, |
| trust_remote_code=True, |
| token=HF_TOKEN, |
| ) |
|
|
| print(f"[startup] Loading model {MODEL_ID}...", flush=True) |
| model = AutoModelForCausalLM.from_pretrained( |
| MODEL_ID, |
| dtype=torch.float32, |
| device_map="cpu", |
| low_cpu_mem_usage=True, |
| trust_remote_code=True, |
| token=HF_TOKEN, |
| ) |
| model.eval() |
| print("[startup] Model loaded successfully!", flush=True) |
|
|
|
|
| def chat_response( |
| message: str, |
| history: list[list[str]], |
| system_prompt: str, |
| temperature: float, |
| top_p: float, |
| max_tokens: int, |
| reasoning_effort: str, |
| ) -> Generator[str, None, None]: |
| messages = [] |
| if system_prompt.strip(): |
| messages.append({"role": "system", "content": system_prompt}) |
|
|
| for user_msg, assistant_msg in history: |
| if user_msg: |
| messages.append({"role": "user", "content": user_msg}) |
| if assistant_msg: |
| messages.append({"role": "assistant", "content": assistant_msg}) |
|
|
| messages.append({"role": "user", "content": message}) |
|
|
| chat_kwargs = {"reasoning_effort": reasoning_effort} if reasoning_effort else {} |
|
|
| try: |
| prompt = tokenizer.apply_chat_template( |
| messages, |
| tokenize=False, |
| add_generation_prompt=True, |
| chat_template_kwargs=chat_kwargs if chat_kwargs else None, |
| ) |
| except Exception: |
| prompt = tokenizer.apply_chat_template( |
| messages, |
| tokenize=False, |
| add_generation_prompt=True, |
| ) |
|
|
| inputs = tokenizer([prompt], return_tensors="pt") |
| inputs.pop("token_type_ids", None) |
|
|
| streamer = TextIteratorStreamer( |
| tokenizer, timeout=30.0, skip_prompt=True, skip_special_tokens=True |
| ) |
|
|
| generate_kwargs = dict( |
| **inputs, |
| streamer=streamer, |
| max_new_tokens=int(max_tokens), |
| do_sample=temperature > 0, |
| temperature=float(temperature) if temperature > 0 else 1.0, |
| top_p=float(top_p), |
| ) |
|
|
| thread = threading.Thread(target=model.generate, kwargs=generate_kwargs) |
| thread.start() |
|
|
| partial_text = "" |
| for new_text in streamer: |
| partial_text += new_text |
| yield partial_text |
|
|
|
|
| demo = gr.ChatInterface( |
| fn=chat_response, |
| title="K2-Horizon-0.9B Chat Demo", |
| description="Interactive demo for [IFM/K2-Horizon-0.9B](https://huggingface.co/IFM/K2-Horizon-0.9B) using PyTorch and Transformers on CPU.", |
| additional_inputs=[ |
| gr.Textbox( |
| value="You are a helpful and harmless assistant.", |
| label="System Prompt", |
| ), |
| gr.Slider( |
| minimum=0.0, |
| maximum=2.0, |
| value=0.6, |
| step=0.1, |
| label="Temperature", |
| ), |
| gr.Slider( |
| minimum=0.1, |
| maximum=1.0, |
| value=0.95, |
| step=0.05, |
| label="Top-P", |
| ), |
| gr.Slider( |
| minimum=128, |
| maximum=8192, |
| value=2048, |
| step=128, |
| label="Max Tokens", |
| ), |
| gr.Dropdown( |
| choices=["high", "medium", "low"], |
| value="high", |
| label="Reasoning Effort", |
| ), |
| ], |
| ) |
|
|
| if __name__ == "__main__": |
| demo.queue().launch(server_name="0.0.0.0", server_port=7860) |
|
|