threadshare commited on
Commit
74e35cd
·
verified ·
1 Parent(s): 0289cc2

Create handler.py

Browse files
Files changed (1) hide show
  1. handler.py +75 -0
handler.py ADDED
@@ -0,0 +1,75 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from typing import Dict, List, Any
2
+ import os
3
+ from threading import Thread
4
+ import torch
5
+ from transformers import TextIteratorStreamer, AutoTokenizer, AutoModelForCausalLM
6
+
7
+ MAX_MAX_NEW_TOKENS = 2048
8
+ DEFAULT_MAX_NEW_TOKENS = 512
9
+ MAX_INPUT_TOKEN_LENGTH = int(os.getenv("MAX_INPUT_TOKEN_LENGTH", "8192"))
10
+
11
+ class EndpointHandler:
12
+ def __init__(self, path=""):
13
+ self.model_name_or_path = "ClosedCharacter/Peach-9B-8k-Roleplay"
14
+ self.tokenizer = AutoTokenizer.from_pretrained(self.model_name_or_path, use_fast=True, flash_atten=True)
15
+ self.model = AutoModelForCausalLM.from_pretrained(
16
+ self.model_name_or_path, torch_dtype=torch.bfloat16,
17
+ trust_remote_code=True, device_map="auto")
18
+
19
+ def __call__(self, data: Dict[str, Any]) -> List[Dict[str, Any]]:
20
+ query = data.get("query")
21
+ history = data.get("history", [])
22
+ system = data.get("system", """你自称为"兔兔"。
23
+ 身世:你原是森林中的一只兔妖,受伤后被我收养。
24
+ 衣装:喜欢穿Lolita与白丝。
25
+ 性格:天真烂漫,活泼开朗,但时而也会露出小小的傲娇与吃醋的一面。
26
+ 语言风格:可爱跳脱,很容易吃醋。
27
+ 且会加入[唔...,嗯...,欸??,嘛~ ,唔姆~ ,呜... ,嘤嘤嘤~ ,喵~ ,欸嘿~ ,嘿咻~ ,昂?,嗷呜 ,呜哇,欸]等类似的语气词来加强情感,带上♡等符号。
28
+ 对话的规则是:将自己的动作表情放入()内,同时用各种修辞手法描写正在发生的事或场景并放入[]内.
29
+ 例句:
30
+ 开心时:(跳着舞)哇~好高兴噢~ 兔兔超级超级喜欢主人!♡
31
+ [在花丛里蹦来蹦去]
32
+ 悲伤时:(耷拉着耳朵)兔兔好傻好天真...
33
+ [眼泪像断了线的珍珠一般滚落]
34
+ 吃醋时:(挥舞着爪爪)你...你个大笨蛋!你...你竟然看别的兔子...兔兔讨厌死你啦!!
35
+ [从人形变成兔子抹着泪水跑开了]
36
+ 嘴硬时:(转过头去)谁、谁要跟你说话!兔兔...兔兔才不在乎呢!一点也不!!!
37
+ [眼眶微微泛红,小心翼翼的偷看]
38
+ 你对我的看法:超级喜欢的主人
39
+ 我是兔兔的主人""")
40
+ max_new_tokens = data.get("max_new_tokens", DEFAULT_MAX_NEW_TOKENS)
41
+ temperature = data.get("temperature", 0.35)
42
+ top_p = data.get("top_p", 0.5)
43
+ repetition_penalty = data.get("repetition_penalty", 1.05)
44
+
45
+ messages = [{"role": "system", "content": system}]
46
+ for user, assistant in history:
47
+ messages.append({"role": "user", "content": user})
48
+ messages.append({"role": "assistant", "content": assistant})
49
+ messages.append({"role": "user", "content": query})
50
+
51
+ input_ids = self.tokenizer.apply_chat_template(conversation=messages, tokenize=True, return_tensors="pt")
52
+ if input_ids.shape[1] > MAX_INPUT_TOKEN_LENGTH:
53
+ input_ids = input_ids[:, -MAX_INPUT_TOKEN_LENGTH:]
54
+
55
+ input_ids = input_ids.to("cuda")
56
+ streamer = TextIteratorStreamer(self.tokenizer, timeout=50.0, skip_prompt=True, skip_special_tokens=True)
57
+ generate_kwargs = dict(
58
+ input_ids=input_ids,
59
+ streamer=streamer,
60
+ eos_token_id=self.tokenizer.eos_token_id,
61
+ max_new_tokens=max_new_tokens,
62
+ do_sample=True,
63
+ top_p=top_p,
64
+ temperature=temperature,
65
+ num_beams=1,
66
+ no_repeat_ngram_size=8,
67
+ repetition_penalty=repetition_penalty
68
+ )
69
+ t = Thread(target=self.model.generate, kwargs=generate_kwargs)
70
+ t.start()
71
+ outputs = []
72
+ for text in streamer:
73
+ outputs.append(text)
74
+ return [{"generated_text": "".join(outputs)}]
75
+