import spaces import json import threading from pathlib import Path import torch from fastapi.responses import HTMLResponse from gradio import Error, Server from transformers import ( AutoModelForCausalLM, AutoTokenizer, LogitsProcessor, LogitsProcessorList, TextIteratorStreamer, ) MODEL_ID = "movingcastles/zero" EOS_IDS = [151645, 151643] MAX_CONTEXT = 16384 MAX_MESSAGE_CHARS = 4000 tokenizer = AutoTokenizer.from_pretrained(MODEL_ID) model = AutoModelForCausalLM.from_pretrained(MODEL_ID, dtype=torch.bfloat16).to("cuda") model.eval() PAGE = (Path(__file__).parent / "index.html").read_text() class PenaltyThenTemperature(LogitsProcessor): def __init__(self, prompt_len, presence_penalty, temperature): self.prompt_len = prompt_len self.presence_penalty = presence_penalty self.temperature = temperature def __call__(self, input_ids, scores): generated = input_ids[:, self.prompt_len :] if self.presence_penalty and generated.shape[1]: seen = torch.zeros_like(scores, dtype=torch.bool) seen.scatter_(1, generated, True) scores = scores - self.presence_penalty * seen.to(scores.dtype) return scores / self.temperature def parse_messages(raw): try: messages = json.loads(raw) except (TypeError, ValueError): raise Error("Malformed conversation.") if not isinstance(messages, list) or not messages: raise Error("Say something first.") clean = [] for i, m in enumerate(messages): role = m.get("role") if isinstance(m, dict) else None content = m.get("content") if isinstance(m, dict) else None expected = "user" if i % 2 == 0 else "assistant" if role != expected or not isinstance(content, str): raise Error("Conversation must alternate between you and Zero.") clean.append({"role": role, "content": content[:MAX_MESSAGE_CHARS]}) if clean[-1]["role"] != "user": raise Error("It is your turn to speak.") return clean def encode(messages, max_tokens): budget = MAX_CONTEXT - max_tokens while True: ids = tokenizer.apply_chat_template( messages, add_generation_prompt=True, return_tensors="pt", return_dict=True )["input_ids"] if ids.shape[1] <= budget or len(messages) <= 1: return ids[:, -budget:] messages = messages[2:] def gpu_duration(input_ids, temperature, presence_penalty, max_tokens): return 8 + int(max_tokens / 40) @spaces.GPU(duration=gpu_duration) def generate(input_ids, temperature, presence_penalty, max_tokens): input_ids = input_ids.to("cuda") streamer = TextIteratorStreamer(tokenizer, skip_prompt=True, skip_special_tokens=True, timeout=60) kwargs = dict( input_ids=input_ids, attention_mask=torch.ones_like(input_ids), max_new_tokens=max_tokens, eos_token_id=EOS_IDS, pad_token_id=EOS_IDS[1], streamer=streamer, repetition_penalty=1.0, ) if temperature > 0: kwargs.update( do_sample=True, temperature=1.0, top_p=1.0, top_k=0, logits_processor=LogitsProcessorList( [PenaltyThenTemperature(input_ids.shape[1], presence_penalty, temperature)] ), ) else: kwargs.update(do_sample=False, temperature=None, top_p=None, top_k=None) def run(): with torch.inference_mode(): model.generate(**kwargs) thread = threading.Thread(target=run, daemon=True) thread.start() text = "" for chunk in streamer: text += chunk yield text thread.join() app = Server(title="Zero") @app.api(name="respond", time_limit=150) def respond(messages: str, temperature: float = 0.7, presence_penalty: float = 1.5, max_tokens: int = 1024) -> str: conversation = parse_messages(messages) temperature = min(max(float(temperature), 0.0), 1.5) presence_penalty = min(max(float(presence_penalty), 0.0), 2.0) max_tokens = min(max(int(max_tokens), 16), 1024) input_ids = encode(conversation, max_tokens) for text in generate(input_ids, temperature, presence_penalty, max_tokens): yield text.strip() @app.get("/", response_class=HTMLResponse) async def homepage(): return PAGE demo = app if __name__ == "__main__": demo.launch(ssr_mode=False, show_error=True)