Spaces:
Running on Zero
Running on Zero
File size: 4,447 Bytes
5f8ba7b 625e2c5 5f8ba7b f072b43 5f8ba7b | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 | 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)
|