@@ -342,6 +342,24 @@ def compute_pad_width(self, spatial_shape: Sequence[int]) -> tuple[tuple[int, in
342342 return spatial_pad .compute_pad_width (spatial_shape )
343343
344344
345+ def _to_int_list (data : Sequence [int ] | int | NdarrayOrTensor ) -> list [int ]:
346+ """Coerce an ROI spec (scalar, sequence, tensor or ndarray) to a list of Python ints."""
347+ if isinstance (data , (str , bytes )):
348+ raise TypeError ("ROI specs must be integers or sequences of integers, not strings." )
349+ return [int (i ) for i in ensure_tuple (data )]
350+
351+
352+ def _broadcast_int_pair (
353+ a : Sequence [int ] | int | NdarrayOrTensor , b : Sequence [int ] | int | NdarrayOrTensor
354+ ) -> tuple [list [int ], list [int ]]:
355+ """Coerce a pair of ROI specs to two equal-length int lists, broadcasting a scalar to match."""
356+ list_a , list_b = _to_int_list (a ), _to_int_list (b )
357+ n = max (len (list_a ), len (list_b ))
358+ if len (list_a ) not in (1 , n ) or len (list_b ) not in (1 , n ):
359+ raise ValueError (f"ROI specs must have matching lengths or be scalar, got { len (list_a )} and { len (list_b )} ." )
360+ return (list_a * n if len (list_a ) == 1 else list_a ), (list_b * n if len (list_b ) == 1 else list_b )
361+
362+
345363class Crop (InvertibleTransform , LazyTransform ):
346364 """
347365 Perform crop operations on the input image.
@@ -379,31 +397,22 @@ def compute_slices(
379397 roi_slices: list of slices for each of the spatial dimensions.
380398
381399 """
382- roi_start_t : torch .Tensor
383-
384400 if roi_slices :
385401 if not all (s .step is None or s .step == 1 for s in roi_slices ):
386402 raise ValueError (f"only slice steps of 1/None are currently supported, got { roi_slices } ." )
387403 return ensure_tuple (roi_slices )
388404 else :
389405 if roi_center is not None and roi_size is not None :
390- roi_center_t = convert_to_tensor (data = roi_center , dtype = torch .int16 , wrap_sequence = True , device = "cpu" )
391- roi_size_t = convert_to_tensor (data = roi_size , dtype = torch .int16 , wrap_sequence = True , device = "cpu" )
392- _zeros = torch .zeros_like (roi_center_t )
393- half = torch .divide (roi_size_t , 2 , rounding_mode = "floor" )
394- roi_start_t = torch .maximum (roi_center_t - half , _zeros )
395- roi_end_t = torch .maximum (roi_start_t + roi_size_t , roi_start_t )
406+ centers , sizes = _broadcast_int_pair (roi_center , roi_size )
407+ starts = [max (c - s // 2 , 0 ) for c , s in zip (centers , sizes )]
408+ ends = [st + s for st , s in zip (starts , sizes )]
396409 else :
397410 if roi_start is None or roi_end is None :
398411 raise ValueError ("please specify either roi_center, roi_size or roi_start, roi_end." )
399- roi_start_t = convert_to_tensor (data = roi_start , dtype = torch .int16 , wrap_sequence = True )
400- roi_start_t = torch .maximum (roi_start_t , torch .zeros_like (roi_start_t ))
401- roi_end_t = convert_to_tensor (data = roi_end , dtype = torch .int16 , wrap_sequence = True )
402- roi_end_t = torch .maximum (roi_end_t , roi_start_t )
403- # convert to slices (accounting for 1d)
404- if roi_start_t .numel () == 1 :
405- return ensure_tuple ([slice (int (roi_start_t .item ()), int (roi_end_t .item ()))])
406- return ensure_tuple ([slice (int (s ), int (e )) for s , e in zip (roi_start_t .tolist (), roi_end_t .tolist ())])
412+ starts , ends = _broadcast_int_pair (roi_start , roi_end )
413+ starts = [max (s , 0 ) for s in starts ]
414+ # clamp each end to its own start so no slice has negative width
415+ return ensure_tuple (slice (s , max (e , s )) for s , e in zip (starts , ends ))
407416
408417 def __call__ ( # type: ignore[override]
409418 self , img : torch .Tensor , slices : tuple [slice , ...], lazy : bool | None = None
0 commit comments