106 lines
3.5 KiB
Python
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()
|