Text Ranking
sentence-transformers
Safetensors
Transformers
multilingual
t5gemma2
text2text-generation
reranker
encoder-decoder
FBNL
Retrieval
RAG
lukann98 commited on
Commit
aef8893
·
verified ·
1 Parent(s): e8eaadc

fix(reranker): avoid re-computing the first batch in predict()'s batch-size probe to reduce additional computational effort

Browse files

Existing Issue:
KaLMReranker.predict() runs a full forward pass on the first batch twice: once as a throwaway OOM-size probe, and once again for the real computation. The probe's result is computed but never used. This roughly doubles reranking latency for any call where the total number of documents fits within a single batch (the common case, since batch_size defaults to 32).

Proposed Fix:
Keep the probe's result and reuse it as the first batch's scores, starting the real loop after the first batch instead of from the beginning.

The former PR introduced a bug when Cuda OOM errors lead to tested_batch_size=1. In the former PR the first batch was then silently dropped. This PR adds the correct logic to prevent this behavior, by checking if scores were computed for the probe batch , see comment in lines 299-303

Files changed (1) hide show
  1. kalm_reranker.py +14 -3
kalm_reranker.py CHANGED
@@ -283,9 +283,10 @@ class KaLMReranker:
283
  sorted_pairs = [validated_pairs[index] for index in length_sorted_indices]
284
 
285
  tested_batch_size = effective_batch_size
 
286
  while tested_batch_size > 1:
287
  try:
288
- self._predict_batch(
289
  sorted_pairs[: min(len(sorted_pairs), tested_batch_size)],
290
  effective_instruction,
291
  )
@@ -295,9 +296,19 @@ class KaLMReranker:
295
  torch.cuda.empty_cache()
296
  tested_batch_size = max(1, tested_batch_size * 3 // 4)
297
 
298
- sorted_scores: List[float] = []
 
 
 
 
 
 
 
 
 
 
299
  try:
300
- for start in range(0, len(sorted_pairs), tested_batch_size):
301
  sorted_scores.extend(
302
  self._predict_batch(
303
  sorted_pairs[start : start + tested_batch_size],
 
283
  sorted_pairs = [validated_pairs[index] for index in length_sorted_indices]
284
 
285
  tested_batch_size = effective_batch_size
286
+ first_batch_scores: Optional[List[float]] = None
287
  while tested_batch_size > 1:
288
  try:
289
+ first_batch_scores = self._predict_batch(
290
  sorted_pairs[: min(len(sorted_pairs), tested_batch_size)],
291
  effective_instruction,
292
  )
 
296
  torch.cuda.empty_cache()
297
  tested_batch_size = max(1, tested_batch_size * 3 // 4)
298
 
299
+ # The while loop's condition (`> 1`) means batch size 1 is never
300
+ # actually probed. If every size down to 2 OOMs, it exits without a
301
+ # successful probe. Only skip ahead to `tested_batch_size` when the
302
+ # probe actually ran; otherwise fall back to starting at 0 like the
303
+ # loop below always did originally, or the first item(s) get dropped.
304
+ if first_batch_scores is None:
305
+ sorted_scores: List[float] = []
306
+ loop_start = 0
307
+ else:
308
+ sorted_scores = list(first_batch_scores)
309
+ loop_start = tested_batch_size
310
  try:
311
+ for start in range(loop_start, len(sorted_pairs), tested_batch_size):
312
  sorted_scores.extend(
313
  self._predict_batch(
314
  sorted_pairs[start : start + tested_batch_size],