Skip to content

Commit 43b46ca

Browse files
committed
Reduce the number of required iterations by introducing a "good enough" threshold
1 parent 0164511 commit 43b46ca

2 files changed

Lines changed: 14 additions & 5 deletions

File tree

‎httomo/method_wrappers/generic.py‎

Lines changed: 11 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -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

‎tests/method_wrappers/test_generic.py‎

Lines changed: 3 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -769,6 +769,7 @@ def test_method(data):
769769
check_slices = lambda slices: memcalc_mock(
770770
dims_shape=(slices, shape[0], shape[1]), dtype=dummy_block.data.dtype
771771
)
772+
threshold = 0.9
772773
if check_slices(1) > available_memory:
773774
# If zero slice fits
774775
assert max_slices == 0
@@ -780,8 +781,8 @@ def test_method(data):
780781
with pytest.raises(Exception):
781782
check_slices(max_slices + 1)
782783
else:
783-
# And one more slice must not fit
784-
assert check_slices(max_slices + 1) > available_memory
784+
# And one more slice must not fit OR above threshold
785+
assert check_slices(max_slices + 1) > available_memory or check_slices(max_slices) >= available_memory * threshold
785786

786787

787788
@pytest.mark.cupy

0 commit comments

Comments
 (0)