Spaces:
Running on Zero
Running on Zero
Download app.py from multimodalart/zero: direct link, hf CLI and curl.
- Browser
- Download file 4.45 kB
-
https://huggingface.co/spaces/multimodalart/zero/resolve/main/app.py
- Command line
-
hf download hf://spaces/multimodalart/zero/app.py
-
curl -L -o app.py https://huggingface.co/spaces/multimodalart/zero/resolve/main/app.py
4.45 kB
| 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) | |
| 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") | |
| 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() | |
| async def homepage(): | |
| return PAGE | |
| demo = app | |
| if __name__ == "__main__": | |
| demo.launch(ssr_mode=False, show_error=True) | |