Text Ranking
sentence-transformers
Safetensors
Transformers
multilingual
t5gemma2
text2text-generation
reranker
encoder-decoder
FBNL
Retrieval
RAG

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

#2
by lukann98 - opened
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],