File size: 24,912 Bytes
987ed1b
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
416
417
418
419
420
421
422
423
424
425
426
427
428
429
430
431
432
433
434
435
436
437
438
439
440
441
442
443
444
445
446
447
448
449
450
451
452
453
454
455
456
457
458
459
460
461
462
463
464
465
466
467
468
469
470
471
472
473
474
475
476
477
478
479
480
481
482
483
484
485
486
487
488
489
490
491
492
493
494
495
496
497
498
499
500
501
502
503
504
505
506
507
508
509
510
511
512
513
514
515
516
517
518
519
520
521
522
523
524
525
526
527
528
529
530
531
532
533
534
535
536
537
538
539
540
541
542
543
544
545
546
547
548
549
550
551
552
553
554
555
556
557
558
559
560
561
562
563
564
565
566
567
568
569
570
571
572
573
import os
import wandb
import numpy as np
import torch
import collections
import pathlib
import tqdm
import h5py
import dill
import math
import wandb.sdk.data_types.video as wv
from diffusion_policy.gym_util.async_vector_env import AsyncVectorEnv
# from diffusion_policy.gym_util.sync_vector_env import SyncVectorEnv
from diffusion_policy.gym_util.multistep_wrapper import MultiStepWrapper
from diffusion_policy.gym_util.video_recording_wrapper import VideoRecordingWrapper, VideoRecorder
from diffusion_policy.model.common.rotation_transformer import RotationTransformer

from diffusion_policy.policy.base_lowdim_policy import BaseLowdimPolicy
from diffusion_policy.common.pytorch_util import dict_apply
from diffusion_policy.env_runner.base_lowdim_runner import BaseLowdimRunner
from diffusion_policy.env.robomimic.robomimic_lowdim_wrapper import RobomimicLowdimWrapper
import robomimic.utils.file_utils as FileUtils
import robomimic.utils.env_utils as EnvUtils
import robomimic.utils.obs_utils as ObsUtils

from termcolor import colored
from diffusion_policy.sampler.single import coherence_sampler, ema_sampler, ac_sampler, sgac_sampler
from diffusion_policy.sampler.multi import contrastive_sampler, bidirectional_sampler
from diffusion_policy.sampler.condition import NoiseGenerator

def create_env(env_meta, obs_keys):
    ObsUtils.initialize_obs_modality_mapping_from_dict(
        {'low_dim': obs_keys})
    env = EnvUtils.create_env_from_metadata(
        env_meta=env_meta,
        render=False, 
        # only way to not show collision geometry
        # is to enable render_offscreen
        # which uses a lot of RAM.
        render_offscreen=False,
        use_image_obs=False,
    )
    return env


