mirror of
https://github.com/huggingface/lerobot.git
synced 2026-07-28 20:26:05 +00:00
chore(mypy): cover annotations and transforms (#3860)
Co-authored-by: Steven Palma <imstevenpmwork@ieee.org>
This commit is contained in:
@@ -494,6 +494,19 @@ ignore_errors = true
|
|||||||
module = "lerobot.envs.*"
|
module = "lerobot.envs.*"
|
||||||
ignore_errors = false
|
ignore_errors = false
|
||||||
|
|
||||||
|
[[tool.mypy.overrides]]
|
||||||
|
module = "lerobot.annotations.*"
|
||||||
|
ignore_errors = false
|
||||||
|
disallow_untyped_defs = true
|
||||||
|
disallow_incomplete_defs = true
|
||||||
|
check_untyped_defs = true
|
||||||
|
|
||||||
|
[[tool.mypy.overrides]]
|
||||||
|
module = "lerobot.transforms.*"
|
||||||
|
ignore_errors = false
|
||||||
|
disallow_untyped_defs = true
|
||||||
|
disallow_incomplete_defs = true
|
||||||
|
check_untyped_defs = true
|
||||||
|
|
||||||
# [[tool.mypy.overrides]]
|
# [[tool.mypy.overrides]]
|
||||||
# module = "lerobot.utils.*"
|
# module = "lerobot.utils.*"
|
||||||
|
|||||||
@@ -384,7 +384,9 @@ class RoboTwinEnv(gym.Env):
|
|||||||
|
|
||||||
self._env: Any | None = None # deferred — created on first reset() inside worker
|
self._env: Any | None = None # deferred — created on first reset() inside worker
|
||||||
self._step_count: int = 0
|
self._step_count: int = 0
|
||||||
self._black_frame = np.zeros((self.observation_height, self.observation_width, 3), dtype=np.uint8)
|
self._black_frame: np.ndarray = np.zeros(
|
||||||
|
(self.observation_height, self.observation_width, 3), dtype=np.uint8
|
||||||
|
)
|
||||||
|
|
||||||
image_spaces = {
|
image_spaces = {
|
||||||
cam: spaces.Box(
|
cam: spaces.Box(
|
||||||
|
|||||||
@@ -373,7 +373,7 @@ class VLABenchEnv(gym.Env):
|
|||||||
|
|
||||||
if action.shape[0] != 7:
|
if action.shape[0] != 7:
|
||||||
# Unknown layout — fall back to zero-pad so the sim doesn't crash.
|
# Unknown layout — fall back to zero-pad so the sim doesn't crash.
|
||||||
padded = np.zeros(ctrl_dim, dtype=np.float64)
|
padded: np.ndarray = np.zeros(ctrl_dim, dtype=np.float64)
|
||||||
padded[: min(action.shape[0], ctrl_dim)] = action[:ctrl_dim]
|
padded[: min(action.shape[0], ctrl_dim)] = action[:ctrl_dim]
|
||||||
return padded
|
return padded
|
||||||
|
|
||||||
|
|||||||
@@ -41,7 +41,7 @@ class RandomSubsetApply(Transform):
|
|||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
transforms: Sequence[Callable],
|
transforms: Sequence[Callable[..., Any]],
|
||||||
p: list[float] | None = None,
|
p: list[float] | None = None,
|
||||||
n_subset: int | None = None,
|
n_subset: int | None = None,
|
||||||
random_order: bool = False,
|
random_order: bool = False,
|
||||||
@@ -50,7 +50,7 @@ class RandomSubsetApply(Transform):
|
|||||||
if not isinstance(transforms, Sequence):
|
if not isinstance(transforms, Sequence):
|
||||||
raise TypeError("Argument transforms should be a sequence of callables")
|
raise TypeError("Argument transforms should be a sequence of callables")
|
||||||
if p is None:
|
if p is None:
|
||||||
p = [1] * len(transforms)
|
p = [1.0] * len(transforms)
|
||||||
elif len(p) != len(transforms):
|
elif len(p) != len(transforms):
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
f"Length of p doesn't match the number of transforms: {len(p)} != {len(transforms)}"
|
f"Length of p doesn't match the number of transforms: {len(p)} != {len(transforms)}"
|
||||||
@@ -69,7 +69,7 @@ class RandomSubsetApply(Transform):
|
|||||||
self.n_subset = n_subset
|
self.n_subset = n_subset
|
||||||
self.random_order = random_order
|
self.random_order = random_order
|
||||||
|
|
||||||
self.selected_transforms = None
|
self.selected_transforms: list[Callable[..., Any]] = []
|
||||||
|
|
||||||
def forward(self, *inputs: Any) -> Any:
|
def forward(self, *inputs: Any) -> Any:
|
||||||
needs_unpacking = len(inputs) > 1
|
needs_unpacking = len(inputs) > 1
|
||||||
@@ -119,7 +119,7 @@ class SharpnessJitter(Transform):
|
|||||||
super().__init__()
|
super().__init__()
|
||||||
self.sharpness = self._check_input(sharpness)
|
self.sharpness = self._check_input(sharpness)
|
||||||
|
|
||||||
def _check_input(self, sharpness):
|
def _check_input(self, sharpness: float | Sequence[float]) -> tuple[float, float]:
|
||||||
if isinstance(sharpness, (int | float)):
|
if isinstance(sharpness, (int | float)):
|
||||||
if sharpness < 0:
|
if sharpness < 0:
|
||||||
raise ValueError("If sharpness is a single number, it must be non negative.")
|
raise ValueError("If sharpness is a single number, it must be non negative.")
|
||||||
@@ -215,7 +215,7 @@ class ImageTransformsConfig:
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
def make_transform_from_config(cfg: ImageTransformConfig):
|
def make_transform_from_config(cfg: ImageTransformConfig) -> Transform:
|
||||||
if cfg.type == "SharpnessJitter":
|
if cfg.type == "SharpnessJitter":
|
||||||
return SharpnessJitter(**cfg.kwargs)
|
return SharpnessJitter(**cfg.kwargs)
|
||||||
|
|
||||||
@@ -236,8 +236,8 @@ class ImageTransforms(Transform):
|
|||||||
super().__init__()
|
super().__init__()
|
||||||
self._cfg = cfg
|
self._cfg = cfg
|
||||||
|
|
||||||
self.weights = []
|
self.weights: list[float] = []
|
||||||
self.transforms = {}
|
self.transforms: dict[str, Transform] = {}
|
||||||
for tf_name, tf_cfg in cfg.tfs.items():
|
for tf_name, tf_cfg in cfg.tfs.items():
|
||||||
if tf_cfg.weight <= 0.0:
|
if tf_cfg.weight <= 0.0:
|
||||||
continue
|
continue
|
||||||
|
|||||||
Reference in New Issue
Block a user