Files
roboimi/roboimi/demos/diana_record_sim_episodes.py
T

106 lines
3.5 KiB
Python

import time
import os
import numpy as np
from roboimi.envs.double_pos_ctrl_env import make_sim_env
from roboimi.demos.diana_air_insert_policy import TestAirInsertPolicy
from roboimi.demos.diana_policy import TestPickAndTransferPolicy
import cv2
from roboimi.utils.act_ex_utils import sample_air_insert_socket_peg_state, sample_transfer_pose
from roboimi.utils.constants import SIM_TASK_CONFIGS
from roboimi.utils.streaming_episode_writer import StreamingEpisodeWriter
import pathlib
HOME_PATH = str(pathlib.Path(__file__).parent.resolve())
DATASET_DIR = HOME_PATH + '/dataset'
def sample_task_state(task_name):
if task_name == 'sim_transfer':
return sample_transfer_pose()
if task_name == 'sim_air_insert_socket_peg':
return sample_air_insert_socket_peg_state()
raise NotImplementedError(f'Unsupported scripted rollout task: {task_name}')
def make_policy(task_name, inject_noise=False, grasp_strategy=None):
if task_name == 'sim_transfer':
return TestPickAndTransferPolicy(inject_noise)
if task_name == 'sim_air_insert_socket_peg':
if grasp_strategy is None:
return TestAirInsertPolicy(inject_noise)
return TestAirInsertPolicy(inject_noise, grasp_strategy=grasp_strategy)
raise NotImplementedError(f'Unsupported scripted rollout task: {task_name}')
def main(task_name='sim_transfer'):
task_cfg = SIM_TASK_CONFIGS[task_name]
dataset_dir = task_cfg['dataset_dir']
num_episodes = 100
inject_noise = False
episode_len = task_cfg['episode_len']
camera_names = task_cfg['camera_names']
image_size = (256, 256)
if task_name in {'sim_transfer', 'sim_air_insert_socket_peg'}:
print(task_name)
else:
raise NotImplementedError
success = []
env = make_sim_env(task_name)
policy = make_policy(task_name, inject_noise=inject_noise)
# 等待osmesa完全启动后再开始收集数据
print("等待osmesa线程启动...")
time.sleep(60)
print("osmesa已就绪,开始收集数据...")
for episode_idx in range(num_episodes):
sum_reward = 0.0
max_reward = float('-inf')
print(f'\n{episode_idx=}')
print('Rollout out EE space scripted policy')
task_state = sample_task_state(task_name)
env.reset(task_state)
episode_writer = StreamingEpisodeWriter(
dataset_path=os.path.join(dataset_dir, f'episode_{episode_idx}.hdf5'),
max_timesteps=episode_len,
camera_names=camera_names,
image_size=image_size,
)
for step in range(episode_len):
raw_action = policy.predict(task_state, step)
env.step(raw_action)
env.render()
sum_reward += env.rew
max_reward = max(max_reward, env.rew)
episode_writer.append(
qpos=env.obs['qpos'],
action=raw_action,
images=env.obs['images'],
)
if max_reward == env.max_reward:
success.append(1)
print(f"{episode_idx=} Successful, {sum_reward=}")
episode_writer.commit()
else:
success.append(0)
print(f"{episode_idx=} Failed")
print(max_reward)
episode_writer.discard()
# del policy
# env.viewer.close()
# del env
print(f'Success: {np.sum(success)} / {len(success)}')
env.exit_flag = True
cv2.destroyAllWindows()
cv2.waitKey(1)
env.cam_thread.join()
env.viewer.close()
if __name__ == '__main__':
main()