Spaces:
Sleeping
Sleeping
Commit ·
58293de
1
Parent(s): 71dd534
Harden inference token ids
Browse files- src/inference.py +16 -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 |
-
|
| 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)
|