diff --git a/rising/transforms/functional/intensity.py b/rising/transforms/functional/intensity.py index 8f01063..5db1256 100644 --- a/rising/transforms/functional/intensity.py +++ b/rising/transforms/functional/intensity.py @@ -283,12 +283,12 @@ def scale_by_value(data: torch.Tensor, value: float, out: Optional[torch.Tensor] def bezier_3rd_order( data: torch.Tensor, maxv: float = 1.0, minv: float = 0.0, out: Optional[torch.Tensor] = None ) -> torch.Tensor: - p0 = torch.zeros((1, 2)) - p1 = torch.rand((1, 2)) - p2 = torch.rand((1, 2)) - p3 = torch.ones((1, 2)) + p0 = torch.zeros((1, 2)).to(data.device) + p1 = torch.rand((1, 2)).to(data.device) + p2 = torch.rand((1, 2)).to(data.device) + p3 = torch.ones((1, 2)).to(data.device) - t = torch.linspace(0.0, 1.0, 1000).unsqueeze(1) + t = torch.linspace(0.0, 1.0, 1000).unsqueeze(1).to(data.device) points = (1 - t * t * t) * p0 + 3 * (1 - t) * (1 - t) * t * p1 + 3 * (1 - t) * t * t * p2 + t * t * t * p3 @@ -298,7 +298,7 @@ def bezier_3rd_order( xvals = points[:, 0] yvals = points[:, 1] - out_flat = Interp1d.apply(xvals, yvals, data.view(-1)) + out_flat = Interp1d.apply(xvals, yvals, data.contiguous().view(-1)) return out_flat.view(data.shape) diff --git a/rising/transforms/kernel.py b/rising/transforms/kernel.py index 2a7646e..26c070d 100644 --- a/rising/transforms/kernel.py +++ b/rising/transforms/kernel.py @@ -107,8 +107,8 @@ def forward(self, **data) -> dict: Returns: dict: dict with transformed data """ - # dtype, device = data[self.keys[0]].dtype, data[self.keys[0]].device - # self.to(dtype) + dtype, device = data[self.keys[0]].dtype, data[self.keys[0]].device + self.to(dtype) for key, padding_mode in zip(self.keys, self.padding_mode): inp_pad = F.pad(data[key], self.padding, mode=padding_mode) diff --git a/rising/transforms/spatial.py b/rising/transforms/spatial.py index b9b5b17..ff1fa12 100644 --- a/rising/transforms/spatial.py +++ b/rising/transforms/spatial.py @@ -73,9 +73,10 @@ def forward(self, **data) -> dict: prob = self.prob seed = torch.random.get_rng_state() + dims = self.dims for key in self.keys: torch.random.set_rng_state(seed) - for dim, p in zip(self.dims, prob): + for dim, p in zip(dims, prob): if torch.rand(1) < p: data[key] = mirror(data[key], dim) return data