@@ -479,6 +479,8 @@ def _calculate_max_slices_iterative(
479479 non_slice_dims_shape : Tuple [int , int ],
480480 available_memory : int ,
481481 ) -> int :
482+ MEM_RATIO_THRESHOLD = 0.9
483+
482484 def get_mem_bytes (current_slices ):
483485 try :
484486 memory_bytes = self ._query .calculate_memory_bytes_for_slices (
@@ -501,14 +503,20 @@ def get_mem_bytes(current_slices):
501503 slices_high = None
502504 memory_bytes = get_mem_bytes (current_slices )
503505 if memory_bytes > available_memory :
506+ # Found upper limit, continue to binary search
504507 slices_high = current_slices
505508 else :
506509 # linear approximation
507510 current_slices = int (current_slices * available_memory / memory_bytes )
508511 while True :
509512 memory_bytes = get_mem_bytes (current_slices )
510513 if memory_bytes > available_memory :
514+ # Found upper limit, continue to binary search
511515 break
516+ elif memory_bytes >= available_memory * MEM_RATIO_THRESHOLD :
517+ # This is "good enough", return
518+ return current_slices
519+
512520 # If linear approximation is not enough, just double every iteration
513521 current_slices *= 2
514522 slices_high = current_slices
@@ -520,10 +528,10 @@ def get_mem_bytes(current_slices):
520528 memory_bytes = get_mem_bytes (current_slices )
521529 if memory_bytes > available_memory :
522530 slices_high = current_slices
523- elif memory_bytes < available_memory :
524- slices_low = current_slices
525- else : # memory_bytes == available_memory
531+ elif memory_bytes >= available_memory * MEM_RATIO_THRESHOLD :
526532 return current_slices
533+ else :
534+ slices_low = current_slices
527535
528536 return slices_low
529537
0 commit comments