| 142 | |
| 143 | class RotationTokenizer: |
| 144 | def __init__( |
| 145 | self, |
| 146 | tokenizer: PreTrainedTokenizerBase, |
| 147 | num_bins: Dict, |
| 148 | bin_policy: Optional[Dict] = None, |
| 149 | array_begin_idx=None, |
| 150 | ): |
| 151 | self.tokenizer = tokenizer |
| 152 | self.num_roll_bins = num_bins["roll_bins"] # M |
| 153 | self.num_pitch_bins = num_bins["pitch_bins"] # N |
| 154 | self.num_yaw_bins = num_bins["yaw_bins"] # P |
| 155 | self.array_begin_idx = array_begin_idx |
| 156 | |
| 157 | # for indexing |
| 158 | self.NP = self.num_pitch_bins * self.num_yaw_bins |
| 159 | |
| 160 | # add special action tokens to language tokenizer |
| 161 | self._vocab_size = self.num_roll_bins * self.num_pitch_bins * self.num_yaw_bins |
| 162 | token_list = [ACTION_TOKEN.format(i + self.array_begin_idx) for i in range(self._vocab_size)] |
| 163 | self.token_array = np.array(token_list) |
| 164 | |
| 165 | num_new_tokens = self.tokenizer.add_tokens(token_list, special_tokens=True) |
| 166 | print(f"Add {num_new_tokens} ROTATION TOKENS to tokenizer, tokenizer vocab size {self.tokenizer.vocab_size} / {len(tokenizer)}") |
| 167 | |
| 168 | self.token_start_idx = self.tokenizer.convert_tokens_to_ids(self.token_array[0]) |
| 169 | self.token_end_idx = self.tokenizer.convert_tokens_to_ids(self.token_array[-1]) |
| 170 | self.set_bins(bin_policy) |
| 171 | |
| 172 | def set_bins(self, bin_policy): |
| 173 | self.roll_bins = np.array(bin_policy["roll_bins"]) |