| 227 | |
| 228 | |
| 229 | def _decode_chatml( |
| 230 | tokens: List[int], |
| 231 | *, |
| 232 | stop_words: List[str], |
| 233 | eod_token_ids: List[int], |
| 234 | tokenizer: PreTrainedTokenizer, |
| 235 | raw_text_len: int, |
| 236 | context_length: int, |
| 237 | verbose: bool = False, |
| 238 | return_end_reason: bool = False, |
| 239 | errors: str='replace' |
| 240 | ): |
| 241 | end_reason = f"Gen length {len(tokens)}" |
| 242 | eod_token_idx = context_length |
| 243 | for eod_token_idx in range(context_length, len(tokens)): |
| 244 | if tokens[eod_token_idx] in eod_token_ids: |
| 245 | end_reason = f"Gen {tokenizer.decode([tokens[eod_token_idx]])!r}" |
| 246 | break |
| 247 | |
| 248 | trim_decode_tokens = tokenizer.decode(tokens[:eod_token_idx], errors=errors)[raw_text_len:] |
| 249 | if verbose: |
| 250 | print("\nRaw Generate w/o EOD:", tokenizer.decode(tokens, errors=errors)[raw_text_len:]) |
| 251 | print("\nRaw Generate:", trim_decode_tokens) |
| 252 | print("\nEnd Reason:", end_reason) |
| 253 | for stop_word in stop_words: |
| 254 | trim_decode_tokens = trim_decode_tokens.replace(stop_word, "").strip() |
| 255 | trim_decode_tokens = trim_decode_tokens.strip() |
| 256 | if verbose: |
| 257 | print("\nGenerate:", trim_decode_tokens) |
| 258 | |
| 259 | if return_end_reason: |
| 260 | return trim_decode_tokens, end_reason |
| 261 | else: |
| 262 | return trim_decode_tokens |
| 263 | |
| 264 | |
| 265 | def decode_tokens( |