Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 1 addition & 1 deletion README.md
Original file line number Diff line number Diff line change
Expand Up @@ -223,7 +223,7 @@ openarm-dataset-convert <input> <output> \
[--camera-format {dir,tar}] # default dir (openarm only); tar packs each \
# camera into one .tar archive \
[--fps INT] # default 30 (lerobot/gr00t only) \
[--smoothing-cutoff FLOAT] # default 1.0 (lerobot/gr00t only) \
[--smoothing-cutoff FLOAT] # default 1.0, 0 disables (lerobot/gr00t only) \
Comment thread
kou marked this conversation as resolved.
[--train-split FLOAT] # default 0.8 (lerobot/gr00t only) \
[--success-only] # lerobot/gr00t only \
[--valid-only] # exclude episodes marked invalid \
Expand Down
2 changes: 1 addition & 1 deletion src/openarm_dataset/convert.py
Original file line number Diff line number Diff line change
Expand Up @@ -53,7 +53,7 @@ def main():
)
parser.add_argument(
"--smoothing-cutoff",
help="Cutoff frequency for smoothing (default: 1.0) if the output format is lerobot_v2.1, lerobot_v3.0 or gr00t",
help="Cutoff frequency for smoothing in Hz (default: 1.0; 0 disables smoothing) if the output format is lerobot_v2.1, lerobot_v3.0 or gr00t",
type=float,
default=1.0,
)
Expand Down
33 changes: 26 additions & 7 deletions src/openarm_dataset/dataset.py
Original file line number Diff line number Diff line change
Expand Up @@ -114,6 +114,12 @@ def _renormalize_orientation(df: pd.DataFrame, attribute: str) -> pd.DataFrame:
return df


def _validate_cutoff(cutoff: float | None) -> None:
"""Raise ValueError if the given smoothing cutoff is negative."""
if cutoff is not None and cutoff < 0:
raise ValueError(f"cutoff must not be negative, got {cutoff}")


class Dataset:
"""OpenArm Dataset."""

Expand Down Expand Up @@ -144,8 +150,15 @@ def __init__(
self._smoothing_cutoff = None
self._kinematics = kinematics

def set_smoothing(self, cutoff: float):
"""Set smoothing."""
def set_smoothing(self, cutoff: float | None):
"""Set smoothing.

Args:
cutoff: Cutoff frequency for smoothing in Hz. If None or 0,
smoothing is disabled.

"""
_validate_cutoff(cutoff)
self._smoothing_cutoff = cutoff

def validate(
Expand Down Expand Up @@ -242,7 +255,9 @@ def load_obs(
episode: Episode to load.
use_unixtime: If True, the DataFrame index is returned as Unix time
(float64) instead of datetime64[ns].
cutoff: If not None, smoothing is applied using this value.
cutoff: Cutoff frequency for smoothing in Hz. If None, the
value set by set_smoothing() is used. If 0, smoothing is
disabled.
state: If not None, return arm states in this representation
("qpos", "pose" or "rot6d"), converting the recorded data
on the fly (qpos to pose via FK, pose to qpos via IK) when
Expand All @@ -261,11 +276,12 @@ def load_obs(
}

"""
_validate_cutoff(cutoff)
return self._load_embodiment_values(
"obs",
episode,
use_unixtime,
cutoff=cutoff or self._smoothing_cutoff,
cutoff=self._smoothing_cutoff if cutoff is None else cutoff,
state=state,
)

Expand All @@ -282,7 +298,9 @@ def load_action(
episode: Episode to load.
use_unixtime: If True, the DataFrame index is returned as Unix time
(float64) instead of datetime64[ns].
cutoff: If not None, smoothing is applied using this value.
cutoff: Cutoff frequency for smoothing in Hz. If None, the
value set by set_smoothing() is used. If 0, smoothing is
disabled.
state: If not None, return arm states in this representation
("qpos", "pose" or "rot6d"), converting the recorded data
on the fly (qpos to pose via FK, pose to qpos via IK) when
Expand All @@ -301,11 +319,12 @@ def load_action(
}

"""
_validate_cutoff(cutoff)
return self._load_embodiment_values(
"action",
episode,
use_unixtime=use_unixtime,
cutoff=cutoff or self._smoothing_cutoff,
cutoff=self._smoothing_cutoff if cutoff is None else cutoff,
state=state,
)

Expand Down Expand Up @@ -510,7 +529,7 @@ def _load_embodiment_values(
component,
use_unixtime=use_unixtime,
)
if cutoff is not None:
if cutoff is not None and cutoff > 0:
Comment thread
kou marked this conversation as resolved.
values = {
key: _renormalize_orientation(
self._apply_smoothing(df, cutoff=cutoff),
Expand Down
38 changes: 38 additions & 0 deletions tests/test_dataset_0_4_0_qpos.py
Original file line number Diff line number Diff line change
Expand Up @@ -70,6 +70,44 @@ def test_load_action(dataset):
assert action["lifter/elevation"].shape == (90, 1)


def test_load_obs_smoothing(dataset):
episode = dataset.meta.episodes[0]
raw = dataset.load_obs(episode)
dataset.set_smoothing(1.0)
smoothed = dataset.load_obs(episode)
assert not smoothed["arms/left/qvel"].equals(raw["arms/left/qvel"])
# 0 overrides the value set by set_smoothing().
for key, df in dataset.load_obs(episode, cutoff=0).items():
pd.testing.assert_frame_equal(df, raw[key])
dataset.set_smoothing(0)
for key, df in dataset.load_obs(episode).items():
pd.testing.assert_frame_equal(df, raw[key])


def test_load_action_smoothing(dataset):
episode = dataset.meta.episodes[0]
raw = dataset.load_action(episode)
dataset.set_smoothing(1.0)
smoothed = dataset.load_action(episode)
assert not smoothed["arms/left/qpos"].equals(raw["arms/left/qpos"])
# 0 overrides the value set by set_smoothing().
for key, df in dataset.load_action(episode, cutoff=0).items():
pd.testing.assert_frame_equal(df, raw[key])
dataset.set_smoothing(0)
for key, df in dataset.load_action(episode).items():
pd.testing.assert_frame_equal(df, raw[key])


def test_negative_cutoff(dataset):
episode = dataset.meta.episodes[0]
with pytest.raises(ValueError, match="cutoff must not be negative"):
dataset.set_smoothing(-1.0)
with pytest.raises(ValueError, match="cutoff must not be negative"):
dataset.load_obs(episode, cutoff=-1.0)
with pytest.raises(ValueError, match="cutoff must not be negative"):
dataset.load_action(episode, cutoff=-1.0)


def test_cameras(dataset):
assert set(dataset.camera_names) == {
"ceiling",
Expand Down
Loading