Skip to content

rslearn.train.transforms.gaussian_noise

gaussian_noise

Gaussian noise augmentation transform.

GaussianNoise

Bases: Transform

Add Gaussian noise to inputs.

Source code in rslearn/train/transforms/gaussian_noise.py
class GaussianNoise(Transform):
    """Add Gaussian noise to inputs."""

    def __init__(
        self,
        std: float,
        selectors: list[str] = ["image"],
        skip_missing: bool = False,
    ):
        """Initialize GaussianNoise.

        Args:
            selectors: inputs to augment.
            std: standard deviation of the noise.
            skip_missing: skip missing selectors.
        """
        super().__init__(skip_missing=skip_missing)
        self.selectors = selectors
        self.std = std

    def apply_image(self, image: RasterImage) -> RasterImage:
        """Add noise."""
        image.image = image.image + torch.randn_like(image.image) * self.std
        return image

    def forward(
        self, input_dict: dict[str, Any], target_dict: dict[str, Any]
    ) -> tuple[dict[str, Any], dict[str, Any]]:
        """Apply transform."""
        self.apply_fn(self.apply_image, input_dict, target_dict, self.selectors)
        return input_dict, target_dict

apply_image

apply_image(image: RasterImage) -> RasterImage

Add noise.

Source code in rslearn/train/transforms/gaussian_noise.py
def apply_image(self, image: RasterImage) -> RasterImage:
    """Add noise."""
    image.image = image.image + torch.randn_like(image.image) * self.std
    return image

forward

forward(input_dict: dict[str, Any], target_dict: dict[str, Any]) -> tuple[dict[str, Any], dict[str, Any]]

Apply transform.

Source code in rslearn/train/transforms/gaussian_noise.py
def forward(
    self, input_dict: dict[str, Any], target_dict: dict[str, Any]
) -> tuple[dict[str, Any], dict[str, Any]]:
    """Apply transform."""
    self.apply_fn(self.apply_image, input_dict, target_dict, self.selectors)
    return input_dict, target_dict