| 195 | yield s[slice_start:] |
| 196 | |
| 197 | class ChatFormat: |
| 198 | def __init__(self, tokenizer: Tokenizer): |
| 199 | self.tokenizer = tokenizer |
| 200 | self.eot_id = tokenizer.special_tokens["<|eot_id|>"] |
| 201 | |
| 202 | def decode(self, tokens: List[int]) -> str: |
| 203 | # Decode the tokens to a string. |
| 204 | decoded_str = self.tokenizer.decode(tokens) |
| 205 | # Remove the special tokens from the decoded string. |
| 206 | decoded_str = decoded_str.replace("<|eot_id|>", "") |
| 207 | return decoded_str |
| 208 | |
| 209 | def encode_header(self, message: Message) -> List[int]: |
| 210 | tokens = [] |
| 211 | if message["role"] == "system": |
| 212 | tokens.extend(self.tokenizer.encode("System: ", bos=False, eos=False)) |
| 213 | elif message["role"] == "user": |
| 214 | tokens.extend(self.tokenizer.encode("User: ", bos=False, eos=False)) |
| 215 | elif message["role"] == "assistant": |
| 216 | tokens.extend(self.tokenizer.encode("Assistant: ", bos=False, eos=False)) |
| 217 | else: |
| 218 | raise NotImplementedError(f"Role {message['role']} not implemented.") |
| 219 | # tokens.append(self.tokenizer.special_tokens["<|start_header_id|>"]) |
| 220 | # tokens.extend(self.tokenizer.encode(message["role"], bos=False, eos=False)) |
| 221 | # tokens.append(self.tokenizer.special_tokens["<|end_header_id|>"]) |
| 222 | # tokens.extend(self.tokenizer.encode("\n\n", bos=False, eos=False)) |
| 223 | return tokens |
| 224 | |
| 225 | def encode_message(self, message: Message, return_target=False) -> List[int]: |
| 226 | tokens, targets = [], [] |
| 227 | headers = self.encode_header(message) |
| 228 | contents = self.tokenizer.encode(message["content"].strip(), bos=False, eos=False) |
| 229 | contents.append(self.tokenizer.special_tokens["<|eot_id|>"]) |
| 230 | tokens = headers + contents |
| 231 | |
| 232 | if message["role"] == "assistant": |
| 233 | targets = [-1] * len(headers) + contents |
| 234 | else: |
| 235 | targets = [-1] * len(tokens) |
| 236 | |
| 237 | if return_target: |
| 238 | return tokens, targets |
| 239 | |
| 240 | return tokens, None |
| 241 | |
| 242 | def encode_dialog_prompt(self, dialog: Dialog, completion=False, return_target=False) -> List[int]: |
| 243 | tokens = [self.tokenizer.special_tokens["<|begin_of_text|>"]] |
| 244 | targets = [-1] |
| 245 | for message in dialog: |
| 246 | _tokens, _targets = self.encode_message(message, return_target=return_target) |
| 247 | tokens.extend(_tokens) |
| 248 | if _targets is not None: |
| 249 | targets.extend(_targets) |
| 250 | # Add the start of an assistant message for the model to complete. |
| 251 | if completion: |
| 252 | tokens.extend(self.encode_header({"role": "assistant", "content": ""})) |
| 253 | |
| 254 | if return_target: |