|
30 | 30 | from unittest.mock import Mock |
31 | 31 |
|
32 | 32 | if cupy_run: |
| 33 | + from tomobar.supp.memory_estimator_helpers import DeviceMemStack |
33 | 34 | from cupyx.scipy.ndimage import median_filter, binary_dilation, uniform_filter1d |
34 | 35 | from cupyx.scipy.fft import fft2, ifft2, fftshift |
35 | 36 | from cupyx.scipy.fftpack import get_fft_plan |
@@ -226,32 +227,8 @@ def _reflect(x: np.ndarray, minx: float, maxx: float) -> np.ndarray: |
226 | 227 | return np.array(out, dtype=x.dtype) |
227 | 228 |
|
228 | 229 |
|
229 | | -class _DeviceMemStack: |
230 | | - def __init__(self) -> None: |
231 | | - self.allocations = [] |
232 | | - self.current = 0 |
233 | | - self.highwater = 0 |
234 | | - |
235 | | - def malloc(self, bytes): |
236 | | - self.allocations.append(bytes) |
237 | | - allocated = self._round_up(bytes) |
238 | | - self.current += allocated |
239 | | - self.highwater = max(self.current, self.highwater) |
240 | | - |
241 | | - def free(self, bytes): |
242 | | - assert bytes in self.allocations |
243 | | - self.allocations.remove(bytes) |
244 | | - self.current -= self._round_up(bytes) |
245 | | - assert self.current >= 0 |
246 | | - |
247 | | - def _round_up(self, size): |
248 | | - ALLOCATION_UNIT_SIZE = 512 |
249 | | - size = (size + ALLOCATION_UNIT_SIZE - 1) // ALLOCATION_UNIT_SIZE |
250 | | - return size * ALLOCATION_UNIT_SIZE |
251 | | - |
252 | | - |
253 | 230 | def _mypad( |
254 | | - x: cp.ndarray, pad: Tuple[int, int, int, int], mem_stack: Optional[_DeviceMemStack] |
| 231 | + x: cp.ndarray, pad: Tuple[int, int, int, int], mem_stack: Optional[DeviceMemStack] |
255 | 232 | ) -> cp.ndarray: |
256 | 233 | """Function to do numpy like padding on Arrays. Only works for 2-D |
257 | 234 | padding. |
@@ -294,7 +271,7 @@ def _conv2d( |
294 | 271 | w: np.ndarray, |
295 | 272 | stride: Tuple[int, int], |
296 | 273 | groups: int, |
297 | | - mem_stack: Optional[_DeviceMemStack], |
| 274 | + mem_stack: Optional[DeviceMemStack], |
298 | 275 | ) -> cp.ndarray: |
299 | 276 | """Convolution (equivalent pytorch.conv2d)""" |
300 | 277 | b, ci, hi, wi = x.shape if not mem_stack else x |
@@ -378,7 +355,7 @@ def _conv_transpose2d( |
378 | 355 | stride: Tuple[int, int], |
379 | 356 | pad: Tuple[int, int], |
380 | 357 | groups: int, |
381 | | - mem_stack: Optional[_DeviceMemStack], |
| 358 | + mem_stack: Optional[DeviceMemStack], |
382 | 359 | ) -> cp.ndarray: |
383 | 360 | """Transposed convolution (equivalent pytorch.conv_transpose2d)""" |
384 | 361 | b, co, ho, wo = x.shape if not mem_stack else x |
@@ -473,7 +450,7 @@ def _afb1d( |
473 | 450 | h0: np.ndarray, |
474 | 451 | h1: np.ndarray, |
475 | 452 | dim: int, |
476 | | - mem_stack: Optional[_DeviceMemStack], |
| 453 | + mem_stack: Optional[DeviceMemStack], |
477 | 454 | ) -> cp.ndarray: |
478 | 455 | """1D analysis filter bank (along one dimension only) of an image |
479 | 456 |
|
@@ -521,7 +498,7 @@ def _sfb1d( |
521 | 498 | g0: np.ndarray, |
522 | 499 | g1: np.ndarray, |
523 | 500 | dim: int, |
524 | | - mem_stack: Optional[_DeviceMemStack], |
| 501 | + mem_stack: Optional[DeviceMemStack], |
525 | 502 | ) -> cp.ndarray: |
526 | 503 | """1D synthesis filter bank of an image Array""" |
527 | 504 |
|
@@ -565,7 +542,7 @@ def __init__(self, wave: str): |
565 | 542 | self.h1_row = np.array(h1_row).astype("float32")[::-1].reshape((1, 1, 1, -1)) |
566 | 543 |
|
567 | 544 | def apply( |
568 | | - self, x: cp.ndarray, mem_stack: Optional[_DeviceMemStack] = None |
| 545 | + self, x: cp.ndarray, mem_stack: Optional[DeviceMemStack] = None |
569 | 546 | ) -> Tuple[cp.ndarray, cp.ndarray]: |
570 | 547 | """Forward pass of the DWT. |
571 | 548 |
|
@@ -627,7 +604,7 @@ def __init__(self, wave: str): |
627 | 604 | def apply( |
628 | 605 | self, |
629 | 606 | coeffs: Tuple[cp.ndarray, cp.ndarray], |
630 | | - mem_stack: Optional[_DeviceMemStack] = None, |
| 607 | + mem_stack: Optional[DeviceMemStack] = None, |
631 | 608 | ) -> cp.ndarray: |
632 | 609 | """ |
633 | 610 | Args: |
@@ -728,7 +705,7 @@ def remove_stripe_fw( |
728 | 705 | sli_shape = [nz, 1, nproj_pad, ni] |
729 | 706 |
|
730 | 707 | if calc_peak_gpu_mem: |
731 | | - mem_stack = _DeviceMemStack() |
| 708 | + mem_stack = DeviceMemStack() |
732 | 709 | # A data copy is assumed when invoking the function |
733 | 710 | mem_stack.malloc(np.prod(data) * np.float32().itemsize) |
734 | 711 | mem_stack.malloc(np.prod(sli_shape) * np.float32().itemsize) |
|
0 commit comments