mirror of
https://github.com/huggingface/lerobot.git
synced 2026-07-29 20:49:42 +00:00
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.
This commit is contained in:
@@ -13,18 +13,34 @@
|
|||||||
# limitations under the License.
|
# limitations under the License.
|
||||||
|
|
||||||
from .transforms import (
|
from .transforms import (
|
||||||
|
CoarseDropout,
|
||||||
|
GammaCorrection,
|
||||||
|
GaussianNoise,
|
||||||
|
GaussianPatchBrightness,
|
||||||
ImageTransformConfig,
|
ImageTransformConfig,
|
||||||
ImageTransforms,
|
ImageTransforms,
|
||||||
ImageTransformsConfig,
|
ImageTransformsConfig,
|
||||||
|
JPEGCompression,
|
||||||
|
MotionBlur,
|
||||||
|
PlanckianJitter,
|
||||||
|
RandomShadow,
|
||||||
RandomSubsetApply,
|
RandomSubsetApply,
|
||||||
SharpnessJitter,
|
SharpnessJitter,
|
||||||
make_transform_from_config,
|
make_transform_from_config,
|
||||||
)
|
)
|
||||||
|
|
||||||
__all__ = [
|
__all__ = [
|
||||||
|
"CoarseDropout",
|
||||||
|
"GammaCorrection",
|
||||||
|
"GaussianNoise",
|
||||||
|
"GaussianPatchBrightness",
|
||||||
"ImageTransformConfig",
|
"ImageTransformConfig",
|
||||||
"ImageTransforms",
|
"ImageTransforms",
|
||||||
"ImageTransformsConfig",
|
"ImageTransformsConfig",
|
||||||
|
"JPEGCompression",
|
||||||
|
"MotionBlur",
|
||||||
|
"PlanckianJitter",
|
||||||
|
"RandomShadow",
|
||||||
"RandomSubsetApply",
|
"RandomSubsetApply",
|
||||||
"SharpnessJitter",
|
"SharpnessJitter",
|
||||||
"make_transform_from_config",
|
"make_transform_from_config",
|
||||||
|
|||||||
@@ -14,6 +14,7 @@
|
|||||||
# See the License for the specific language governing permissions and
|
# See the License for the specific language governing permissions and
|
||||||
# limitations under the License.
|
# limitations under the License.
|
||||||
import collections
|
import collections
|
||||||
|
import math
|
||||||
from collections.abc import Callable, Sequence
|
from collections.abc import Callable, Sequence
|
||||||
from dataclasses import dataclass, field
|
from dataclasses import dataclass, field
|
||||||
from typing import Any
|
from typing import Any
|
||||||
@@ -144,6 +145,336 @@ class SharpnessJitter(Transform):
|
|||||||
return self._call_kernel(F.adjust_sharpness, inpt, sharpness_factor=sharpness_factor)
|
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
|
@dataclass
|
||||||
class ImageTransformConfig:
|
class ImageTransformConfig:
|
||||||
"""
|
"""
|
||||||
@@ -216,16 +547,17 @@ class ImageTransformsConfig:
|
|||||||
|
|
||||||
|
|
||||||
def make_transform_from_config(cfg: ImageTransformConfig) -> Transform:
|
def make_transform_from_config(cfg: ImageTransformConfig) -> Transform:
|
||||||
if cfg.type == "SharpnessJitter":
|
if cfg.type in _CUSTOM_TRANSFORMS:
|
||||||
return SharpnessJitter(**cfg.kwargs)
|
return _CUSTOM_TRANSFORMS[cfg.type](**cfg.kwargs)
|
||||||
|
|
||||||
transform_cls = getattr(v2, cfg.type, None)
|
transform_cls = getattr(v2, cfg.type, None)
|
||||||
if isinstance(transform_cls, type) and issubclass(transform_cls, Transform):
|
if isinstance(transform_cls, type) and issubclass(transform_cls, Transform):
|
||||||
return transform_cls(**cfg.kwargs)
|
return transform_cls(**cfg.kwargs)
|
||||||
|
|
||||||
|
valid_custom = ", ".join(sorted(_CUSTOM_TRANSFORMS.keys()))
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
f"Transform '{cfg.type}' is not valid. It must be a class in "
|
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}."
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -28,9 +28,17 @@ from lerobot.scripts.lerobot_imgtransform_viz import (
|
|||||||
save_each_transform,
|
save_each_transform,
|
||||||
)
|
)
|
||||||
from lerobot.transforms import (
|
from lerobot.transforms import (
|
||||||
|
CoarseDropout,
|
||||||
|
GammaCorrection,
|
||||||
|
GaussianNoise,
|
||||||
|
GaussianPatchBrightness,
|
||||||
ImageTransformConfig,
|
ImageTransformConfig,
|
||||||
ImageTransforms,
|
ImageTransforms,
|
||||||
ImageTransformsConfig,
|
ImageTransformsConfig,
|
||||||
|
JPEGCompression,
|
||||||
|
MotionBlur,
|
||||||
|
PlanckianJitter,
|
||||||
|
RandomShadow,
|
||||||
RandomSubsetApply,
|
RandomSubsetApply,
|
||||||
SharpnessJitter,
|
SharpnessJitter,
|
||||||
make_transform_from_config,
|
make_transform_from_config,
|
||||||
@@ -455,3 +463,63 @@ def test_save_each_transform(img_tensor_factory, tmp_path):
|
|||||||
assert (transform_dir / file_name).exists(), (
|
assert (transform_dir / file_name).exists(), (
|
||||||
f"{file_name} was not found in {transform} directory."
|
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"))
|
||||||
|
|||||||
Reference in New Issue
Block a user