diff --git a/monai/transforms/spatial/array.py b/monai/transforms/spatial/array.py index 2d3edaeed5..0635e843c6 100644 --- a/monai/transforms/spatial/array.py +++ b/monai/transforms/spatial/array.py @@ -1861,6 +1861,7 @@ def __init__( shear_range: RandRange = None, translate_range: RandRange = None, scale_range: RandRange = None, + translate_relative: bool = False, device: torch.device | None = None, dtype: DtypeLike = np.float32, lazy: bool = False, @@ -1889,6 +1890,9 @@ def __init__( scale_range: scaling range with format matching `rotate_range`. it defines the range to randomly select the scale factor to translate for every spatial dims. A value of 1.0 is added to the result. This allows 0 to correspond to no change (i.e., a scaling of 1.0). + translate_relative: if True, `translate_range` values are interpreted as fractions of the + corresponding spatial dimension size (e.g. 0.5 allows translating up to half the image + size), instead of absolute voxels. Defaults to False. device: device to store the output grid data. dtype: data type for the grid computation. Defaults to ``np.float32``. If ``None``, use the data type of input data (if `grid` is provided). @@ -1907,6 +1911,7 @@ def __init__( self.shear_range = ensure_tuple(shear_range) self.translate_range = ensure_tuple(translate_range) self.scale_range = ensure_tuple(scale_range) + self.translate_relative = translate_relative self.rotate_params: list[float] | None = None self.shear_params: list[float] | None = None @@ -1956,10 +1961,18 @@ def __call__( if randomize: self.randomize() lazy_ = self.lazy if lazy is None else lazy + translate_params = self.translate_params + if self.translate_relative and translate_params is not None: + # interpret the sampled parameters as fractions of the spatial dims + sp_size = ( + spatial_size if spatial_size is not None else (tuple(grid.shape[1:]) if grid is not None else None) + ) + if sp_size is not None: + translate_params = [p * d for p, d in zip(translate_params, sp_size)] affine_grid = AffineGrid( rotate_params=self.rotate_params, shear_params=self.shear_params, - translate_params=self.translate_params, + translate_params=translate_params, scale_params=self.scale_params, device=self.device, dtype=self.dtype, @@ -2455,6 +2468,7 @@ def __init__( shear_range: RandRange = None, translate_range: RandRange = None, scale_range: RandRange = None, + translate_relative: bool = False, spatial_size: Sequence[int] | int | None = None, mode: str | int = GridSampleMode.BILINEAR, padding_mode: str = GridSamplePadMode.REFLECTION, @@ -2488,6 +2502,9 @@ def __init__( scale_range: scaling range with format matching `rotate_range`. it defines the range to randomly select the scale factor to translate for every spatial dims. A value of 1.0 is added to the result. This allows 0 to correspond to no change (i.e., a scaling of 1.0). + translate_relative: if True, `translate_range` values are interpreted as fractions of the + corresponding spatial dimension size (e.g. 0.5 allows translating up to half the image + size), instead of absolute voxels. Defaults to False. spatial_size: output image spatial size. if `spatial_size` and `self.spatial_size` are not defined, or smaller than 1, the transform will use the spatial size of `img`. @@ -2532,6 +2549,7 @@ def __init__( shear_range=shear_range, translate_range=translate_range, scale_range=scale_range, + translate_relative=translate_relative, device=device, lazy=lazy, ) diff --git a/monai/transforms/spatial/dictionary.py b/monai/transforms/spatial/dictionary.py index 5830737711..4bd822457d 100644 --- a/monai/transforms/spatial/dictionary.py +++ b/monai/transforms/spatial/dictionary.py @@ -1048,6 +1048,7 @@ def __init__( shear_range: Sequence[tuple[float, float] | float] | float | None = None, translate_range: Sequence[tuple[float, float] | float] | float | None = None, scale_range: Sequence[tuple[float, float] | float] | float | None = None, + translate_relative: bool = False, mode: SequenceStr = GridSampleMode.BILINEAR, padding_mode: SequenceStr = GridSamplePadMode.REFLECTION, cache_grid: bool = False, @@ -1088,6 +1089,9 @@ def __init__( scale_range: scaling range with format matching `rotate_range`. it defines the range to randomly select the scale factor to translate for every spatial dims. A value of 1.0 is added to the result. This allows 0 to correspond to no change (i.e., a scaling of 1.0). + translate_relative: if True, `translate_range` values are interpreted as fractions of the + corresponding spatial dimension size (e.g. 0.5 allows translating up to half the image + size), instead of absolute voxels. Defaults to False. mode: {``"bilinear"``, ``"nearest"``} or spline interpolation order 0-5 (integers). Interpolation mode to calculate output values. Defaults to ``"bilinear"``. See also: https://pytorch.org/docs/stable/generated/torch.nn.functional.grid_sample.html @@ -1131,6 +1135,7 @@ def __init__( shear_range=shear_range, translate_range=translate_range, scale_range=scale_range, + translate_relative=translate_relative, spatial_size=spatial_size, cache_grid=cache_grid, device=device,