Spaces:
Sleeping
Sleeping
Commit ·
dc9f449
1
Parent(s): 1f017cb
Clamp inference sequence length to BERT limit
Browse files- src/inference.py +11 -2
src/inference.py
CHANGED
|
@@ -186,11 +186,20 @@ class AspectPredictor:
|
|
| 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 predict(self, review_text: str, product_meta: Dict,
|
| 190 |
return_attention: bool = True) -> Dict:
|
| 191 |
enc = self.tokenizer(
|
| 192 |
review_text,
|
| 193 |
-
max_length=
|
| 194 |
truncation=True,
|
| 195 |
padding="max_length",
|
| 196 |
return_tensors="pt",
|
|
@@ -234,7 +243,7 @@ class AspectPredictor:
|
|
| 234 |
chunk = reviews[i:i + batch_size]
|
| 235 |
enc = self.tokenizer(
|
| 236 |
[r["review_text"] for r in chunk],
|
| 237 |
-
max_length=
|
| 238 |
truncation=True,
|
| 239 |
padding="max_length",
|
| 240 |
return_tensors="pt",
|
|
|
|
| 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)
|
| 191 |
+
config = getattr(bert, "config", None)
|
| 192 |
+
max_positions = int(getattr(config, "max_position_embeddings", cfg.MAX_LENGTH) or cfg.MAX_LENGTH)
|
| 193 |
+
tokenizer_max = int(getattr(self.tokenizer, "model_max_length", cfg.MAX_LENGTH) or cfg.MAX_LENGTH)
|
| 194 |
+
if tokenizer_max > 100000:
|
| 195 |
+
tokenizer_max = cfg.MAX_LENGTH
|
| 196 |
+
return max(8, min(int(cfg.MAX_LENGTH), max_positions, tokenizer_max))
|
| 197 |
+
|
| 198 |
def predict(self, review_text: str, product_meta: Dict,
|
| 199 |
return_attention: bool = True) -> Dict:
|
| 200 |
enc = self.tokenizer(
|
| 201 |
review_text,
|
| 202 |
+
max_length=self._max_inference_length(),
|
| 203 |
truncation=True,
|
| 204 |
padding="max_length",
|
| 205 |
return_tensors="pt",
|
|
|
|
| 243 |
chunk = reviews[i:i + batch_size]
|
| 244 |
enc = self.tokenizer(
|
| 245 |
[r["review_text"] for r in chunk],
|
| 246 |
+
max_length=self._max_inference_length(),
|
| 247 |
truncation=True,
|
| 248 |
padding="max_length",
|
| 249 |
return_tensors="pt",
|