augustoFranke commited on
Commit
5ab302c
·
verified ·
1 Parent(s): cb96101

Warm up every bucket at load so the first request per size is not ~10x slower

Browse files
Files changed (1) hide show
  1. src/gliner_decide_coreml/router.py +12 -0
src/gliner_decide_coreml/router.py CHANGED
@@ -9,6 +9,7 @@ from dataclasses import dataclass
9
  from pathlib import Path
10
 
11
  import coremltools as ct
 
12
 
13
  from .decoding import decode, head_probabilities, merge_chunks
14
  from .encoding import Encoded, encode, load_collator, task_labels, to_bucket
@@ -57,6 +58,17 @@ class DecideRouter:
57
  for length in BUCKETS
58
  }
59
  self.chunk_overlap_words = chunk_overlap_words
 
 
 
 
 
 
 
 
 
 
 
60
 
61
  def count_tokens(self, text: str, tasks: dict) -> int:
62
  return encode(self.collator, text, tasks).length
 
9
  from pathlib import Path
10
 
11
  import coremltools as ct
12
+ import numpy as np
13
 
14
  from .decoding import decode, head_probabilities, merge_chunks
15
  from .encoding import Encoded, encode, load_collator, task_labels, to_bucket
 
58
  for length in BUCKETS
59
  }
60
  self.chunk_overlap_words = chunk_overlap_words
61
+ self._warm_up()
62
+
63
+ def _warm_up(self):
64
+ """Run each function once: the first prediction per function is ~10x slower (device setup)."""
65
+ for length, model in self.models.items():
66
+ model.predict({
67
+ "input_ids": np.full((1, length), self.pad_id, dtype=np.int32),
68
+ "attention_mask": np.ones((1, length), dtype=np.int32),
69
+ "marker_indices": np.zeros((1, MAX_HEADS, MAX_OPTIONS), dtype=np.int32),
70
+ "marker_mask": np.ones((1, MAX_HEADS, MAX_OPTIONS), dtype=np.float32),
71
+ })
72
 
73
  def count_tokens(self, text: str, tasks: dict) -> int:
74
  return encode(self.collator, text, tasks).length