@@ -34,7 +34,7 @@ def import_mpl_plt():
3434 return mpl , plt
3535
3636
37- def rotate (image , radiological = True , n = - 1 ):
37+ def rotate (image : np . ndarray , * , radiological : bool = True , n : int = - 1 ) -> np . ndarray :
3838 # Rotate for visualization purposes
3939 image = np .rot90 (image , n , axes = (0 , 1 ))
4040 if radiological :
@@ -93,13 +93,14 @@ def plot_volume(
9393 else :
9494 data = image .data [np .newaxis , channel ]
9595 data = rearrange (data , 'c x y z -> x y z c' )
96+ data_numpy : np .ndarray = data .cpu ().numpy ()
9697
9798 if indices is None :
98- indices = np .array (data .shape [:3 ]) // 2
99+ indices = np .array (data_numpy .shape [:3 ]) // 2
99100 i , j , k = indices
100- slice_x = rotate (data [i , :, :], radiological = radiological )
101- slice_y = rotate (data [:, j , :], radiological = radiological )
102- slice_z = rotate (data [:, :, k ], radiological = radiological )
101+ slice_x = rotate (data_numpy [i , :, :], radiological = radiological )
102+ slice_y = rotate (data_numpy [:, j , :], radiological = radiological )
103+ slice_z = rotate (data_numpy [:, :, k ], radiological = radiological )
103104
104105 if isinstance (cmap , dict ):
105106 slices = slice_x , slice_y , slice_z
@@ -118,8 +119,15 @@ def plot_volume(
118119 sr , sa , ss = image .spacing
119120 imshow_kwargs ['origin' ] = 'lower'
120121
121- if percentiles is not None and not is_label :
122- p1 , p2 = np .percentile (data , percentiles )
122+ if not is_label :
123+ displayed_data = np .concatenate (
124+ [
125+ slice_x .flatten (),
126+ slice_y .flatten (),
127+ slice_z .flatten (),
128+ ]
129+ )
130+ p1 , p2 = np .percentile (displayed_data , percentiles )
123131 imshow_kwargs ['vmin' ] = p1
124132 imshow_kwargs ['vmax' ] = p2
125133
0 commit comments