Compare commits

...

7 Commits

Author SHA1 Message Date
Steven Palma 6c60f4651b update image 2026-07-29 16:12:39 +02:00
Steven Palma 117216b29b refactor(transforms): several updates 2026-07-29 15:43:13 +02:00
liyux 2422d43fb1 tune showcase: softer shadow, dropout, jitter intensity 2026-07-29 15:27:31 +02:00
liyux 53a5cffb4c tune showcase to balanced augmentation intensity 2026-07-29 15:27:31 +02:00
liyux 2a3f40c673 update showcase with better sample frame 2026-07-29 15:27:31 +02:00
liyux 2ddfb5a376 add augmentation showcase image for PR 2026-07-29 15:27:31 +02:00
Yuxian LI d8b0a6b17f 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.
2026-07-29 15:27:31 +02:00
4 changed files with 645 additions and 3 deletions
Binary file not shown.

After

Width:  |  Height:  |  Size: 682 KiB

+16
View File
@@ -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",
+471 -3
View File
@@ -14,11 +14,13 @@
# 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
import torch
from torchvision.io import decode_image, encode_jpeg
from torchvision.transforms import v2
from torchvision.transforms.v2 import (
Transform,
@@ -144,6 +146,471 @@ 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))
elif isinstance(std, Sequence) and len(std) == 2:
self.std = (float(std[0]), float(std[1]))
else:
raise TypeError("std must be a number or a sequence with length 2.")
if not 0.0 <= self.std[0] <= self.std[1]:
raise ValueError(f"std must satisfy 0 <= min <= max, but got {self.std}.")
def make_params(self, flat_inputs: list[Any]) -> dict[str, Any]:
return {
"std": torch.empty(1).uniform_(self.std[0], self.std[1]).item(),
"seed": torch.randint(0, torch.iinfo(torch.int64).max, ()).item(),
}
def transform(self, inpt: Any, params: dict[str, Any]) -> Any:
if isinstance(inpt, torch.Tensor) and inpt.is_floating_point():
generator = torch.Generator(device=inpt.device).manual_seed(params["seed"])
noise = torch.randn(inpt.shape, device=inpt.device, dtype=inpt.dtype, generator=generator)
return (inpt + noise * (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: An odd kernel size or a range containing at least one odd kernel size.
"""
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)
elif isinstance(kernel_size, Sequence) and len(kernel_size) == 2:
self.kernel_size = (int(kernel_size[0]), int(kernel_size[1]))
else:
raise TypeError("kernel_size must be an int or a sequence with length 2.")
if not 1 <= self.kernel_size[0] <= self.kernel_size[1]:
raise ValueError(f"kernel_size must satisfy 1 <= min <= max, but got {self.kernel_size}.")
self._first_odd_kernel_size = self.kernel_size[0] + (self.kernel_size[0] + 1) % 2
if self._first_odd_kernel_size > self.kernel_size[1]:
raise ValueError(f"kernel_size range must contain an odd value, but got {self.kernel_size}.")
def make_params(self, flat_inputs: list[Any]) -> dict[str, Any]:
num_odd_sizes = (self.kernel_size[1] - self._first_odd_kernel_size) // 2 + 1
size_index = int(torch.randint(0, num_odd_sizes, ()).item())
ks = self._first_odd_kernel_size + 2 * size_index
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
if inpt.ndim < 3:
raise ValueError(f"MotionBlur expects [..., C, H, W] input, but got shape {inpt.shape}.")
kernel_size = params["kernel_size"]
radius = kernel_size // 2
angle = math.radians(params["angle"])
positions = torch.linspace(-radius, radius, kernel_size, device=inpt.device)
x_coords = (positions * math.cos(angle)).round().to(torch.long) + radius
y_coords = (positions * math.sin(angle)).round().to(torch.long) + radius
kernel = torch.zeros((kernel_size, kernel_size), device=inpt.device, dtype=inpt.dtype)
kernel[y_coords, x_coords] = 1
kernel /= kernel.sum()
channels, height, width = inpt.shape[-3:]
flat_input = inpt.reshape(-1, channels, height, width)
depthwise_kernel = kernel.expand(channels, 1, kernel_size, kernel_size)
padded = torch.nn.functional.pad(flat_input, (radius,) * 4, mode="replicate")
output = torch.nn.functional.conv2d(padded, depthwise_kernel, groups=channels)
return output.reshape(inpt.shape).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)
elif isinstance(quality, Sequence) and len(quality) == 2:
self.quality = (int(quality[0]), int(quality[1]))
else:
raise TypeError("quality must be an int or a sequence with length 2.")
if not 1 <= self.quality[0] <= self.quality[1] <= 100:
raise ValueError(f"quality must satisfy 1 <= min <= max <= 100, but got {self.quality}.")
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
if inpt.ndim < 3:
raise ValueError(f"JPEGCompression expects [..., C, H, W] input, but got shape {inpt.shape}.")
channels, height, width = inpt.shape[-3:]
if channels not in (1, 3):
raise ValueError(f"JPEGCompression expects 1 or 3 channels, but got {channels}.")
flat_input = inpt.reshape(-1, channels, height, width)
flat_uint8 = (flat_input.clamp(0.0, 1.0) * 255).round().to(torch.uint8).cpu()
decoded_frames = [
decode_image(encode_jpeg(frame, quality=params["quality"])) for frame in flat_uint8.unbind()
]
output = torch.stack(decoded_frames).to(device=inpt.device, dtype=inpt.dtype) / 255.0
return output.reshape(inpt.shape)
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)
elif isinstance(num_patches, Sequence) and len(num_patches) == 2:
self.num_patches = (int(num_patches[0]), int(num_patches[1]))
else:
raise TypeError("num_patches must be an int or a sequence with length 2.")
if not 1 <= self.num_patches[0] <= self.num_patches[1]:
raise ValueError(f"num_patches must satisfy 1 <= min <= max, but got {self.num_patches}.")
if not isinstance(sigma_range, Sequence) or len(sigma_range) != 2:
raise TypeError("sigma_range must be a sequence with length 2.")
self.sigma_range = (float(sigma_range[0]), float(sigma_range[1]))
if not 0.0 < self.sigma_range[0] <= self.sigma_range[1]:
raise ValueError(f"sigma_range must satisfy 0 < min <= max, but got {self.sigma_range}.")
if not isinstance(factor_range, Sequence) or len(factor_range) != 2:
raise TypeError("factor_range must be a sequence with length 2.")
self.factor_range = (float(factor_range[0]), float(factor_range[1]))
if not 0.0 <= self.factor_range[0] <= self.factor_range[1]:
raise ValueError(f"factor_range must satisfy 0 <= min <= max, but got {self.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)
broadcast_shape = (1,) * (inpt.ndim - 2) + (h, w)
return (inpt * mask.reshape(broadcast_shape)).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))
elif isinstance(opacity, Sequence) and len(opacity) == 2:
self.opacity = (float(opacity[0]), float(opacity[1]))
else:
raise TypeError("opacity must be a number or a sequence with length 2.")
if not 0.0 <= self.opacity[0] <= self.opacity[1] <= 1.0:
raise ValueError(f"opacity must satisfy 0 <= min <= max <= 1, but got {self.opacity}.")
def make_params(self, flat_inputs: list[Any]) -> dict[str, Any]:
return {
"opacity": torch.empty(1).uniform_(self.opacity[0], self.opacity[1]).item(),
"start": torch.rand(1).item(),
"width": torch.empty(1).uniform_(1 / 3, 2 / 3).item(),
"direction": -1.0 if torch.rand(1).item() < 0.5 else 1.0,
}
def transform(self, inpt: Any, params: dict[str, Any]) -> Any:
if not isinstance(inpt, torch.Tensor) or not inpt.is_floating_point():
return inpt
if inpt.ndim < 3:
raise ValueError(f"RandomShadow expects [..., C, H, W] input, but got shape {inpt.shape}.")
h, w = inpt.shape[-2:]
band_width = max(1, min(w, round(params["width"] * w)))
x_start = round(params["start"] * (w - band_width))
x_end = x_start + band_width
mask = torch.ones(h, w, device=inpt.device, dtype=inpt.dtype)
mask[:, x_start:x_end] = 1.0 + params["direction"] * params["opacity"]
smoothing_size = min(8, h, w)
if smoothing_size > 1:
batched_mask = mask[None, None]
small = torch.nn.functional.avg_pool2d(batched_mask, smoothing_size, stride=smoothing_size)
mask = torch.nn.functional.interpolate(small, size=(h, w), mode="bilinear", align_corners=False)[
0, 0
]
broadcast_shape = (1,) * (inpt.ndim - 2) + (h, w)
return (inpt * mask.reshape(broadcast_shape)).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__()
if not isinstance(max_holes, int):
raise TypeError("max_holes must be an int.")
if max_holes < 1:
raise ValueError(f"max_holes must be at least 1, but got {max_holes}.")
if not 0.0 < max_height_frac <= 1.0:
raise ValueError(f"max_height_frac must be in (0, 1], but got {max_height_frac}.")
if not 0.0 < max_width_frac <= 1.0:
raise ValueError(f"max_width_frac must be in (0, 1], but got {max_width_frac}.")
if not 0.0 <= fill_value <= 1.0:
raise ValueError(f"fill_value must be in [0, 1], but got {fill_value}.")
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())
sizes = torch.rand(n, 2)
sizes[:, 0] *= self.max_height_frac
sizes[:, 1] *= self.max_width_frac
return {"sizes": sizes.tolist(), "positions": torch.rand(n, 2).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
if inpt.ndim < 3:
raise ValueError(f"CoarseDropout expects [..., C, H, W] input, but got shape {inpt.shape}.")
h, w = inpt.shape[-2:]
result = inpt.clone()
for (height_frac, width_frac), (y_frac, x_frac) in zip(
params["sizes"], params["positions"], strict=True
):
hole_h = max(1, min(h, round(height_frac * h)))
hole_w = max(1, min(w, round(width_frac * w)))
y = round(y_frac * (h - hole_h))
x = round(x_frac * (w - hole_w))
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)):
gamma = float(gamma)
if gamma <= 0:
raise ValueError(f"gamma must be positive, but got {gamma}.")
self.gamma = (min(gamma, 1.0 / gamma), max(gamma, 1.0 / gamma))
elif isinstance(gamma, Sequence) and len(gamma) == 2:
self.gamma = (float(gamma[0]), float(gamma[1]))
else:
raise TypeError("gamma must be a number or a sequence with length 2.")
if not 0.0 < self.gamma[0] <= self.gamma[1]:
raise ValueError(f"gamma must satisfy 0 < min <= max, but got {self.gamma}.")
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
# From the paper authors' MIT-licensed reference implementation:
# https://github.com/TheZino/PlanckianJitter
_PLANCKIAN_BLACKBODY_COEFFICIENTS = (
(0.6743, 0.4029, 0.0013),
(0.6281, 0.4241, 0.1665),
(0.5919, 0.4372, 0.2513),
(0.5623, 0.4457, 0.3154),
(0.5376, 0.4515, 0.3672),
(0.5163, 0.4555, 0.4103),
(0.4979, 0.4584, 0.4468),
(0.4816, 0.4604, 0.4782),
(0.4672, 0.4619, 0.5053),
(0.4542, 0.4630, 0.5289),
(0.4426, 0.4638, 0.5497),
(0.4320, 0.4644, 0.5681),
(0.4223, 0.4648, 0.5844),
(0.4135, 0.4651, 0.5990),
(0.4054, 0.4653, 0.6121),
(0.3980, 0.4654, 0.6239),
(0.3911, 0.4655, 0.6346),
(0.3847, 0.4656, 0.6444),
(0.3787, 0.4656, 0.6532),
(0.3732, 0.4656, 0.6613),
(0.3680, 0.4655, 0.6688),
(0.3632, 0.4655, 0.6756),
(0.3586, 0.4655, 0.6820),
(0.3544, 0.4654, 0.6878),
(0.3503, 0.4653, 0.6933),
)
_PLANCKIAN_MIN_TEMPERATURE = 3_000
_PLANCKIAN_MAX_TEMPERATURE = 15_000
_PLANCKIAN_TEMPERATURE_STEP = 500
class PlanckianJitter(Transform):
"""Simulate color temperature shift along the Planckian locus.
Samples one black-body temperature and applies the corresponding correlated red
and blue channel scaling while preserving the green channel. Coefficients between
the tabulated 500 K intervals are linearly interpolated.
Reference: Zini et al., "Planckian Jitter", CVPR 2022 Workshop.
Args:
temperature: A fixed color temperature or range in Kelvin. Supported values
are between 3000 K and 15000 K.
"""
def __init__(self, temperature: int | Sequence[int] = (3_000, 15_000)) -> None:
super().__init__()
if isinstance(temperature, int):
self.temperature = (temperature, temperature)
elif isinstance(temperature, Sequence) and len(temperature) == 2:
self.temperature = (int(temperature[0]), int(temperature[1]))
else:
raise TypeError("temperature must be an int or a sequence with length 2.")
if not (
_PLANCKIAN_MIN_TEMPERATURE
<= self.temperature[0]
<= self.temperature[1]
<= _PLANCKIAN_MAX_TEMPERATURE
):
raise ValueError(
"temperature must satisfy "
f"{_PLANCKIAN_MIN_TEMPERATURE} <= min <= max <= {_PLANCKIAN_MAX_TEMPERATURE}, "
f"but got {self.temperature}."
)
def make_params(self, flat_inputs: list[Any]) -> dict[str, Any]:
temperature = int(torch.randint(self.temperature[0], self.temperature[1] + 1, ()).item())
return {"temperature": temperature}
def transform(self, inpt: Any, params: dict[str, Any]) -> Any:
if not isinstance(inpt, torch.Tensor) or not inpt.is_floating_point():
return inpt
if inpt.ndim < 3 or inpt.shape[-3] != 3:
raise ValueError(f"PlanckianJitter expects [..., 3, H, W] input, but got shape {inpt.shape}.")
table_position = (params["temperature"] - _PLANCKIAN_MIN_TEMPERATURE) / _PLANCKIAN_TEMPERATURE_STEP
left_index = math.floor(table_position)
right_index = min(left_index + 1, len(_PLANCKIAN_BLACKBODY_COEFFICIENTS) - 1)
interpolation_weight = table_position - left_index
left = torch.tensor(
_PLANCKIAN_BLACKBODY_COEFFICIENTS[left_index],
device=inpt.device,
dtype=inpt.dtype,
)
right = torch.tensor(
_PLANCKIAN_BLACKBODY_COEFFICIENTS[right_index],
device=inpt.device,
dtype=inpt.dtype,
)
coefficients = torch.lerp(left, right, interpolation_weight)
scale = torch.stack(
(
coefficients[0] / coefficients[1],
coefficients.new_tensor(1.0),
coefficients[2] / coefficients[1],
)
)
broadcast_shape = (1,) * (inpt.ndim - 3) + (3, 1, 1)
return (inpt * scale.reshape(broadcast_shape)).clamp(0.0, 1.0)
_CUSTOM_TRANSFORMS: dict[str, type[Transform]] = {
"SharpnessJitter": SharpnessJitter,
"GaussianNoise": GaussianNoise,
"MotionBlur": MotionBlur,
"JPEGCompression": JPEGCompression,
"GaussianPatchBrightness": GaussianPatchBrightness,
"RandomShadow": RandomShadow,
"CoarseDropout": CoarseDropout,
"GammaCorrection": GammaCorrection,
"PlanckianJitter": PlanckianJitter,
}
@dataclass
class ImageTransformConfig:
"""
@@ -216,16 +683,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}."
)
+158
View File
@@ -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,153 @@ 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, {"temperature": (3_000, 15_000)}),
]
@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"))
@pytest.mark.parametrize("name,cls,kwargs", ROBOTICS_TRANSFORMS, ids=[t[0] for t in ROBOTICS_TRANSFORMS])
@pytest.mark.parametrize("shape", [(4, 3, 32, 32), (2, 4, 3, 16, 16)])
def test_robotics_transform_supports_temporal_batches(name, cls, kwargs, shape):
img = torch.rand(shape)
out = cls(**kwargs)(img)
assert out.shape == img.shape, f"{name} changed shape: {img.shape} -> {out.shape}"
assert out.min() >= 0
assert out.max() <= 1
@pytest.mark.parametrize(
"cls,kwargs",
[
(GaussianNoise, {"std": (25.0, 25.0)}),
(MotionBlur, {"kernel_size": 5}),
(JPEGCompression, {"quality": 10}),
(
GaussianPatchBrightness,
{"num_patches": 1, "sigma_range": (0.2, 0.2), "factor_range": (0.5, 0.5)},
),
(RandomShadow, {"opacity": 0.5}),
(CoarseDropout, {"max_holes": 1, "fill_value": 0.0}),
(GammaCorrection, {"gamma": (2.0, 2.0)}),
(PlanckianJitter, {"temperature": 3_000}),
],
)
def test_robotics_transform_is_not_silent_noop(cls, kwargs):
img = torch.rand(3, 32, 32)
out = cls(**kwargs)(img)
assert not torch.equal(out, img)
@pytest.mark.parametrize(
"transform",
[
GaussianNoise(std=25),
RandomShadow(opacity=0.5),
CoarseDropout(max_holes=4),
],
)
def test_robotics_transform_random_params_are_reused(transform):
img = torch.rand(3, 32, 32)
params = transform.make_params([img])
torch.testing.assert_close(transform.transform(img, params), transform.transform(img, params))
def test_motion_blur_kernel_size_stays_in_configured_range():
transform = MotionBlur(kernel_size=(4, 10))
sampled_sizes = {transform.make_params([])["kernel_size"] for _ in range(100)}
assert sampled_sizes <= {5, 7, 9}
assert sampled_sizes
def test_gamma_correction_scalar_below_one_defines_symmetric_range():
transform = GammaCorrection(gamma=0.5)
assert transform.gamma == (0.5, 2.0)
assert transform(torch.rand(3, 8, 8)).shape == (3, 8, 8)
def test_planckian_jitter_uses_correlated_temperature_coefficients():
img = torch.full((2, 3, 8, 8), 0.25)
out = PlanckianJitter(temperature=3_000)(img)
torch.testing.assert_close(out[:, 1], img[:, 1])
assert torch.all(out[:, 0] > out[:, 1])
assert torch.all(out[:, 2] < out[:, 1])
def test_random_shadow_supports_small_images():
img = torch.rand(3, 7, 7)
assert RandomShadow()(img).shape == img.shape
@pytest.mark.parametrize(
"cls,kwargs",
[
(GaussianNoise, {"std": (-1.0, 1.0)}),
(MotionBlur, {"kernel_size": 4}),
(JPEGCompression, {"quality": (0, 75)}),
(GaussianPatchBrightness, {"sigma_range": (0.0, 0.25)}),
(RandomShadow, {"opacity": (0.3, 1.1)}),
(CoarseDropout, {"max_holes": 0}),
(GammaCorrection, {"gamma": 0.0}),
(PlanckianJitter, {"temperature": (2_000, 6_500)}),
],
)
def test_robotics_transform_rejects_invalid_config(cls, kwargs):
with pytest.raises(ValueError):
cls(**kwargs)