class RobomimicLowdimRunner(BaseLowdimRunner):
    """
    Robomimic envs already enforces number of steps.
    """

    def __init__(
            self,
            output_dir,
            dataset_path,
            obs_keys,
            n_train=10,
            n_train_vis=3,
            train_start_idx=0,
            n_test=22,
            n_test_vis=6,
            test_start_seed=10000,
            max_steps=400,
            n_obs_steps=2,
            n_action_steps=8,
            n_latency_steps=0,
            render_hw=(256,256),
            render_camera_name='agentview',
            fps=10,
            crf=22,
            past_action=False,
            abs_action=False,
            tqdm_interval_sec=5.0,
            n_envs=None,
            perturb_level=0.0,
            return_intermediate_state=False,
            use_oracle_ac=False,
            oracle_ac_config=None,
            collect_data=False,
        ):
        """
        Assuming:
        n_obs_steps=2
        n_latency_steps=3
        n_action_steps=4
        o: obs
        i: inference
        a: action
        Batch t:
        |o|o| | | | | | |
        | |i|i|i| | | | |
        | | | | |a|a|a|a|
        Batch t+1
        | | | | |o|o| | | | | | |
        | | | | | |i|i|i| | | | |
        | | | | | | | | |a|a|a|a|
        """

        super().__init__(output_dir)
        self.return_intermediate_state = return_intermediate_state
        self.use_oracle_ac = use_oracle_ac
        self.oracle_ac_config = oracle_ac_config
        self.collect_data = collect_data

        # # reset render size
        # factor = 3
        # render_hw[0] *= factor
        # render_hw[1] *= factor

        if n_envs is None:
            n_envs = n_train + n_test

        # handle latency step
        # to mimic latency, we request n_latency_steps additional steps 
        # of past observations, and the discard the last n_latency_steps
        env_n_obs_steps = n_obs_steps + n_latency_steps
        self.env_n_action_steps = n_action_steps
        _env_n_action_steps = 1 if self.return_intermediate_state else self.env_n_action_steps

        # assert n_obs_steps <= n_action_steps
        dataset_path = os.path.expanduser(dataset_path)
        robosuite_fps = 20
        steps_per_render = max(robosuite_fps // fps, 1)

        # read from dataset
        env_meta = FileUtils.get_env_metadata_from_dataset(
            dataset_path)
        rotation_transformer = None
        if abs_action:
            env_meta['env_kwargs']['controller_configs']['control_delta'] = False
            rotation_transformer = RotationTransformer('axis_angle', 'rotation_6d')
        if self.collect_data:
            env_meta['env_kwargs']['reward_shaping'] = True

        def env_fn():
            robomimic_env = create_env(
                    env_meta=env_meta, 
                    obs_keys=obs_keys
                )
            # hard reset doesn't influence lowdim env
            # robomimic_env.env.hard_reset = False
            return MultiStepWrapper(
                    VideoRecordingWrapper(
                        RobomimicLowdimWrapper(
                            env=robomimic_env,
                            obs_keys=obs_keys,
                            init_state=None,
                            render_hw=render_hw,
                            render_camera_name=render_camera_name
                        ),
                        video_recoder=VideoRecorder.create_h264(
                            fps=fps,
                            codec='h264',
                            input_pix_fmt='rgb24',
                            crf=crf,
                            thread_type='FRAME',
                            thread_count=1
                        ),
                        file_path=None,
                        steps_per_render=steps_per_render if not self.collect_data else 1
                    ),
                    n_obs_steps=env_n_obs_steps,
                    n_action_steps=_env_n_action_steps,
                    max_episode_steps=max_steps
                )

        env_fns = [env_fn] * n_envs
        env_seeds = list()
        env_prefixs = list()
        env_init_fn_dills = list()

        # train
        with h5py.File(dataset_path, 'r') as f:
            for i in range(n_train):
                train_idx = train_start_idx + i
                enable_render = i < n_train_vis
                init_state = f[f'data/demo_{train_idx}/states'][0]

                def init_fn(env, init_state=init_state, 
                    enable_render=enable_render):
                    # setup rendering
                    # video_wrapper
                    assert isinstance(env.env, VideoRecordingWrapper)
                    env.env.video_recoder.stop()
                    env.env.file_path = None
                    if enable_render:
                        filename = pathlib.Path(output_dir).joinpath(
                            'media', wv.util.generate_id() + ".mp4")
                        filename.parent.mkdir(parents=False, exist_ok=True)
                        filename = str(filename)
                        env.env.file_path = filename

                    # switch to init_state reset
                    assert isinstance(env.env.env, RobomimicLowdimWrapper)
                    env.env.env.init_state = init_state

                env_seeds.append(train_idx)
                env_prefixs.append('train/')
                env_init_fn_dills.append(dill.dumps(init_fn))
        
        # test
        for i in range(n_test):
            seed = test_start_seed + i
            enable_render = i < n_test_vis

            def init_fn(env, seed=seed, 
                enable_render=enable_render):
                # setup rendering
                # video_wrapper
                assert isinstance(env.env, VideoRecordingWrapper)
                env.env.video_recoder.stop()
                env.env.file_path = None
                if enable_render:
                    if self.collect_data:
                        filename = pathlib.Path(output_dir).joinpath('media', f"episode_{seed - test_start_seed}.mp4")
                    else:
                        filename = pathlib.Path(output_dir).joinpath('media', f"{seed}_" + wv.util.generate_id() + ".mp4")
                    filename.parent.mkdir(parents=False, exist_ok=True)
                    filename = str(filename)
                    env.env.file_path = filename

                # switch to seed reset
                assert isinstance(env.env.env, RobomimicLowdimWrapper)
                env.env.env.init_state = None
                env.seed(seed)

            env_seeds.append(seed)
            env_prefixs.append('test/')
            env_init_fn_dills.append(dill.dumps(init_fn))
        
        env = AsyncVectorEnv(env_fns)
        # env = SyncVectorEnv(env_fns)

        self.env_meta = env_meta
        self.env = env
        self.env_fns = env_fns
        self.env_seeds = env_seeds
        self.env_prefixs = env_prefixs
        self.env_init_fn_dills = env_init_fn_dills
        self.fps = fps
        self.crf = crf
        self.n_obs_steps = n_obs_steps
        self.n_action_steps = n_action_steps
        self.n_latency_steps = n_latency_steps
        self.env_n_obs_steps = env_n_obs_steps
        self.past_action = past_action
        self.max_steps = max_steps
        self.rotation_transformer = rotation_transformer
        self.abs_action = abs_action
        self.tqdm_interval_sec = tqdm_interval_sec
        self.sampler = None
        self.n_samples = 0
        self.nmode = 0
        self.weak = None
        self.noise = 0.0
        self.decay = 1.0
        self.disruptor = None

    def set_sampler(self, sampler, nsample=1, nmode=1, noise=0.0, decay=1.0, tau=0.99):
        self.sampler = sampler
        self.n_samples = nsample
        self.nmode = nmode
        self.noise = noise
        self.decay = decay
        self.tau = tau
        if noise > 0:
            self.disruptor = NoiseGenerator(self.noise)
        print(colored(f'Set sampler: {sampler} {nsample}/{nmode}', 'yellow'))

    def set_reference(self, weak):
        self.weak = weak

    def run(self, policy: BaseLowdimPolicy):
        device = policy.device
        dtype = policy.dtype
        env = self.env
        
        # plan for rollout
        n_envs = len(self.env_fns)
        n_inits = len(self.env_init_fn_dills)
        n_chunks = math.ceil(n_inits / n_envs)

        # allocate data
        all_video_paths = [None] * n_inits
        all_rewards = [None] * n_inits
        all_steps_until_done = [None] * n_inits
        all_calls_until_done = np.ones((n_inits,), dtype=int)  # default for querying at least one
        all_infos = [None] * n_inits

        if self.collect_data:
            collect_observations = [[] for _ in range(n_inits)]
            collect_actions = [[] for _ in range(n_inits)]
            collect_rewards = [[] for _ in range(n_inits)]
            collect_terminals = [[] for _ in range(n_inits)]
            collect_infos = [[] for _ in range(n_inits)]
        else:
            collect_observations = collect_actions = collect_rewards = collect_terminals = collect_infos =  None

        for chunk_idx in range(n_chunks):
            start = chunk_idx * n_envs
            end = min(n_inits, start + n_envs)
            this_global_slice = slice(start, end)
            this_n_active_envs = end - start
            this_local_slice = slice(0,this_n_active_envs)

            if self.use_oracle_ac:
                raise NotImplementedError
            else:
                oracle_ac = None
            
            this_init_fns = self.env_init_fn_dills[this_global_slice]
            n_diff = n_envs - len(this_init_fns)
            if n_diff > 0:
                this_init_fns.extend([self.env_init_fn_dills[0]]*n_diff)
            assert len(this_init_fns) == n_envs

            # init envs
            env.call_each('run_dill_function', args_list=[(x,) for x in this_init_fns])

            # start rollout
            obs = env.reset()
            past_action = None
            policy.reset()

            env_name = self.env_meta['env_name']
            pbar = tqdm.tqdm(total=self.max_steps, desc=f"Eval {env_name}Lowdim {chunk_idx+1}/{n_chunks}", leave=False)
            done = False
            while not done:
                # create obs dict
                np_obs_dict = {
                    # handle n_latency_steps by discarding the last n_latency_steps
                    'obs': obs[:,-policy.n_obs_steps:].astype(np.float32)
                }
                if self.sampler in ['sg', 'sgac']:
                    prev_obs_dict = {
                        # handle n_latency_steps by discarding the last n_latency_steps
                        'obs': obs[:, -policy.n_obs_steps-1:-1].astype(np.float32)
                    }
                if self.past_action and (past_action is not None):
                    # TODO: not tested
                    np_obs_dict['past_action'] = past_action[:,-(self.n_obs_steps-1):].astype(np.float32)
                
                # device transfer
                obs_dict = dict_apply(np_obs_dict, lambda x: torch.from_numpy(x).to(device=device))
                # run policy
                with torch.no_grad():
                    if self.sampler == 'random':
                        action_dict = policy.predict_action(obs_dict)
                    elif self.sampler == 'ema':
                        if 'action_prior' not in locals():
                            action_prior = None
                        action_dict = ema_sampler(policy, action_prior, obs_dict, self.decay)
                        action_prior = action_dict['action_pred'][:, self.n_action_steps:]
                    elif self.sampler == 'contrast':
                        action_dict = contrastive_sampler(policy, self.weak, obs_dict, self.n_samples, self.nmode, self.sampler)
                    elif self.sampler == 'coherence':
                        if 'action_prior' not in locals():
                            action_prior = None
                        action_dict = coherence_sampler(policy, action_prior, obs_dict, self.n_samples, self.decay)
                        action_prior = action_dict['action_pred'][:, self.n_action_steps:]
                    elif self.sampler == 'bid':
                        if 'action_prior' not in locals():
                            action_prior = None
                        action_dict = bidirectional_sampler(policy, self.weak, obs_dict, action_prior, self.n_samples, self.decay, self.nmode)
                        action_prior = action_dict['action_pred'][:, self.n_action_steps:]                        
                    elif self.sampler == 'sg':
                        action_dict = policy.predict_action(obs_dict, prev_obs_dict)
                    elif self.sampler == 'ac':
                        if 'action_prior' not in locals():
                            action_prior = None
                        action_dict = ac_sampler(policy, action_prior, obs_dict, self.tau)
                        action_prior = action_dict['action_pred'][:, self.n_action_steps:]
                    elif self.sampler == 'sgac':
                        if 'action_prior' not in locals():
                            action_prior = None
                            action_dict = sgac_sampler(policy, action_prior, obs_dict, obs_dict, self.tau)
                        else:
                            action_dict = sgac_sampler(policy, action_prior, obs_dict, prev_obs_dict, self.tau)
                        action_prior = action_dict['action_pred'][:, self.n_action_steps:]
                    else:
                        action_dict = policy.predict_action(obs_dict)

                # device_transfer
                np_action_dict = dict_apply(action_dict, lambda x: x.detach().to('cpu').numpy())

                # handle latency_steps, we discard the first n_latency_steps actions to simulate latency
                action = np_action_dict['action'][:,self.n_latency_steps:]
                if not np.all(np.isfinite(action)):
                    print(action)
                    raise RuntimeError("Nan or Inf action")

                # noise
                if self.noise > 0.0:
                    noise_cum = self.disruptor.step(np_action_dict['action_pred'])
                    action += noise_cum[:, :action.shape[1]] * 0.1

                # step env
                env_action = action
                if self.abs_action:
                    env_action = self.undo_transform_action(action)

                if self.return_intermediate_state:  # Expose intermediate states during executing sequence of actions
                    if self.use_oracle_ac:
                        # At this point, always need to update action queue
                        if oracle_ac.first_time:
                            oracle_ac.update_action_chunk(env_action, replanning_mask=None)  # fill action for all envs at reset
                        else:
                            oracle_ac.update_action_chunk(env_action, replanning_mask=replanning_mask)

                        total_executed_steps = 0
                        while True:
                            single_step_action = oracle_ac.get_action()
                            obs, reward, done, info = env.step(single_step_action)
                            total_executed_steps += 1
                            replanning_mask = oracle_ac.compute_mask_to_replan(obs, reward, info, done, config=self.oracle_ac_config)
                            if replanning_mask.any():
                                break

                        query_mask = 1 - done  # 1 means query, 0 means no query
                        all_calls_until_done[start:end] = all_calls_until_done[start:end] + replanning_mask.astype(int)[0:end - start] * query_mask[0:end - start]
                        done = np.all(done)
                        past_action = action
                        # update pbar
                        pbar.update(total_executed_steps)

                    else:
                        for a_idx in range(self.n_action_steps):
                            single_step_action = env_action[:, a_idx:a_idx + 1, :]
                            obs, reward, done, info = env.step(single_step_action)

                            # Record data if in collect_data mode
                            if self.collect_data:
                                single_step_action_raw = action[:, a_idx:a_idx + 1, :]
                                for i in range(n_envs):
                                    collect_observations[chunk_idx * n_envs + i].append(obs[i, 0, ...])
                                    collect_actions[chunk_idx * n_envs + i].append(single_step_action_raw[i, 0, ...])
                                    collect_terminals[chunk_idx * n_envs + i].append(done[i])

                        query_mask = 1 - done  # 1 means query, 0 means no query
                        all_calls_until_done[start:end] = all_calls_until_done[start:end] + query_mask[0:end - start]
                        done = np.all(done)
                        past_action = action
                        # update pbar
                        pbar.update(action.shape[1])
                else:
                    obs, reward, done, info = env.step(env_action)
                    query_mask = 1 - done  # 1 means query, 0 means no query
                    all_calls_until_done[start:end] = all_calls_until_done[start:end] + query_mask[0:end - start]
                    done = np.all(done)
                    past_action = action

                    # update pbar
                    pbar.update(action.shape[1])
            pbar.close()

            # collect data for this round
            all_video_paths[this_global_slice] = env.render()[this_local_slice]
            all_rewards[this_global_slice] = env.call('get_attr', 'reward')[this_local_slice]
            all_steps_until_done[this_global_slice] = env.call('get_attr', 'step_elapsed')[this_local_slice]
            all_infos[this_global_slice] = env.call('get_attr', 'all_infos')[this_local_slice]
            if self.collect_data:
                for i in range(n_envs):
                    episode_reward = np.array(all_rewards[chunk_idx * n_envs + i])
                    collect_rewards[chunk_idx * n_envs + i].extend(episode_reward)
                    collect_infos[chunk_idx * n_envs + i].extend(all_infos[chunk_idx * n_envs + i])

        # log
        max_rewards = collections.defaultdict(list)
        successes = collections.defaultdict(list)
        env_step_till_max_reward = collections.defaultdict(list)
        env_step_till_done = collections.defaultdict(list)
        policy_step_till_done = collections.defaultdict(list)
        log_data = dict()
        # results reported in the paper are generated using the commented out line below
        # which will only report and average metrics from first n_envs initial condition and seeds
        # fortunately this won't invalidate our conclusion since
        # 1. This bug only affects the variance of metrics, not their mean
        # 2. All baseline methods are evaluated using the same code
        # to completely reproduce reported numbers, uncomment this line:
        # for i in range(len(self.env_fns)):
        # and comment out this line
        for i in range(n_inits):
            seed = self.env_seeds[i]
            prefix = self.env_prefixs[i]
            max_reward = np.max(all_rewards[i])
            success = float(max_reward == 1.0)

            max_rewards[prefix].append(max_reward)
            successes[prefix].append(success)
            env_step_till_max_reward[prefix].append(np.argmax(all_rewards[i]))
            env_step_till_done[prefix].append(all_steps_until_done[i])
            policy_step_till_done[prefix].append(all_calls_until_done[i])

            log_data[prefix+f'sim_max_reward_{seed}'] = max_reward
            log_data[prefix + f'sim_success_{seed}'] = success
            log_data[prefix + f'sim_step_to_max_reward_{seed}'] = float(np.argmax(all_rewards[i]))
            log_data[prefix + f'sim_step_to_success_{seed}'] = float(all_steps_until_done[i])
            log_data[prefix + f'sim_policy_call_to_success_{seed}'] = float(all_calls_until_done[i])

            # visualize sim
            video_path = all_video_paths[i]
            if video_path is not None:
                sim_video = wandb.Video(video_path)
                log_data[prefix+f'sim_video_{seed}'] = sim_video

        # log aggregate metrics
        for prefix, value in max_rewards.items():
            name = prefix+'mean_score'
            value = np.mean(value)
            log_data[name] = value

        for prefix, value in successes.items():
            name = prefix + 'mean_success'
            value = np.mean(value)
            log_data[name] = value

        for prefix, value in env_step_till_max_reward.items():
            name = prefix + 'mean_env_step_till_max_reward'
            value = np.mean(value)
            log_data[name] = value

        for prefix, value in env_step_till_done.items():
            name = prefix + 'mean_env_step_till_done'
            value = np.mean(value)
            log_data[name] = value

        for prefix, value in policy_step_till_done.items():
            name = prefix + 'mean_policy_step_till_done'
            value = np.mean(value)
            log_data[name] = value

        if self.collect_data:
            final_observations, final_actions, final_rewards, final_terminals, final_infos = [], [], [], [], []

            for i in range(n_inits):
                idx = np.argmax(collect_terminals[i]) + 1  # Find that first done
                final_observations.append(collect_observations[i][:idx + 1])  # include final obs of last action, thus +1
                final_actions.append(collect_actions[i][:idx])
                final_rewards.append(collect_rewards[i][:idx])
                final_terminals.append(collect_terminals[i][:idx])
                final_infos.append(collect_infos[i][:idx + 1])  # include final obs of last action, thus +1

            episode_data = {
                'observations': final_observations,
                'actions': final_actions,
                'rewards': final_rewards,
                'terminals': final_terminals,
                'infos': final_infos,
            }
            return log_data, episode_data
        else:
            return log_data

    def undo_transform_action(self, action):
        raw_shape = action.shape
        if raw_shape[-1] == 20:
            # dual arm
            action = action.reshape(-1,2,10)

        d_rot = action.shape[-1] - 4
        pos = action[...,:3]
        rot = action[...,3:3+d_rot]
        gripper = action[...,[-1]]
        rot = self.rotation_transformer.inverse(rot)
        uaction = np.concatenate([
            pos, rot, gripper
        ], axis=-1)

        if raw_shape[-1] == 20:
            # dual arm
            uaction = uaction.reshape(*raw_shape[:-1], 14)

        return uaction