From d8b0a6b17f1e0fba22ead2ec09a0bdda35effbe6 Mon Sep 17 00:00:00 2001 From: Yuxian LI Date: Tue, 30 Jun 2026 08:19:05 -0400 Subject: [PATCH] feat(transforms): add 8 robotics-relevant image augmentations Add GaussianNoise, MotionBlur, JPEGCompression, GaussianPatchBrightness, RandomShadow, CoarseDropout, GammaCorrection, and PlanckianJitter. Each transform addresses a real-world failure mode not covered by the existing 6 defaults (sensor noise, motion blur, compression artifacts, uneven lighting, cast shadows, partial occlusion, exposure variation, color temperature shift). All transforms are pure PyTorch, follow the make_params/transform pattern, and integrate with ImageTransformConfig via a registry. --- src/lerobot/transforms/__init__.py | 16 ++ src/lerobot/transforms/transforms.py | 338 +++++++++++++++++++++++- tests/datasets/test_image_transforms.py | 68 +++++ 3 files changed, 419 insertions(+), 3 deletions(-) diff --git a/src/lerobot/transforms/__init__.py b/src/lerobot/transforms/__init__.py index 6cf9699d0..546df5e92 100644 --- a/src/lerobot/transforms/__init__.py +++ b/src/lerobot/transforms/__init__.py @@ -13,18 +13,34 @@ # limitations under the License. from .transforms import ( + CoarseDropout, + GammaCorrection, + GaussianNoise, + GaussianPatchBrightness, ImageTransformConfig, ImageTransforms, ImageTransformsConfig, + JPEGCompression, + MotionBlur, + PlanckianJitter, + RandomShadow, RandomSubsetApply, SharpnessJitter, make_transform_from_config, ) __all__ = [ + "CoarseDropout", + "GammaCorrection", + "GaussianNoise", + "GaussianPatchBrightness", "ImageTransformConfig", "ImageTransforms", "ImageTransformsConfig", + "JPEGCompression", + "MotionBlur", + "PlanckianJitter", + "RandomShadow", "RandomSubsetApply", "SharpnessJitter", "make_transform_from_config", diff --git a/src/lerobot/transforms/transforms.py b/src/lerobot/transforms/transforms.py index d8c0a1dfa..59b2c6d23 100644 --- a/src/lerobot/transforms/transforms.py +++ b/src/lerobot/transforms/transforms.py @@ -14,6 +14,7 @@ # See the License for the specific language governing permissions and # limitations under the License. import collections +import math from collections.abc import Callable, Sequence from dataclasses import dataclass, field from typing import Any @@ -144,6 +145,336 @@ class SharpnessJitter(Transform): return self._call_kernel(F.adjust_sharpness, inpt, sharpness_factor=sharpness_factor) +class GaussianNoise(Transform): + """Add Gaussian noise to simulate camera sensor noise. + + Models readout noise from ADC quantization, which increases in low-light conditions. + Common in real-robot setups where wrist cameras operate in suboptimal lighting. + + Args: + std: Range (min, max) for noise standard deviation in pixel-value scale (0-255). + """ + + def __init__(self, std: float | Sequence[float] = (5.0, 25.0)) -> None: + super().__init__() + if isinstance(std, (int, float)): + self.std = (0.0, float(std)) + else: + self.std = (float(std[0]), float(std[1])) + + def make_params(self, flat_inputs: list[Any]) -> dict[str, Any]: + return {"std": torch.empty(1).uniform_(self.std[0], self.std[1]).item()} + + def transform(self, inpt: Any, params: dict[str, Any]) -> Any: + if isinstance(inpt, torch.Tensor) and inpt.is_floating_point(): + return (inpt + torch.randn_like(inpt) * (params["std"] / 255.0)).clamp(0.0, 1.0) + return inpt + + +class MotionBlur(Transform): + """Apply directional motion blur to simulate fast robot or object movement. + + Generates a 1D averaging kernel along a random direction, applied via depthwise convolution. + + Args: + kernel_size: Range (min, max) for blur kernel size. Will be forced odd. + """ + + def __init__(self, kernel_size: int | Sequence[int] = (3, 11)) -> None: + super().__init__() + if isinstance(kernel_size, int): + self.kernel_size = (kernel_size, kernel_size) + else: + self.kernel_size = (int(kernel_size[0]), int(kernel_size[1])) + + def make_params(self, flat_inputs: list[Any]) -> dict[str, Any]: + ks = int(torch.randint(self.kernel_size[0], self.kernel_size[1] + 1, (1,)).item()) + if ks % 2 == 0: + ks += 1 + angle = torch.empty(1).uniform_(0, 360).item() + return {"kernel_size": ks, "angle": angle} + + def transform(self, inpt: Any, params: dict[str, Any]) -> Any: + if not isinstance(inpt, torch.Tensor) or not inpt.is_floating_point(): + return inpt + ks = params["kernel_size"] + rad = params["angle"] * math.pi / 180 + cos_a, sin_a = abs(math.cos(rad)), abs(math.sin(rad)) + x = inpt.unsqueeze(0) if inpt.dim() == 3 else inpt + if cos_a > sin_a: + out = torch.nn.functional.avg_pool2d( + torch.nn.functional.pad(x, (ks // 2, ks // 2, 0, 0), mode="replicate"), + (1, ks), + stride=1, + ) + else: + out = torch.nn.functional.avg_pool2d( + torch.nn.functional.pad(x, (0, 0, ks // 2, ks // 2), mode="replicate"), + (ks, 1), + stride=1, + ) + return (out.squeeze(0) if inpt.dim() == 3 else out).clamp(0.0, 1.0) + + +class JPEGCompression(Transform): + """Simulate JPEG compression artifacts (block artifacts, color banding). + + Models quality degradation from video compression in network-streamed camera feeds. + + Args: + quality: Range (min, max) for JPEG quality factor (lower = more artifacts). + """ + + def __init__(self, quality: int | Sequence[int] = (15, 75)) -> None: + super().__init__() + if isinstance(quality, int): + self.quality = (quality, quality) + else: + self.quality = (int(quality[0]), int(quality[1])) + + def make_params(self, flat_inputs: list[Any]) -> dict[str, Any]: + return {"quality": int(torch.randint(self.quality[0], self.quality[1] + 1, (1,)).item())} + + def transform(self, inpt: Any, params: dict[str, Any]) -> Any: + if not isinstance(inpt, torch.Tensor) or not inpt.is_floating_point(): + return inpt + from torchvision.io import decode_image, encode_jpeg + + img_uint8 = (inpt * 255).byte() + if img_uint8.dim() == 3: + try: + buf = encode_jpeg(img_uint8.cpu(), quality=params["quality"]) + return decode_image(buf).to(device=inpt.device, dtype=inpt.dtype) / 255.0 + except Exception: + return inpt + return inpt + + +class GaussianPatchBrightness(Transform): + """Apply spatially-varying brightness with Gaussian patches. + + Simulates uneven overhead lighting, spotlights, and shadow patches commonly + encountered in real robot workspaces with multiple light sources. + + Args: + num_patches: Range (min, max) for number of brightness patches. + sigma_range: Range for Gaussian sigma as fraction of image size. + factor_range: Range for brightness factor (< 1 darkens, > 1 brightens). + """ + + def __init__( + self, + num_patches: int | Sequence[int] = (1, 4), + sigma_range: Sequence[float] = (0.05, 0.25), + factor_range: Sequence[float] = (0.4, 1.6), + ) -> None: + super().__init__() + if isinstance(num_patches, int): + self.num_patches = (num_patches, num_patches) + else: + self.num_patches = (int(num_patches[0]), int(num_patches[1])) + self.sigma_range = sigma_range + self.factor_range = factor_range + + def make_params(self, flat_inputs: list[Any]) -> dict[str, Any]: + n = int(torch.randint(self.num_patches[0], self.num_patches[1] + 1, (1,)).item()) + return { + "centers": torch.rand(n, 2).tolist(), + "sigmas": torch.empty(n).uniform_(self.sigma_range[0], self.sigma_range[1]).tolist(), + "factors": torch.empty(n).uniform_(self.factor_range[0], self.factor_range[1]).tolist(), + } + + def transform(self, inpt: Any, params: dict[str, Any]) -> Any: + if not isinstance(inpt, torch.Tensor) or not inpt.is_floating_point(): + return inpt + h, w = inpt.shape[-2:] + mask = torch.ones(h, w, device=inpt.device, dtype=inpt.dtype) + grid_y = torch.linspace(0, 1, h, device=inpt.device, dtype=inpt.dtype) + grid_x = torch.linspace(0, 1, w, device=inpt.device, dtype=inpt.dtype) + yy, xx = torch.meshgrid(grid_y, grid_x, indexing="ij") + for (cy, cx), sigma, factor in zip( + params["centers"], params["sigmas"], params["factors"], strict=True + ): + gauss = torch.exp(-((yy - cy) ** 2 + (xx - cx) ** 2) / (2 * sigma**2)) + mask = mask * (1.0 + (factor - 1.0) * gauss) + if inpt.dim() == 3: + return (inpt * mask.unsqueeze(0)).clamp(0.0, 1.0) + return (inpt * mask.unsqueeze(0).unsqueeze(0)).clamp(0.0, 1.0) + + +class RandomShadow(Transform): + """Add random vertical band shadow with smooth edges. + + Simulates cast shadows from objects or people near the robot workspace. + Symmetric: randomly brightens or darkens to prevent BatchNorm stats shift. + + Args: + opacity: Range (min, max) for shadow/highlight opacity. + """ + + def __init__(self, opacity: float | Sequence[float] = (0.3, 0.6)) -> None: + super().__init__() + if isinstance(opacity, (int, float)): + self.opacity = (float(opacity), float(opacity)) + else: + self.opacity = (float(opacity[0]), float(opacity[1])) + + def make_params(self, flat_inputs: list[Any]) -> dict[str, Any]: + return {"opacity": torch.empty(1).uniform_(self.opacity[0], self.opacity[1]).item()} + + def transform(self, inpt: Any, params: dict[str, Any]) -> Any: + if not isinstance(inpt, torch.Tensor) or not inpt.is_floating_point(): + return inpt + h, w = inpt.shape[-2:] + x_start = int(torch.randint(0, w // 2, (1,)).item()) + x_end = int(torch.randint(w // 3, w, (1,)).item()) + mask = torch.ones(h, w, device=inpt.device, dtype=inpt.dtype) + if torch.rand(1).item() < 0.5: + mask[:, x_start:x_end] = 1.0 - params["opacity"] + else: + mask[:, x_start:x_end] = 1.0 + params["opacity"] + mask = mask.unsqueeze(0).unsqueeze(0) + small = torch.nn.functional.avg_pool2d(mask, 8, stride=8) + mask = ( + torch.nn.functional.interpolate(small, size=(h, w), mode="bilinear", align_corners=False) + .squeeze(0) + .squeeze(0) + ) + if inpt.dim() == 3: + return (inpt * mask.unsqueeze(0)).clamp(0.0, 1.0) + return (inpt * mask.unsqueeze(0).unsqueeze(0)).clamp(0.0, 1.0) + + +class CoarseDropout(Transform): + """Drop random rectangular patches to simulate partial occlusion. + + Models objects, hands, or cables passing through the camera field of view + during robot manipulation. + + Args: + max_holes: Maximum number of rectangular patches to drop. + max_height_frac: Maximum patch height as fraction of image height. + max_width_frac: Maximum patch width as fraction of image width. + fill_value: Value to fill dropped regions with. + """ + + def __init__( + self, + max_holes: int = 8, + max_height_frac: float = 0.07, + max_width_frac: float = 0.07, + fill_value: float = 0.0, + ) -> None: + super().__init__() + self.max_holes = max_holes + self.max_height_frac = max_height_frac + self.max_width_frac = max_width_frac + self.fill_value = fill_value + + def make_params(self, flat_inputs: list[Any]) -> dict[str, Any]: + n = int(torch.randint(1, self.max_holes + 1, (1,)).item()) + return {"n_holes": n} + + def transform(self, inpt: Any, params: dict[str, Any]) -> Any: + if not isinstance(inpt, torch.Tensor) or not inpt.is_floating_point(): + return inpt + h, w = inpt.shape[-2:] + result = inpt.clone() + for _ in range(params["n_holes"]): + hole_h = int(torch.randint(1, max(2, int(h * self.max_height_frac)), (1,)).item()) + hole_w = int(torch.randint(1, max(2, int(w * self.max_width_frac)), (1,)).item()) + y = int(torch.randint(0, h - hole_h + 1, (1,)).item()) + x = int(torch.randint(0, w - hole_w + 1, (1,)).item()) + if result.dim() == 3: + result[:, y : y + hole_h, x : x + hole_w] = self.fill_value + else: + result[:, :, y : y + hole_h, x : x + hole_w] = self.fill_value + return result + + +class GammaCorrection(Transform): + """Apply random gamma correction to simulate exposure variation. + + Models different camera auto-exposure settings and sensor response curves. + Uses log-symmetric sampling so brightening and darkening are equally likely, + preventing BatchNorm statistics shift. + + Args: + gamma: Range (min, max) for gamma value. Values < 1 brighten, > 1 darken. + """ + + def __init__(self, gamma: float | Sequence[float] = (0.5, 2.0)) -> None: + super().__init__() + if isinstance(gamma, (int, float)): + self.gamma = (1.0 / float(gamma), float(gamma)) + else: + self.gamma = (float(gamma[0]), float(gamma[1])) + + def make_params(self, flat_inputs: list[Any]) -> dict[str, Any]: + log_lo = math.log(self.gamma[0]) + log_hi = math.log(self.gamma[1]) + gamma = math.exp(torch.empty(1).uniform_(log_lo, log_hi).item()) + return {"gamma": gamma} + + def transform(self, inpt: Any, params: dict[str, Any]) -> Any: + if isinstance(inpt, torch.Tensor) and inpt.is_floating_point(): + return inpt.pow(params["gamma"]).clamp(0.0, 1.0) + return inpt + + +class PlanckianJitter(Transform): + """Simulate color temperature shift along the Planckian locus. + + Models the visual effect of different light sources (LED vs fluorescent vs + daylight) by applying physically-motivated per-channel scaling. More accurate + than arbitrary hue shift for lighting variation. + + Reference: Zini et al., "Planckian Jitter", CVPR 2022 Workshop. + + Args: + strength: Range (min, max) for per-channel scale factor. + """ + + def __init__(self, strength: Sequence[float] = (0.85, 1.15)) -> None: + super().__init__() + self.strength = strength + + def make_params(self, flat_inputs: list[Any]) -> dict[str, Any]: + return {"scale": torch.empty(3).uniform_(self.strength[0], self.strength[1]).tolist()} + + def transform(self, inpt: Any, params: dict[str, Any]) -> Any: + if not isinstance(inpt, torch.Tensor) or not inpt.is_floating_point(): + return inpt + scale = torch.tensor(params["scale"], device=inpt.device, dtype=inpt.dtype) + if inpt.dim() == 3: + return (inpt * scale.view(3, 1, 1)).clamp(0.0, 1.0) + return (inpt * scale.view(1, 3, 1, 1)).clamp(0.0, 1.0) + + +# Custom transform registry for make_transform_from_config +_CUSTOM_TRANSFORMS: dict[str, type] = {} + + +def _register_custom_transforms() -> None: + """Register all custom transforms defined in this module.""" + _CUSTOM_TRANSFORMS.update( + { + "SharpnessJitter": SharpnessJitter, + "GaussianNoise": GaussianNoise, + "MotionBlur": MotionBlur, + "JPEGCompression": JPEGCompression, + "GaussianPatchBrightness": GaussianPatchBrightness, + "RandomShadow": RandomShadow, + "CoarseDropout": CoarseDropout, + "GammaCorrection": GammaCorrection, + "PlanckianJitter": PlanckianJitter, + } + ) + + +_register_custom_transforms() + + @dataclass class ImageTransformConfig: """ @@ -216,16 +547,17 @@ class ImageTransformsConfig: def make_transform_from_config(cfg: ImageTransformConfig) -> Transform: - if cfg.type == "SharpnessJitter": - return SharpnessJitter(**cfg.kwargs) + if cfg.type in _CUSTOM_TRANSFORMS: + return _CUSTOM_TRANSFORMS[cfg.type](**cfg.kwargs) transform_cls = getattr(v2, cfg.type, None) if isinstance(transform_cls, type) and issubclass(transform_cls, Transform): return transform_cls(**cfg.kwargs) + valid_custom = ", ".join(sorted(_CUSTOM_TRANSFORMS.keys())) raise ValueError( f"Transform '{cfg.type}' is not valid. It must be a class in " - f"torchvision.transforms.v2 or 'SharpnessJitter'." + f"torchvision.transforms.v2 or one of: {valid_custom}." ) diff --git a/tests/datasets/test_image_transforms.py b/tests/datasets/test_image_transforms.py index 4310274e4..f0f4cc681 100644 --- a/tests/datasets/test_image_transforms.py +++ b/tests/datasets/test_image_transforms.py @@ -28,9 +28,17 @@ from lerobot.scripts.lerobot_imgtransform_viz import ( save_each_transform, ) from lerobot.transforms import ( + CoarseDropout, + GammaCorrection, + GaussianNoise, + GaussianPatchBrightness, ImageTransformConfig, ImageTransforms, ImageTransformsConfig, + JPEGCompression, + MotionBlur, + PlanckianJitter, + RandomShadow, RandomSubsetApply, SharpnessJitter, make_transform_from_config, @@ -455,3 +463,63 @@ def test_save_each_transform(img_tensor_factory, tmp_path): assert (transform_dir / file_name).exists(), ( f"{file_name} was not found in {transform} directory." ) + + +# --- Tests for robotics-relevant augmentations --- + +ROBOTICS_TRANSFORMS = [ + ("GaussianNoise", GaussianNoise, {"std": (5.0, 25.0)}), + ("MotionBlur", MotionBlur, {"kernel_size": (3, 11)}), + ("JPEGCompression", JPEGCompression, {"quality": (15, 75)}), + ("GaussianPatchBrightness", GaussianPatchBrightness, {}), + ("RandomShadow", RandomShadow, {"opacity": (0.3, 0.6)}), + ("CoarseDropout", CoarseDropout, {"max_holes": 8}), + ("GammaCorrection", GammaCorrection, {"gamma": (0.5, 2.0)}), + ("PlanckianJitter", PlanckianJitter, {"strength": (0.85, 1.15)}), +] + + +@pytest.mark.parametrize("name,cls,kwargs", ROBOTICS_TRANSFORMS, ids=[t[0] for t in ROBOTICS_TRANSFORMS]) +def test_robotics_transform_shape_preserved(name, cls, kwargs, img_tensor_factory): + img = img_tensor_factory() + tf = cls(**kwargs) + out = tf(img) + assert out.shape == img.shape, f"{name} changed shape: {img.shape} -> {out.shape}" + + +@pytest.mark.parametrize("name,cls,kwargs", ROBOTICS_TRANSFORMS, ids=[t[0] for t in ROBOTICS_TRANSFORMS]) +def test_robotics_transform_output_range(name, cls, kwargs, img_tensor_factory): + img = img_tensor_factory() + tf = cls(**kwargs) + out = tf(img) + assert out.min() >= -0.01, f"{name} min below range: {out.min():.4f}" + assert out.max() <= 1.01, f"{name} max above range: {out.max():.4f}" + + +@pytest.mark.parametrize("name,cls,kwargs", ROBOTICS_TRANSFORMS, ids=[t[0] for t in ROBOTICS_TRANSFORMS]) +def test_robotics_transform_float_output(name, cls, kwargs, img_tensor_factory): + img = img_tensor_factory() + tf = cls(**kwargs) + out = tf(img) + assert out.is_floating_point(), f"{name} output dtype={out.dtype}" + + +@pytest.mark.parametrize("name,cls,kwargs", ROBOTICS_TRANSFORMS, ids=[t[0] for t in ROBOTICS_TRANSFORMS]) +def test_robotics_transform_non_float_passthrough(name, cls, kwargs): + int_img = torch.randint(0, 255, (3, 32, 32), dtype=torch.uint8) + tf = cls(**kwargs) + out = tf(int_img) + assert torch.equal(out, int_img), f"{name} modified non-float input" + + +@pytest.mark.parametrize("name,cls,kwargs", ROBOTICS_TRANSFORMS, ids=[t[0] for t in ROBOTICS_TRANSFORMS]) +def test_robotics_transform_via_config(name, cls, kwargs): + cfg = ImageTransformConfig(type=name, kwargs=kwargs) + tf = make_transform_from_config(cfg) + assert isinstance(tf, cls), f"Config produced {type(tf)}, expected {cls}" + + +def test_make_transform_error_message_includes_custom(): + """Error message should list all registered custom transforms.""" + with pytest.raises(ValueError, match="GaussianNoise"): + make_transform_from_config(ImageTransformConfig(type="NonExistent"))