| 241 | self.freeze() |
| 242 | |
| 243 | def add_token(self): |
| 244 | self.modifier_token_id = [] |
| 245 | token_embeds1 = self.transformer.get_input_embeddings().weight.data |
| 246 | for each_modifier_token in self.modifier_token: |
| 247 | num_added_tokens = self.tokenizer.add_tokens(each_modifier_token) |
| 248 | modifier_token_id = self.tokenizer.convert_tokens_to_ids(each_modifier_token) |
| 249 | self.modifier_token_id.append(modifier_token_id) |
| 250 | |
| 251 | self.transformer.resize_token_embeddings(len(self.tokenizer)) |
| 252 | token_embeds = self.transformer.get_input_embeddings().weight.data |
| 253 | token_embeds[self.modifier_token_id[-1]] = torch.nn.Parameter(token_embeds[42170], requires_grad=True) |
| 254 | if len(self.modifier_token) == 2: |
| 255 | token_embeds[self.modifier_token_id[-2]] = torch.nn.Parameter(token_embeds[47629], requires_grad=True) |
| 256 | if len(self.modifier_token) == 3: |
| 257 | token_embeds[self.modifier_token_id[-3]] = torch.nn.Parameter(token_embeds[43514], requires_grad=True) |
| 258 | |
| 259 | def custom_forward(self, hidden_states, input_ids): |
| 260 | r""" |