feat(record): 支持 --async-save 异步保存 episode
This commit is contained in:
parent
68d7898603
commit
83a1c46441
@ -1,6 +1,11 @@
|
|||||||
import yaml
|
import yaml
|
||||||
|
import copy
|
||||||
import time
|
import time
|
||||||
|
import queue
|
||||||
import argparse
|
import argparse
|
||||||
|
import logging
|
||||||
|
import shutil
|
||||||
|
import threading
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
import ufactory_lerobot # patch
|
import ufactory_lerobot # patch
|
||||||
from lerobot.scripts.lerobot_record import *
|
from lerobot.scripts.lerobot_record import *
|
||||||
@ -9,6 +14,181 @@ from ufactory_lerobot.teleoperators.base_teleop import UFBaseTeleop
|
|||||||
from ufactory_lerobot.utils.utils import instantiate_from_dict, init_keyboard_listener
|
from ufactory_lerobot.utils.utils import instantiate_from_dict, init_keyboard_listener
|
||||||
|
|
||||||
|
|
||||||
|
def _get_dataset_writer(dataset):
|
||||||
|
return getattr(dataset, "writer", None)
|
||||||
|
|
||||||
|
|
||||||
|
def _get_episode_buffer(dataset):
|
||||||
|
try:
|
||||||
|
return dataset.episode_buffer
|
||||||
|
except AttributeError:
|
||||||
|
pass
|
||||||
|
|
||||||
|
writer = _get_dataset_writer(dataset)
|
||||||
|
if writer is not None and hasattr(writer, "episode_buffer"):
|
||||||
|
return writer.episode_buffer
|
||||||
|
raise RuntimeError("Unable to access dataset episode buffer for async save.")
|
||||||
|
|
||||||
|
|
||||||
|
def _set_episode_buffer(dataset, episode_buffer):
|
||||||
|
updated = False
|
||||||
|
writer = _get_dataset_writer(dataset)
|
||||||
|
if writer is not None and hasattr(writer, "episode_buffer"):
|
||||||
|
writer.episode_buffer = episode_buffer
|
||||||
|
updated = True
|
||||||
|
|
||||||
|
try:
|
||||||
|
getattr(dataset, "episode_buffer")
|
||||||
|
except AttributeError:
|
||||||
|
pass
|
||||||
|
else:
|
||||||
|
try:
|
||||||
|
dataset.episode_buffer = episode_buffer
|
||||||
|
updated = True
|
||||||
|
except AttributeError:
|
||||||
|
pass
|
||||||
|
|
||||||
|
if not updated:
|
||||||
|
raise RuntimeError("Unable to replace dataset episode buffer for async save.")
|
||||||
|
|
||||||
|
|
||||||
|
def _to_int(value):
|
||||||
|
if isinstance(value, (list, tuple)):
|
||||||
|
return int(value[0])
|
||||||
|
if hasattr(value, "item"):
|
||||||
|
try:
|
||||||
|
return int(value.item())
|
||||||
|
except (TypeError, ValueError):
|
||||||
|
pass
|
||||||
|
try:
|
||||||
|
return int(value)
|
||||||
|
except (TypeError, ValueError):
|
||||||
|
return int(value[0])
|
||||||
|
|
||||||
|
|
||||||
|
def _episode_buffer_size(episode_buffer):
|
||||||
|
return _to_int(episode_buffer.get("size", 0))
|
||||||
|
|
||||||
|
|
||||||
|
def _episode_buffer_index(episode_buffer):
|
||||||
|
return _to_int(episode_buffer["episode_index"])
|
||||||
|
|
||||||
|
|
||||||
|
def _current_episode_index(dataset):
|
||||||
|
try:
|
||||||
|
return _episode_buffer_index(_get_episode_buffer(dataset))
|
||||||
|
except Exception:
|
||||||
|
return dataset.num_episodes
|
||||||
|
|
||||||
|
|
||||||
|
def _create_empty_episode_buffer(dataset, episode_index, template_episode_buffer):
|
||||||
|
writer = _get_dataset_writer(dataset)
|
||||||
|
|
||||||
|
if writer is not None and hasattr(writer, "_create_episode_buffer"):
|
||||||
|
episode_buffer = writer._create_episode_buffer()
|
||||||
|
elif hasattr(dataset, "create_episode_buffer"):
|
||||||
|
episode_buffer = dataset.create_episode_buffer(episode_index=episode_index)
|
||||||
|
elif hasattr(dataset, "_create_episode_buffer"):
|
||||||
|
episode_buffer = dataset._create_episode_buffer()
|
||||||
|
else:
|
||||||
|
episode_buffer = copy.deepcopy(template_episode_buffer)
|
||||||
|
for key, value in list(episode_buffer.items()):
|
||||||
|
if key == "size":
|
||||||
|
episode_buffer[key] = 0
|
||||||
|
elif key == "episode_index":
|
||||||
|
continue
|
||||||
|
elif isinstance(value, list):
|
||||||
|
episode_buffer[key] = []
|
||||||
|
else:
|
||||||
|
episode_buffer[key] = []
|
||||||
|
|
||||||
|
if _episode_buffer_index(episode_buffer) != episode_index:
|
||||||
|
episode_buffer["episode_index"] = episode_index
|
||||||
|
return episode_buffer
|
||||||
|
|
||||||
|
|
||||||
|
def _create_next_episode_buffer(dataset, current_episode_buffer):
|
||||||
|
current_episode_index = _episode_buffer_index(current_episode_buffer)
|
||||||
|
return _create_empty_episode_buffer(dataset, current_episode_index + 1, current_episode_buffer)
|
||||||
|
|
||||||
|
|
||||||
|
class AsyncEpisodeSaver:
|
||||||
|
_STOP = object()
|
||||||
|
|
||||||
|
def __init__(self, dataset):
|
||||||
|
self.dataset = dataset
|
||||||
|
self._queue = queue.Queue()
|
||||||
|
self._total_cnts = 0
|
||||||
|
self._finish_cnts = 0
|
||||||
|
self._exception = None
|
||||||
|
self._thread = threading.Thread(target=self._run, name="uf-async-episode-saver", daemon=True)
|
||||||
|
self._thread.start()
|
||||||
|
|
||||||
|
def submit_current_episode(self):
|
||||||
|
self._raise_if_failed()
|
||||||
|
episode_buffer = _get_episode_buffer(self.dataset)
|
||||||
|
if _episode_buffer_size(episode_buffer) == 0:
|
||||||
|
raise RuntimeError("Cannot async save an empty episode buffer.")
|
||||||
|
|
||||||
|
episode_index = _episode_buffer_index(episode_buffer)
|
||||||
|
next_episode_buffer = _create_next_episode_buffer(self.dataset, episode_buffer)
|
||||||
|
_set_episode_buffer(self.dataset, next_episode_buffer)
|
||||||
|
self._queue.put((episode_index, episode_buffer))
|
||||||
|
return episode_index
|
||||||
|
|
||||||
|
def wait_idle(self):
|
||||||
|
self._queue.join()
|
||||||
|
self._raise_if_failed()
|
||||||
|
|
||||||
|
def close(self):
|
||||||
|
self._queue.join()
|
||||||
|
self._queue.put(self._STOP)
|
||||||
|
self._queue.join()
|
||||||
|
self._thread.join()
|
||||||
|
self._raise_if_failed()
|
||||||
|
|
||||||
|
def _run(self):
|
||||||
|
while True:
|
||||||
|
item = self._queue.get()
|
||||||
|
try:
|
||||||
|
if item is self._STOP:
|
||||||
|
return
|
||||||
|
episode_index, episode_buffer = item
|
||||||
|
print(f'[Async] saving episode {episode_index}')
|
||||||
|
try:
|
||||||
|
self.dataset.save_episode(episode_data=episode_buffer)
|
||||||
|
except TypeError as exc:
|
||||||
|
if "episode_data" in str(exc):
|
||||||
|
raise RuntimeError(
|
||||||
|
"--async-save requires LeRobotDataset.save_episode(episode_data=...)."
|
||||||
|
) from exc
|
||||||
|
raise
|
||||||
|
self._delete_saved_image_dirs(episode_index)
|
||||||
|
print(f'[Async] save episode {episode_index} finish')
|
||||||
|
except BaseException as exc:
|
||||||
|
self._exception = exc
|
||||||
|
print(f'[Async] episode {episode_index} save failed, {exc}')
|
||||||
|
finally:
|
||||||
|
self._queue.task_done()
|
||||||
|
|
||||||
|
def _delete_saved_image_dirs(self, episode_index):
|
||||||
|
writer = _get_dataset_writer(self.dataset)
|
||||||
|
meta = getattr(writer, "_meta", getattr(self.dataset, "meta", None))
|
||||||
|
image_keys = getattr(meta, "image_keys", [])
|
||||||
|
image_dir_owner = writer if writer is not None and hasattr(writer, "_get_image_file_dir") else self.dataset
|
||||||
|
if not image_keys or not hasattr(image_dir_owner, "_get_image_file_dir"):
|
||||||
|
return
|
||||||
|
|
||||||
|
for cam_key in image_keys:
|
||||||
|
img_dir = image_dir_owner._get_image_file_dir(episode_index, cam_key)
|
||||||
|
if img_dir.is_dir():
|
||||||
|
shutil.rmtree(img_dir)
|
||||||
|
|
||||||
|
def _raise_if_failed(self):
|
||||||
|
if self._exception is not None:
|
||||||
|
raise RuntimeError("Async episode save failed.") from self._exception
|
||||||
|
|
||||||
|
|
||||||
@safe_stop_image_writer
|
@safe_stop_image_writer
|
||||||
def record_loop(
|
def record_loop(
|
||||||
robot: Robot,
|
robot: Robot,
|
||||||
@ -165,7 +345,7 @@ def record_loop(
|
|||||||
|
|
||||||
|
|
||||||
@parser.wrap()
|
@parser.wrap()
|
||||||
def record(cfg: RecordConfig) -> LeRobotDataset:
|
def record(cfg: RecordConfig, async_save: bool = False) -> LeRobotDataset:
|
||||||
init_logging()
|
init_logging()
|
||||||
logging.info(pformat(asdict(cfg)))
|
logging.info(pformat(asdict(cfg)))
|
||||||
if cfg.display_data:
|
if cfg.display_data:
|
||||||
@ -293,9 +473,12 @@ def record(cfg: RecordConfig) -> LeRobotDataset:
|
|||||||
if is_uf_teleop:
|
if is_uf_teleop:
|
||||||
teleop.set_teleop_enabled(True)
|
teleop.set_teleop_enabled(True)
|
||||||
is_recorded = True
|
is_recorded = True
|
||||||
print("\n********** Episode Record Loop Start **********")
|
print('\n********** Episode Record Loop Start **********')
|
||||||
|
|
||||||
frame_callback = None
|
frame_callback = None
|
||||||
|
async_episode_saver = AsyncEpisodeSaver(dataset) if async_save else None
|
||||||
|
if async_episode_saver is not None:
|
||||||
|
print('Async episode saving is enabled.')
|
||||||
|
|
||||||
with VideoEncodingManager(dataset):
|
with VideoEncodingManager(dataset):
|
||||||
recorded_episodes = 0
|
recorded_episodes = 0
|
||||||
@ -325,7 +508,7 @@ def record(cfg: RecordConfig) -> LeRobotDataset:
|
|||||||
robot.configure()
|
robot.configure()
|
||||||
obs = robot.get_observation()
|
obs = robot.get_observation()
|
||||||
teleop.set_teleop_enabled(True, obs)
|
teleop.set_teleop_enabled(True, obs)
|
||||||
log_say(f"Recording episode {dataset.num_episodes}", cfg.play_sounds)
|
log_say(f"Recording episode {_current_episode_index(dataset)}", cfg.play_sounds)
|
||||||
record_loop(
|
record_loop(
|
||||||
robot=robot,
|
robot=robot,
|
||||||
events=events,
|
events=events,
|
||||||
@ -353,8 +536,14 @@ def record(cfg: RecordConfig) -> LeRobotDataset:
|
|||||||
events["exit_early"] = False
|
events["exit_early"] = False
|
||||||
if is_uf_teleop:
|
if is_uf_teleop:
|
||||||
teleop.set_teleop_enabled(False)
|
teleop.set_teleop_enabled(False)
|
||||||
if dataset.episode_buffer:
|
episode_buffer = _get_episode_buffer(dataset)
|
||||||
|
if _episode_buffer_size(episode_buffer) > 0:
|
||||||
|
if async_episode_saver is None:
|
||||||
dataset.clear_episode_buffer()
|
dataset.clear_episode_buffer()
|
||||||
|
else:
|
||||||
|
episode_index = _episode_buffer_index(episode_buffer)
|
||||||
|
empty_episode_buffer = _create_empty_episode_buffer(dataset, episode_index, episode_buffer)
|
||||||
|
_set_episode_buffer(dataset, empty_episode_buffer)
|
||||||
is_recorded = False
|
is_recorded = False
|
||||||
if is_evt:
|
if is_evt:
|
||||||
print('⌨ [ESC] Exit [Space] Start [←] Reset [→] Save')
|
print('⌨ [ESC] Exit [Space] Start [←] Reset [→] Save')
|
||||||
@ -364,11 +553,18 @@ def record(cfg: RecordConfig) -> LeRobotDataset:
|
|||||||
continue
|
continue
|
||||||
|
|
||||||
if is_recorded and not events['stop_recording']:
|
if is_recorded and not events['stop_recording']:
|
||||||
log_say(f"Save episode {dataset.num_episodes}", cfg.play_sounds)
|
episode_index = _current_episode_index(dataset)
|
||||||
|
log_say(f"Save episode {episode_index}", cfg.play_sounds)
|
||||||
if is_uf_teleop:
|
if is_uf_teleop:
|
||||||
teleop.set_teleop_enabled(False)
|
teleop.set_teleop_enabled(False)
|
||||||
|
if async_episode_saver is None:
|
||||||
dataset.save_episode()
|
dataset.save_episode()
|
||||||
log_say(f"[Finish] Save episode {dataset.num_episodes}", cfg.play_sounds)
|
log_say(f"[Finish] Save episode {episode_index}", cfg.play_sounds)
|
||||||
|
else:
|
||||||
|
queued_episode_index = async_episode_saver.submit_current_episode()
|
||||||
|
if queued_episode_index is not None:
|
||||||
|
log_say(f"[Queued] Save episode {queued_episode_index}", cfg.play_sounds)
|
||||||
|
|
||||||
recorded_episodes += 1
|
recorded_episodes += 1
|
||||||
is_recorded = False
|
is_recorded = False
|
||||||
if is_evt:
|
if is_evt:
|
||||||
@ -377,6 +573,10 @@ def record(cfg: RecordConfig) -> LeRobotDataset:
|
|||||||
input('⌨ Press Enter to record at the next episode >>>>> ')
|
input('⌨ Press Enter to record at the next episode >>>>> ')
|
||||||
is_recorded = True
|
is_recorded = True
|
||||||
|
|
||||||
|
if async_episode_saver is not None:
|
||||||
|
print('Waiting for pending async episode saves.')
|
||||||
|
async_episode_saver.close()
|
||||||
|
|
||||||
print("\n********** Episode Record Loop Exit **********")
|
print("\n********** Episode Record Loop Exit **********")
|
||||||
|
|
||||||
robot.disconnect()
|
robot.disconnect()
|
||||||
@ -401,6 +601,10 @@ def main():
|
|||||||
action='store_true', # specify --resume if resume needs to be True
|
action='store_true', # specify --resume if resume needs to be True
|
||||||
default=False,
|
default=False,
|
||||||
help='Whether contitue recording on existing dataset (default: False)')
|
help='Whether contitue recording on existing dataset (default: False)')
|
||||||
|
parser.add_argument('-a', '--async-save',
|
||||||
|
action='store_true',
|
||||||
|
default=False,
|
||||||
|
help='Enable async background saving (default: False)')
|
||||||
args = parser.parse_args()
|
args = parser.parse_args()
|
||||||
try:
|
try:
|
||||||
with open(Path(args.config).expanduser(), 'r') as f:
|
with open(Path(args.config).expanduser(), 'r') as f:
|
||||||
@ -412,7 +616,7 @@ def main():
|
|||||||
config = instantiate_from_dict(cfg)
|
config = instantiate_from_dict(cfg)
|
||||||
|
|
||||||
record_cfg = RecordConfig(resume=args.resume, play_sounds=False, robot=config["RobotConfig"], dataset=config["DatasetRecordConfig"], teleop=config["TeleoperatorConfig"])
|
record_cfg = RecordConfig(resume=args.resume, play_sounds=False, robot=config["RobotConfig"], dataset=config["DatasetRecordConfig"], teleop=config["TeleoperatorConfig"])
|
||||||
record(record_cfg)
|
record(record_cfg, async_save=args.async_save)
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
|
|||||||
Loading…
Reference in New Issue
Block a user