| from transformers import PreTrainedTokenizer |
| from typing import List, Optional, Dict |
|
|
| |
| |
| |
| NUM_NODES = 100 |
|
|
|
|
| class STokenizer(PreTrainedTokenizer): |
| def __init__(self, num_nodes=NUM_NODES): |
| |
| |
| |
| |
| self.vocab = {str(i): i for i in range(0, num_nodes)} |
| n = num_nodes |
| self.vocab['<|start-latent|>'] = n |
| self.vocab['<|end-latent|>'] = n + 1 |
| self.vocab['<|latent|>'] = n + 2 |
| self.vocab['|'] = n + 3 |
| self.vocab['[Q]'] = n + 4 |
| self.vocab['[R]'] = n + 5 |
| self.vocab['[A]'] = n + 6 |
| |
| self.vocab['<eos>'] = n + 7 |
| self.vocab['<|no-answer|>'] = n + 8 |
| |
| |
| self.vocab['<unk>'] = n + 9 |
| |
| |
| self.ids_to_tokens = {v: k for k, v in self.vocab.items()} |
| |
| |
| super().__init__( |
| pad_token="<eos>", eos_token="<eos>", bos_token="<eos>", unk_token="<unk>" |
| ) |
|
|
| def get_vocab(self) -> Dict[str, int]: |
| """Returns the vocabulary as a dict""" |
| return self.vocab.copy() |
|
|
| @property |
| def vocab_size(self) -> int: |
| return len(self.vocab) |
| |
| def _tokenize(self, text: str) -> List[str]: |
| |
| tokens = [] |
| for token in text.replace("\n", " ").strip().split(): |
| if token in self.vocab: |
| tokens.append(token) |
| else: |
| raise ValueError(f"Token {token} not in vocabulary") |
| return tokens |
| |
| def _convert_token_to_id(self, token: str) -> int: |
| |
| return self.vocab[token] |
| |
| def _convert_id_to_token(self, index: int) -> str: |
| |
| |
| return self.ids_to_tokens.get(int(index), "<unk>") |
| |
| def convert_tokens_to_string(self, tokens: List[str]) -> str: |
| |
| return ' '.join(tokens) |
| |
| def build_inputs_with_special_tokens(self, token_ids_0: List[int], |
| token_ids_1: Optional[List[int]] = None) -> List[int]: |
| |
| if token_ids_1 is None: |
| return token_ids_0 + [self.vocab['<eos>']] |
| return token_ids_0 + [self.vocab['<eos>']] + token_ids_1 + [self.vocab['<eos>']] |
|
|
| def get_special_tokens_mask(self, token_ids_0: List[int], |
| token_ids_1: Optional[List[int]] = None, |
| already_has_special_tokens: bool = False) -> List[int]: |
| |
| if already_has_special_tokens: |
| return [1 if token_id in [self.vocab['<eos>'], self.vocab['<pad>'], self.vocab['<unk>']] |
| else 0 for token_id in token_ids_0] |
| if token_ids_1 is None: |
| return [0] * len(token_ids_0) + [1] |
| return [0] * len(token_ids_0) + [1] + [0] * len(token_ids_1) + [1] |
|
|