Lavender825 commited on
Commit
58293de
·
1 Parent(s): 71dd534

Harden inference token ids

Browse files
Files changed (2) hide show
  1. src/inference.py +16 -2
  2. src/models.py +6 -0
src/inference.py CHANGED
@@ -173,6 +173,21 @@ class AspectPredictor:
173
  self.model, self.tokenizer, self.meta_encoder, self.device = load_meta_acsa(
174
  checkpoint_dir=checkpoint_dir, device=device,
175
  )
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
176
 
177
  def _prepare_input_ids(self, input_ids: torch.Tensor) -> torch.Tensor:
178
  bert = getattr(self.model, "bert", None)
@@ -183,8 +198,7 @@ class AspectPredictor:
183
  vocab_size = int(word_embeddings.num_embeddings)
184
  if vocab_size <= 0:
185
  return input_ids
186
- unk_id = int(getattr(self.tokenizer, "unk_token_id", 0) or 0)
187
- return torch.where(input_ids >= vocab_size, torch.full_like(input_ids, unk_id), input_ids)
188
 
189
  def _max_inference_length(self) -> int:
190
  bert = getattr(self.model, "bert", None)
 
173
  self.model, self.tokenizer, self.meta_encoder, self.device = load_meta_acsa(
174
  checkpoint_dir=checkpoint_dir, device=device,
175
  )
176
+ self._align_tokenizer_and_embeddings()
177
+
178
+ def _align_tokenizer_and_embeddings(self) -> None:
179
+ bert = getattr(self.model, "bert", None)
180
+ embeddings = getattr(bert, "embeddings", None)
181
+ word_embeddings = getattr(embeddings, "word_embeddings", None)
182
+ if bert is None or word_embeddings is None:
183
+ return
184
+ try:
185
+ tokenizer_size = len(self.tokenizer)
186
+ except Exception:
187
+ tokenizer_size = 0
188
+ vocab_size = int(getattr(word_embeddings, "num_embeddings", 0) or 0)
189
+ if tokenizer_size > vocab_size > 0 and hasattr(bert, "resize_token_embeddings"):
190
+ bert.resize_token_embeddings(tokenizer_size)
191
 
192
  def _prepare_input_ids(self, input_ids: torch.Tensor) -> torch.Tensor:
193
  bert = getattr(self.model, "bert", None)
 
198
  vocab_size = int(word_embeddings.num_embeddings)
199
  if vocab_size <= 0:
200
  return input_ids
201
+ return input_ids.clamp(min=0, max=vocab_size - 1)
 
202
 
203
  def _max_inference_length(self) -> int:
204
  bert = getattr(self.model, "bert", None)
src/models.py CHANGED
@@ -61,6 +61,12 @@ def _load_bert_model(bert_name: str):
61
 
62
 
63
  def _safe_bert_forward(bert, input_ids, attention_mask, output_attentions=False):
 
 
 
 
 
 
64
  embeddings = getattr(bert, "embeddings", None)
65
  word_embeddings = getattr(embeddings, "word_embeddings", None)
66
  position_embeddings = getattr(embeddings, "position_embeddings", None)
 
61
 
62
 
63
  def _safe_bert_forward(bert, input_ids, attention_mask, output_attentions=False):
64
+ input_ids = input_ids.long()
65
+ if attention_mask is None:
66
+ attention_mask = torch.ones_like(input_ids)
67
+ elif attention_mask.size(1) != input_ids.size(1):
68
+ attention_mask = attention_mask[:, :input_ids.size(1)]
69
+
70
  embeddings = getattr(bert, "embeddings", None)
71
  word_embeddings = getattr(embeddings, "word_embeddings", None)
72
  position_embeddings = getattr(embeddings, "position_embeddings", None)