Views
No views yet
pip install peft transformers bitsandbytes ipykernel rapidfuzz1import json
2from dataclasses import dataclass
3from enum import Enum
4from typing import List, Dict, Tuple, Literal
5
6class Roles(Enum):
7 system = "system"
8 user = "user"
9 assistant = "assistant"
10 tool = "tool"
11
12class MessagesFormatterType(Enum):
13 """
14 Enum representing different types of predefined messages formatters.
15 """
16
17 MISTRAL = 1
18
19@dataclass
20class PromptMarkers:
21 start: str
22 end: str
23
24class MessagesFormatter:
25 def __init__(
26 self,
27 pre_prompt: str,
28 prompt_markers: Dict[Roles, PromptMarkers],
29 include_sys_prompt_in_first_user_message: bool,
30 default_stop_sequences: List[str],
31 use_user_role_for_function_call_result: bool = True,
32 strip_prompt: bool = True,
33 bos_token: str = "<s>",
34 eos_token: str = "</s>"
35 ):
36 self.pre_prompt = pre_prompt
37 self.prompt_markers = prompt_markers
38 self.include_sys_prompt_in_first_user_message = include_sys_prompt_in_first_user_message
39 self.default_stop_sequences = default_stop_sequences
40 self.use_user_role_for_function_call_result = use_user_role_for_function_call_result
41 self.strip_prompt = strip_prompt
42 self.bos_token = bos_token
43 self.eos_token = eos_token
44 self.added_system_prompt = False
45
46 def get_bos_token(self) -> str:
47 return self.bos_token
48
49 def format_conversation(
50 self,
51 messages: List[Dict[str, str]],
52 response_role: Literal[Roles.user, Roles.assistant] | None = None,
53 ) -> Tuple[str, Roles]:
54 formatted_messages = self.pre_prompt
55 last_role = Roles.assistant
56 self.added_system_prompt = False
57 for message in messages:
58 role = Roles(message["role"])
59 content = self._format_message_content(message["content"], role)
60
61 if role == Roles.system:
62 formatted_messages += self._format_system_message(content)
63 last_role = Roles.system
64 elif role == Roles.user:
65 formatted_messages += self._format_user_message(content)
66 last_role = Roles.user
67 elif role == Roles.assistant:
68 formatted_messages += self._format_assistant_message(content)
69 last_role = Roles.assistant
70 elif role == Roles.tool:
71 formatted_messages += self._format_tool_message(content)
72 last_role = Roles.tool
73
74 return self._format_response(formatted_messages, last_role, response_role)
75
76 def _format_message_content(self, content: str, role: Roles) -> str:
77 if self.strip_prompt:
78 return content.strip()
79 return content
80
81 def _format_system_message(self, content: str) -> str:
82 formatted_message = self.prompt_markers[Roles.system].start + content + self.prompt_markers[Roles.system].end
83 self.added_system_prompt = True
84 if self.include_sys_prompt_in_first_user_message:
85 formatted_message = self.prompt_markers[Roles.user].start + formatted_message
86 return formatted_message
87
88 def _format_user_message(self, content: str) -> str:
89 if self.include_sys_prompt_in_first_user_message and self.added_system_prompt:
90 self.added_system_prompt = False
91 return content + self.prompt_markers[Roles.user].end
92 return self.prompt_markers[Roles.user].start + content + self.prompt_markers[Roles.user].end
93
94 def _format_assistant_message(self, content: str) -> str:
95 return self.prompt_markers[Roles.assistant].start + content + self.prompt_markers[Roles.assistant].end
96
97 def _format_tool_message(self, content: str) -> str:
98 if isinstance(content, list):
99 content = "\n".join(json.dumps(m, indent=2) for m in content)
100 if self.use_user_role_for_function_call_result:
101 return self._format_user_message(content)
102 else:
103 return self.prompt_markers[Roles.tool].start + content + self.prompt_markers[Roles.tool].end
104
105 def _format_response(
106 self,
107 formatted_messages: str,
108 last_role: Roles,
109 response_role: Literal[Roles.user, Roles.assistant] | None = None,
110 ) -> Tuple[str, Roles]:
111 if response_role is None:
112 response_role = Roles.assistant if last_role != Roles.assistant else Roles.user
113
114 prompt_start = self.prompt_markers[response_role].start.strip() if self.strip_prompt else self.prompt_markers[
115 response_role].start
116 return formatted_messages + prompt_start, response_role
117
118mixtral_prompt_markers = {
119 Roles.system: PromptMarkers("", """\n\n"""),
120 Roles.user: PromptMarkers("""[INST] """, """ [/INST]"""),
121 Roles.assistant: PromptMarkers("""""", """</s>"""),
122 Roles.tool: PromptMarkers("", ""),
123}
124
125mixtral_formatter = MessagesFormatter(
126 "",
127 mixtral_prompt_markers,
128 True,
129 ["</s>"],
130)
131
132from transformers import TextStreamer, AutoTokenizer, AutoModelForCausalLM
133from peft import PeftModel
134tokenizer = AutoTokenizer.from_pretrained("svjack/DPO_Genshin_Impact_Mistral_Plot_Engine_Step_Json_Short_merged",)
135mis_model = AutoModelForCausalLM.from_pretrained("svjack/DPO_Genshin_Impact_Mistral_Plot_Engine_Step_Json_Short_merged", load_in_4bit = True)
136mis_model = mis_model.eval()
137
138streamer = TextStreamer(tokenizer)
139
140def mistral_hf_predict(messages, mis_model = mis_model,
141 tokenizer = tokenizer, streamer = streamer,
142 do_sample = True,
143 top_p = 0.95,
144 top_k = 40,
145 max_new_tokens = 512,
146 max_input_length = 3500,
147 temperature = 0.9,
148 repetition_penalty = 1.0,
149 device = "cuda"):
150
151 #encodeds = tokenizer.apply_chat_template(messages, return_tensors="pt")
152 #model_inputs = encodeds.to(device)
153 prompt, _ = mixtral_formatter.format_conversation(messages)
154 model_inputs = tokenizer.encode(prompt, return_tensors="pt").to(device)
155
156 generated_ids = mis_model.generate(model_inputs, max_new_tokens=max_new_tokens,
157 do_sample=do_sample,
158 streamer = streamer,
159 top_p = top_p,
160 top_k = top_k,
161 temperature = temperature,
162 repetition_penalty = repetition_penalty,
163 )
164 out = tokenizer.batch_decode(generated_ids)[0].split("[/INST]")[-1].replace("</s>", "").strip()
165 return out
166
167from rapidfuzz import fuzz
168from IPython.display import clear_output
169def run_step_infer_times(x, times = 5, temperature = 0.01,
170 repetition_penalty = 1.0,
171 sim_val = 70
172 ):
173 req = []
174 for _ in range(times):
175 clear_output(wait = True)
176 out = mistral_hf_predict([
177 {
178 "role": "system",
179 "content": ""
180 },
181 {
182 "role": "user",
183 "content": x
184 },
185 ],
186 repetition_penalty = repetition_penalty,
187 temperature = temperature,
188 max_new_tokens = 2070,
189 max_input_length = 6000,
190 )
191 if req:
192 val = max(map(lambda x: fuzz.ratio(x, out), req))
193 #print(val)
194 #print(req)
195 if val < sim_val:
196 req.append(out.strip())
197 x = x.strip() + "\n" + out.strip()
198 else:
199 req.append(out.strip())
200 x = x.strip() + "\n" + out.strip()
201 return req
202
203out_l = run_step_infer_times(
204'''
205故事标题:归乡
206故事背景:在须弥城门口,派蒙与纳西妲偶遇并帮助一只昏迷的元素生命找寻家园。过程中揭示了这只生物并非普通的蕈兽,而是元素生物,并且它们曾受到过‘末日’的影响,家园被侵蚀。纳西妲回忆起晶体里的力量可能与一个预言有关,为了拯救它们的家园,她必须解决‘禁忌知识’问题,但这个过程对她自身也会产生干扰。
207参与角色:派蒙、纳西妲、浮游水蕈兽、旅行者
208''',
209 temperature=0.1,
210 repetition_penalty = 1.0,
211 times = 10
212)
213clear_output(wait = True)
214
215print("\n".join(out_l)){'参与者1': '派蒙', '参与者2': '纳西妲', '当前故事背景': '在须弥城门口,派蒙发现了一只昏迷的浮游水蕈兽,并询问它的家园。纳西妲确认这是元素生物,并解释了它们的特殊性和‘末日’的影响。纳西妲提出要帮助它们找到家园,但这需要解决‘禁忌知识’问题,这对她个人也有所挑战。'}
{'参与者1': '派蒙', '参与者2': '浮游水蕈兽', '当前故事背景': '派蒙询问浮游水蕈兽是否记得家园,它似乎对此感到疑惑,并且纳西妲解释了它们的特殊性和‘末日’后的变化。'}
{'参与者1': '纳西妲', '参与者2': '旅行者', '当前故事背景': '纳西妲提到‘禁忌知识’与她的力量有关,她需要解决这个问题以帮助元素生物。旅行者对此表示愿意协助,并在后续的对话中提到了‘晶体’的力量可能与纳西妲的预言有关。'}
{'参与者1': '纳西妲', '参与者2': '派蒙', '当前故事背景': '纳西妲提到‘晶体’中的力量可能与她的预言有联系,这让派蒙感到惊讶,并询问具体内容。'}
{'参与者1': '纳西妲', '参与者2': '旅行者', '当前故事背景': '纳西妲提出要去寻找‘禁忌知识’,旅行者表示愿意帮忙,两人准备前往目标地点。'}
{'参与者1': '派蒙', '参与者2': '纳西妲', '当前故事背景': '在寻找过程中,派蒙对纳西妲的力量和预言感到疑惑,纳西妲解释了这是为了帮助元素生物。'}
{'参与者1': '旅行者', '参与者2': '纳西妲', '当前故事背景': '旅行者表示愿意帮助解决‘禁忌知识’问题,两人的合作关系在故事中得到了展现。'}