Lavender825 commited on
Commit
dc9f449
·
1 Parent(s): 1f017cb

Clamp inference sequence length to BERT limit

Browse files
Files changed (1) hide show
  1. 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=cfg.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=cfg.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",