serialexperimentsleon commited on
Commit
09ddfef
·
verified ·
1 Parent(s): 92d93f4

Add files using upload-large-folder tool

Browse files
.gitattributes CHANGED
@@ -199,3 +199,4 @@ wj867hap/eval_video/2_eval.mp4 filter=lfs diff=lfs merge=lfs -text
199
  wj867hap/eval_video/0_eval.mp4 filter=lfs diff=lfs merge=lfs -text
200
  3o6eq09y/episode_rosbags/episode_4_2024-12-31-21-17-59.bag filter=lfs diff=lfs merge=lfs -text
201
  nkdxms1x/episode_rosbags/episode_0_2024-12-31-21-03-48.bag filter=lfs diff=lfs merge=lfs -text
 
 
199
  wj867hap/eval_video/0_eval.mp4 filter=lfs diff=lfs merge=lfs -text
200
  3o6eq09y/episode_rosbags/episode_4_2024-12-31-21-17-59.bag filter=lfs diff=lfs merge=lfs -text
201
  nkdxms1x/episode_rosbags/episode_0_2024-12-31-21-03-48.bag filter=lfs diff=lfs merge=lfs -text
202
+ 3o6eq09y/episode_rosbags/episode_3_2024-12-31-21-17-17.bag filter=lfs diff=lfs merge=lfs -text
3o6eq09y/episode_rosbags/episode_3_2024-12-31-21-17-17.bag ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:a57f5b875bacf8950e34cf0e2c271563d03d68bdb62322fcdba680f69ed3f65e
3
+ size 1412768785
ksj6ie6d/.hydra/config.yaml ADDED
@@ -0,0 +1,350 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ root_dir: /home/${oc.env:USER}/fish_leon
2
+ nstep: 3
3
+ seed: 41
4
+ dataset_shuffle_seed: ${seed}
5
+ device: cuda
6
+ save_video: true
7
+ save_buffer: true
8
+ use_tb: true
9
+ baseline: false
10
+ use_wandb: true
11
+ eval: true
12
+ process_contact_features: ${eval}
13
+ obs_type: pixels
14
+ use_color: true
15
+ use_depth: true
16
+ use_masks: false
17
+ mask_list:
18
+ - EE_obj_mask
19
+ mask_representation: channels
20
+ crop_hw:
21
+ - 144
22
+ - 144
23
+ crop_down_offset: 48
24
+ color_crop_type: null
25
+ depth_crop_type: null
26
+ segmask_crop_type: null
27
+ add_crop_binary_mask: false
28
+ add_coord_conv_map: false
29
+ use_context_color: false
30
+ use_context_depth: false
31
+ use_context_segmask: false
32
+ context_color_crop_type: null
33
+ context_depth_crop_type: null
34
+ context_segmask_crop_type: null
35
+ context_add_crop_binary_mask: false
36
+ context_add_coord_conv_map: false
37
+ use_contact_map: false
38
+ use_sdf_maps: false
39
+ use_normals_maps: false
40
+ which_objects: both
41
+ max_contact_prob: 0.1
42
+ max_depth: 2.0
43
+ grasped_dtc_max_value: 0.105
44
+ env_dtc_max_value: 0.425
45
+ grasped_normals_mask_max_dtc_value: 0.105
46
+ env_normals_mask_max_dtc_value: 0.425
47
+ clamp_dtc: true
48
+ dtc_adaptive_normalization: false
49
+ mask_normals_within_sdf: true
50
+ adaptive_normals_mask: true
51
+ learnable_contact_preprocess_params: false
52
+ contact_model_name: local_multitask_outhd64all_home_crop_h144w144d48_ctxt_seed_183386_epoch_9
53
+ contact_estimation_model_ckpt_path: ~/fish_leon/contact_estimation/artifacts/175604_2/checkpoints/epoch=09-val_loss=0.00.ckpt
54
+ num_eval: 5
55
+ debug_timestamps: false
56
+ open_loop: false
57
+ action_trajectories: true
58
+ stop_after_action: false
59
+ interpolation_frequency: 25
60
+ policy_frequency: 5
61
+ wait_for_new_camera_frames: true
62
+ random_start: false
63
+ eval_starts: ${root_dir}/FISH/eval_starts/${suite.name}_${obs_type}/${task_name}
64
+ train_demo_idxs_list_or_num: null
65
+ num_valid_demos: null
66
+ val_num_groups: 3
67
+ name_of_expert_demo: 112_240x320_all_twodim_left_to_right_annotated_start_idx_5hz_zstd7_EE_pxl_coords_expert_demos_imp_act
68
+ expert_dataset_dirpath: ${root_dir}/FISH/expert_demos/${suite.name}/${task_name}/${name_of_expert_demo}
69
+ expert_dataset: ${expert_dataset_dirpath}/demos.zarr
70
+ action_key: ${oc.if_else:${action_trajectories}, 'action_trajectory_${interpolation_frequency}hz',
71
+ 'action'}
72
+ semantic_demo_grouping_name: semantic_demo_grouping.yaml
73
+ semantic_demo_grouping: ${expert_dataset_dirpath}/${semantic_demo_grouping_name}
74
+ expert_dataset_config: ${expert_dataset_dirpath}/demo_config.yaml
75
+ bc_regularize: false
76
+ bc_weight_type: qfilter
77
+ load_checkpoint: ${agent.load_checkpoint}
78
+ wandb_run_id: '1007_0'
79
+ true_action_history: false
80
+ wandb_notes: null
81
+ checkpoint_epoch: 12000
82
+ load_residual_weight: false
83
+ checkpoint_root_dir: /home/${oc.env:USER}/fish_leon/FISH
84
+ checkpoint_weight_dir: ${checkpoint_root_dir}/exp_local/${suite.name}_${obs_type}/${task_name}/${wandb_run_id}
85
+ residual_weight: ${root_dir}/FISH/weights/${suite.name}_${obs_type}/${task_name}/weight.pt
86
+ experiment_dir: ./exp_local/${suite.name}_${obs_type}/${task_name}/${wandb_run_id}
87
+ final_experiment_dir: ${experiment_dir}/${oc.generate_run_id:}
88
+ agent:
89
+ _target_: agent.diffusion_policy.DiffusionPolicyAgent
90
+ name: diffusion_policy
91
+ load_checkpoint: ${eval}
92
+ device: ${device}
93
+ n_obs_steps: ${.config.policy_cfg.n_obs_steps}
94
+ suite_name: ${suite.name}
95
+ obs_type: ${obs_type}
96
+ enable_arm: ${eval}
97
+ enable_camera: ${eval}
98
+ use_tb: ${use_tb}
99
+ desired_image_shape:
100
+ - 13
101
+ - 180
102
+ - 240
103
+ orig_cam_shape:
104
+ - 3
105
+ - 240
106
+ - 320
107
+ config:
108
+ _target_: agent.diffusion_policy.DiffusionPolicyAgentConfig
109
+ compile: false
110
+ device: ${device}
111
+ cam_resize_shape: ${agent.desired_image_shape}
112
+ orig_cam_shape: ${agent.orig_cam_shape}
113
+ policy_frequency: ${policy_frequency}
114
+ interpolation_frequency: ${interpolation_frequency}
115
+ policy_cfg:
116
+ _target_: lerobot.common.policies.diffusion.configuration_diffusion.DiffusionConfig
117
+ n_obs_steps: 1
118
+ horizon: 36
119
+ n_action_steps: ${agent.config.policy_cfg.horizon}
120
+ input_shapes:
121
+ observation.image: ${agent.config.cam_resize_shape}
122
+ context_observation.image: ${agent.config.cam_resize_shape}
123
+ observation.state:
124
+ - 8
125
+ observation.action_history:
126
+ - 7
127
+ output_shapes:
128
+ action:
129
+ - 7
130
+ input_normalization_modes:
131
+ observation.image: mean_std
132
+ observation.state: min_max
133
+ observation.action_history: min_max
134
+ output_normalization_modes:
135
+ action: min_max
136
+ vision_backbone: resnet18
137
+ pretrained_backbone_weights: null
138
+ transforms:
139
+ - _target_: torchaug.transforms.RandomAffine
140
+ degrees:
141
+ - -5
142
+ - 5
143
+ translate:
144
+ - 0.05
145
+ - 0.05
146
+ batch_transform: true
147
+ num_chunks: -1
148
+ batch_inplace: true
149
+ - _target_: torchaug.transforms.RandomColorJitter
150
+ brightness: 0.3
151
+ contrast: 0.4
152
+ saturation: 0.5
153
+ hue: 0.08
154
+ batch_transform: true
155
+ num_chunks: -1
156
+ batch_inplace: true
157
+ use_group_norm: true
158
+ spatial_softmax_num_keypoints: 32
159
+ action_history_encoder_config:
160
+ _target_: lerobot.common.policies.diffusion.configuration_diffusion.Unet1dEncoderConfig
161
+ in_channels: 7
162
+ out_channels: 32
163
+ history_length: ${agent.config.policy_cfg.n_action_steps}
164
+ kernel_size: ${agent.config.policy_cfg.kernel_size}
165
+ downsample_kernel_size: 3
166
+ downsample_stride: 2
167
+ downsample_padding: 1
168
+ down_dims:
169
+ - 256
170
+ - 512
171
+ - 1024
172
+ kernel_size: 5
173
+ n_groups: 8
174
+ diffusion_step_embed_dim: 128
175
+ use_film_scale_modulation: true
176
+ noise_scheduler_type: DDIM
177
+ beta_schedule: squaredcos_cap_v2
178
+ beta_start: 0.0001
179
+ beta_end: 0.02
180
+ prediction_type: epsilon
181
+ clip_sample: true
182
+ clip_sample_range: 1.0
183
+ num_train_timesteps: 50
184
+ num_inference_steps: 10
185
+ do_mask_loss_for_padding: false
186
+ train_cfg:
187
+ _target_: utils.TrainConfig
188
+ lr: 0.0001
189
+ lr_scheduler: cosine
190
+ lr_warmup_steps: 500
191
+ adam_betas:
192
+ - 0.95
193
+ - 0.999
194
+ adam_eps: 1.0e-08
195
+ adam_weight_decay: 1.0e-06
196
+ grad_clip_norm: 10
197
+ offline_steps: ${num_train_frames_diffusion}
198
+ use_amp: true
199
+ observation_cfg:
200
+ _target_: agent.encoder.VisualFeatureSet
201
+ use_depth: ${use_depth}
202
+ use_color: ${use_color}
203
+ mask_input_dict:
204
+ _target_: agent.encoder.MaskInputDict
205
+ enable: ${use_masks}
206
+ representation: ${mask_representation}
207
+ mask_list: ${mask_list}
208
+ crop_input_config:
209
+ _target_: agent.encoder.CropInputConfig
210
+ color_crop_type: ${color_crop_type}
211
+ depth_crop_type: ${depth_crop_type}
212
+ segmask_crop_type: ${segmask_crop_type}
213
+ crop_hw: ${crop_hw}
214
+ crop_down_offset: ${crop_down_offset}
215
+ add_crop_binary_mask: ${add_crop_binary_mask}
216
+ add_coord_conv_map: ${add_coord_conv_map}
217
+ context_input_config:
218
+ _target_: agent.encoder.ContextInputConfig
219
+ use_color: ${use_context_color}
220
+ use_depth: ${use_context_depth}
221
+ mask_input_dict:
222
+ _target_: agent.encoder.MaskInputDict
223
+ enable: ${use_context_segmask}
224
+ representation: ${mask_representation}
225
+ mask_list: ${mask_list}
226
+ crop_input_config:
227
+ _target_: agent.encoder.CropInputConfig
228
+ color_crop_type: ${context_color_crop_type}
229
+ depth_crop_type: ${context_depth_crop_type}
230
+ segmask_crop_type: ${context_segmask_crop_type}
231
+ crop_hw: ${crop_hw}
232
+ crop_down_offset: ${crop_down_offset}
233
+ add_crop_binary_mask: ${context_add_crop_binary_mask}
234
+ add_coord_conv_map: ${context_add_coord_conv_map}
235
+ mask_soft_approx_scheduler_config:
236
+ _target_: agent.encoder.MaskSoftApproxSchedulerConfig
237
+ num_steps: 40000
238
+ initial_value: 10.0
239
+ final_value: 1000.0
240
+ interpolation_scheme: constant
241
+ use_contact_map: ${use_contact_map}
242
+ use_sdf_maps: ${use_sdf_maps}
243
+ use_normals_maps: ${use_normals_maps}
244
+ which_objects: ${which_objects}
245
+ grasped_dtc_max_value: ${grasped_dtc_max_value}
246
+ env_dtc_max_value: ${env_dtc_max_value}
247
+ grasped_normals_mask_max_dtc_value: ${grasped_normals_mask_max_dtc_value}
248
+ env_normals_mask_max_dtc_value: ${env_normals_mask_max_dtc_value}
249
+ clamp_dtc: ${clamp_dtc}
250
+ max_contact_prob: ${max_contact_prob}
251
+ mask_normals_within_sdf: ${mask_normals_within_sdf}
252
+ dtc_adaptive_normalization: ${dtc_adaptive_normalization}
253
+ adaptive_normals_mask: ${adaptive_normals_mask}
254
+ max_depth: ${max_depth}
255
+ image_shape: ${agent.desired_image_shape}
256
+ learnable_contact_preprocess_params: ${learnable_contact_preprocess_params}
257
+ learning_rate: ${agent.config.train_cfg.lr}
258
+ weight_decay: 0.0
259
+ contact_model_name: ${contact_model_name}
260
+ zero_centered: false
261
+ suite:
262
+ suite: frankagym
263
+ name: frankagym
264
+ frame_stack: ${agent.n_obs_steps}
265
+ action_repeat: 1
266
+ discount: 0.99
267
+ hidden_dim: 1024
268
+ num_train_frames: 2010
269
+ num_seed_frames: 260
270
+ num_train_epochs: 5000
271
+ validate_every_epochs: 100
272
+ validate_diffusion_on_action_loss_every_epochs: 500
273
+ train_eval_diffusion_on_action_loss_every_epochs: 500
274
+ check_topk_every_epochs: 10
275
+ save_snapshot_every_epochs: 5000
276
+ eval_every_frames: 2000
277
+ num_eval_episodes: 5
278
+ save_snapshot: true
279
+ wait_for_user_to_start_episode: true
280
+ task_make_fn:
281
+ _target_: suite.frankagym.make
282
+ name: ${task_name}
283
+ height: 240
284
+ width: 320
285
+ frame_stack: ${suite.frame_stack}
286
+ action_repeat: ${suite.action_repeat}
287
+ seed: ${seed}
288
+ enable_arm: ${agent.enable_arm}
289
+ enable_gripper: ${enable_gripper}
290
+ start_with_gripper_open: ${start_with_gripper_open}
291
+ enable_camera: ${agent.enable_camera}
292
+ path_to_depth_extrinsics: ${path_to_depth_extrinsics}
293
+ contact_estimation_model_ckpt_path: ${contact_estimation_model_ckpt_path}
294
+ x_limit: ${x_limit}
295
+ y_limit: ${y_limit}
296
+ z_limit: ${z_limit}
297
+ device: ${device}
298
+ interpolation_frequency: ${interpolation_frequency}
299
+ policy_frequency: ${policy_frequency}
300
+ debug_timestamps: ${debug_timestamps}
301
+ stop_after_action: ${stop_after_action}
302
+ open_loop: ${open_loop}
303
+ wait_for_new_camera_frames: ${wait_for_new_camera_frames}
304
+ action_key: ${action_key}
305
+ action_trajectory_horizon: ${agent.config.policy_cfg.horizon}
306
+ action_trajectories: ${action_trajectories}
307
+ path_to_zarr_dataset: ${expert_dataset}
308
+ agent_policy_cfg: ???
309
+ true_action_history: ${true_action_history}
310
+ num_train_frames_bc: 50000
311
+ num_train_frames_drq: 1100000
312
+ stddev_schedule_drq: linear(1.0,0.1,100000)
313
+ task_name: FrankaInsertion-v1
314
+ num_train_frames_vinn: 25000
315
+ num_train_frames_diffusion: 1000000
316
+ num_train_epochs_bc: 5000
317
+ num_train_epochs_diffusion: 5000
318
+ validate_every_epochs_bc: 5
319
+ validate_every_epochs_diffusion: 25
320
+ validate_diffusion_on_action_loss_every_epochs: 50
321
+ train_eval_diffusion_on_action_loss_every_epochs: 500
322
+ check_topk_every_epochs: 5
323
+ check_topk_every_epochs_diffusion: ${validate_diffusion_on_action_loss_every_epochs}
324
+ save_snapshot_every_epochs_diffusion: 5000
325
+ x_limit:
326
+ - 0.2
327
+ - 0.7
328
+ y_limit:
329
+ - -0.4
330
+ - 0.4
331
+ z_limit:
332
+ - -0.05
333
+ - 0.55
334
+ home_displacement:
335
+ - 0.55
336
+ - 0.0
337
+ - 0.55
338
+ - 180.0
339
+ - 0.0
340
+ - 0.0
341
+ enable_gripper: true
342
+ start_with_gripper_open: true
343
+ offset_mask:
344
+ - 1
345
+ - 1
346
+ - 1
347
+ - 1
348
+ - 1
349
+ - 1
350
+ path_to_depth_extrinsics: ~/fish_leon/FISH/cfgs/camera_poses/camera_poses_L515/20240904-122305/color_tf_world.npy
ksj6ie6d/.hydra/hydra.yaml ADDED
@@ -0,0 +1,169 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ hydra:
2
+ run:
3
+ dir: ${final_experiment_dir}
4
+ sweep:
5
+ dir: ${final_experiment_dir}
6
+ subdir: ${hydra.job.num}
7
+ launcher:
8
+ submitit_folder: ${final_experiment_dir}/.slurm
9
+ timeout_min: 60
10
+ cpus_per_task: null
11
+ gpus_per_node: null
12
+ tasks_per_node: 1
13
+ mem_gb: null
14
+ nodes: 1
15
+ name: ${hydra.job.name}
16
+ stderr_to_stdout: false
17
+ _target_: hydra_plugins.hydra_submitit_launcher.submitit_launcher.LocalLauncher
18
+ sweeper:
19
+ _target_: hydra._internal.core_plugins.basic_sweeper.BasicSweeper
20
+ max_batch_size: null
21
+ params: null
22
+ help:
23
+ app_name: ${hydra.job.name}
24
+ header: '${hydra.help.app_name} is powered by Hydra.
25
+
26
+ '
27
+ footer: 'Powered by Hydra (https://hydra.cc)
28
+
29
+ Use --hydra-help to view Hydra specific help
30
+
31
+ '
32
+ template: '${hydra.help.header}
33
+
34
+ == Configuration groups ==
35
+
36
+ Compose your configuration from those groups (group=option)
37
+
38
+
39
+ $APP_CONFIG_GROUPS
40
+
41
+
42
+ == Config ==
43
+
44
+ Override anything in the config (foo.bar=value)
45
+
46
+
47
+ $CONFIG
48
+
49
+
50
+ ${hydra.help.footer}
51
+
52
+ '
53
+ hydra_help:
54
+ template: 'Hydra (${hydra.runtime.version})
55
+
56
+ See https://hydra.cc for more info.
57
+
58
+
59
+ == Flags ==
60
+
61
+ $FLAGS_HELP
62
+
63
+
64
+ == Configuration groups ==
65
+
66
+ Compose your configuration from those groups (For example, append hydra/job_logging=disabled
67
+ to command line)
68
+
69
+
70
+ $HYDRA_CONFIG_GROUPS
71
+
72
+
73
+ Use ''--cfg hydra'' to Show the Hydra config.
74
+
75
+ '
76
+ hydra_help: ???
77
+ hydra_logging:
78
+ version: 1
79
+ formatters:
80
+ simple:
81
+ format: '[%(asctime)s][HYDRA] %(message)s'
82
+ handlers:
83
+ console:
84
+ class: logging.StreamHandler
85
+ formatter: simple
86
+ stream: ext://sys.stdout
87
+ root:
88
+ level: INFO
89
+ handlers:
90
+ - console
91
+ loggers:
92
+ logging_example:
93
+ level: DEBUG
94
+ disable_existing_loggers: false
95
+ job_logging:
96
+ version: 1
97
+ formatters:
98
+ simple:
99
+ format: '[%(asctime)s][%(name)s][%(levelname)s] - %(message)s'
100
+ handlers:
101
+ console:
102
+ class: logging.StreamHandler
103
+ formatter: simple
104
+ stream: ext://sys.stdout
105
+ file:
106
+ class: logging.FileHandler
107
+ formatter: simple
108
+ filename: ${hydra.runtime.output_dir}/${hydra.job.name}.log
109
+ root:
110
+ level: INFO
111
+ handlers:
112
+ - console
113
+ - file
114
+ disable_existing_loggers: false
115
+ env: {}
116
+ mode: RUN
117
+ searchpath: []
118
+ callbacks: {}
119
+ output_subdir: .hydra
120
+ overrides:
121
+ hydra:
122
+ - hydra.mode=RUN
123
+ task:
124
+ - agent=diffusion
125
+ - suite=frankagym
126
+ - suite/frankagym_task@_global_=insertion
127
+ job:
128
+ name: eval_robot
129
+ chdir: true
130
+ override_dirname: agent=diffusion,suite/frankagym_task@_global_=insertion,suite=frankagym
131
+ id: ???
132
+ num: ???
133
+ config_name: config_eval
134
+ env_set: {}
135
+ env_copy: []
136
+ config:
137
+ override_dirname:
138
+ kv_sep: '='
139
+ item_sep: ','
140
+ exclude_keys: []
141
+ runtime:
142
+ version: 1.3.2
143
+ version_base: '1.1'
144
+ cwd: /home/leonmkim/fish_leon/FISH
145
+ config_sources:
146
+ - path: hydra.conf
147
+ schema: pkg
148
+ provider: hydra
149
+ - path: /home/leonmkim/fish_leon/FISH/cfgs
150
+ schema: file
151
+ provider: main
152
+ - path: ''
153
+ schema: structured
154
+ provider: schema
155
+ output_dir: /home/leonmkim/fish_leon/FISH/exp_local/frankagym_pixels/FrankaInsertion-v1/1007_0/ksj6ie6d
156
+ choices:
157
+ suite: frankagym
158
+ suite/frankagym_task@_global_: insertion
159
+ agent: diffusion
160
+ hydra/env: default
161
+ hydra/callbacks: null
162
+ hydra/job_logging: default
163
+ hydra/hydra_logging: default
164
+ hydra/hydra_help: default
165
+ hydra/help: default
166
+ hydra/sweeper: basic
167
+ hydra/launcher: submitit_local
168
+ hydra/output: default
169
+ verbose: false
ksj6ie6d/.hydra/overrides.yaml ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ - agent=diffusion
2
+ - suite=frankagym
3
+ - suite/frankagym_task@_global_=insertion
ksj6ie6d/tb/events.out.tfevents.1735697318.leonmkim-ROG-Strix-G15CS-G15CS.1741133.0 ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:95740a1c217bacdbf8973af251ea8461fbed9e4a24fd87b124da865d5b9dbc17
3
+ size 88
ksj6ie6d/wandb/run-20241231_210837-ksj6ie6d/files/code/FISH/eval_robot.py ADDED
@@ -0,0 +1,599 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ #%%
2
+ import warnings
3
+ import os
4
+
5
+ os.environ['MKL_SERVICE_FORCE_INTEL'] = '1'
6
+ os.environ['MUJOCO_GL'] = 'egl'
7
+ from pathlib import Path
8
+ #%%
9
+ import hydra
10
+ import numpy as np
11
+ import torch
12
+
13
+ import utils
14
+ from utils import get_feature_dirname_from_configs
15
+
16
+ from video import VideoRecorder
17
+ import pickle
18
+ import time
19
+ import threading
20
+ import shutil
21
+ from logger import Logger
22
+
23
+ import wandb
24
+ from omegaconf import OmegaConf, open_dict
25
+
26
+ from replay_buffer_robot import RosbagEvalReplayBufferStorage
27
+ from lerobot.common.utils.utils import _relative_path_between
28
+
29
+ torch.backends.cudnn.benchmark = True
30
+ warnings.filterwarnings('ignore', category=DeprecationWarning)
31
+
32
+ # import specs for replay buffer
33
+ from dm_env import specs
34
+
35
+ import sys, signal
36
+ import yaml
37
+
38
+ import binomial_cis as bc
39
+
40
+ # get path of current file
41
+ current_path = os.path.dirname(os.path.realpath(__file__))
42
+ sys.path.append(os.path.join(current_path, os.pardir))
43
+ # from contact_estimation.src.utils.viz_utils import normalized_surface_normal_to_rgb, depth_map_to_im, grasped_env_dtc_map_to_im, contact_prob_map_to_im, desaturate_color_image, masked_overlay_im_list
44
+
45
+ def make_agent(obs_spec, action_spec, cfg):
46
+ cfg.obs_shape = obs_spec['pixels'].shape
47
+ dataset_statistics = None # this will be loaded from the checkpoint
48
+ try:
49
+ cfg.action_shape = action_spec.shape
50
+ except:
51
+ pass
52
+ return hydra.utils.instantiate(cfg, dataset_statistics)
53
+
54
+ class Workspace:
55
+ def __init__(self, cfg):
56
+ self.work_dir = Path.cwd()
57
+ print(f'workspace: {self.work_dir}')
58
+
59
+ signal.signal(signal.SIGINT, self.signal_handler)
60
+
61
+ self.cfg = cfg
62
+ self.loading_uncompiled_checkpoint_with_compile = False
63
+ self.loading_compiled_checkpoint_with_no_compile = False
64
+
65
+ snapshot_path = Path(self.cfg.checkpoint_weight_dir) / f'snapshot_{self.cfg.checkpoint_epoch}.pt'
66
+ self.load_checkpoint_conf(snapshot_path=snapshot_path)
67
+
68
+ # load config for action trajectories
69
+ utils.set_seed_everywhere(self.cfg.seed)
70
+ self.device = torch.device(self.cfg.device)
71
+ self.setup()
72
+
73
+ # self.agent = make_agent(self.eval_env.observation_spec(),
74
+ # self.eval_env.action_spec(), self.cfg.agent)
75
+ self.timer = utils.Timer()
76
+ # self._global_step = 0
77
+ self._global_episode = 0
78
+ self._global_epoch = 0
79
+ self.num_episode_successes = 0
80
+
81
+ self.alpha_range = [.01, .025, .05, .1]
82
+
83
+ # Need to convert hydra config to primitive container for wandb https://docs.wandb.ai/guides/integrations/hydra
84
+ with open_dict(self.cfg):
85
+ self.cfg.feature_type = get_feature_dirname_from_configs(
86
+ hydra.utils.instantiate(self.cfg.agent.config.observation_cfg),
87
+ self.cfg.agent.config.policy_cfg.input_shapes,
88
+ hydra.utils.instantiate(self.cfg.agent.config.policy_cfg.action_history_encoder_config) if 'observation.action_history' in self.cfg.agent.config.policy_cfg.input_shapes else None,
89
+ )
90
+
91
+ wandb_config = OmegaConf.to_container(
92
+ self.cfg, resolve=True, throw_on_missing=True
93
+ )
94
+ # must be called before any tf summary writer is created
95
+ if self.cfg.use_wandb:
96
+ # get the run id from the final_experiment_dir directory
97
+ run_id = os.path.basename(os.path.normpath(self.cfg.final_experiment_dir))
98
+ wandb.init(project='extrinsic_contact_downstream', entity='serialexperimentsleon', job_type='eval', sync_tensorboard=self.cfg.use_tb, config=wandb_config, id=run_id)
99
+
100
+ self.logger = Logger(self.work_dir, use_tb=self.cfg.use_tb, use_wandb=self.cfg.use_wandb)
101
+
102
+ # if not self.loading_uncompiled_checkpoint_with_compile and self.cfg.agent.config.compile:
103
+ # self.agent.compile_modules()
104
+
105
+ # self.load_checkpoint(snapshot_path=snapshot_path)
106
+
107
+ # if self.loading_uncompiled_checkpoint_with_compile: # need to call compile after loading the checkpoint
108
+ # self.agent.compile_modules()
109
+
110
+ print(f"loaded agent with feature_type: {self.cfg.feature_type}")
111
+
112
+ def check_for_key_press(self):
113
+ while self.continue_keypress_thread:
114
+ inp = input("Press 'r' to restart current episode, 'n' to stop current episode and skip to next, 'q' to break entire eval\n")
115
+ if inp == 'n':
116
+ self.preempt_episode = True
117
+ print("preempting episode")
118
+ elif inp in ['', '0', '1']: # enter key
119
+ if inp in ['0', '1']:
120
+ self.num_episode_successes += int(inp)
121
+ self.proceed_after_env_reset_event.set()
122
+ print("proceeding to start episode!")
123
+ elif inp == 'q':
124
+ self.proceed_after_env_reset_event.set()
125
+ self.preempt_episode = True
126
+ self.exit_eval = True
127
+ self.continue_keypress_thread = False # will stop the keypress thread
128
+ print("quitting eval")
129
+ break
130
+ elif inp == 'r':
131
+ print('restarting episode')
132
+ self.preempt_episode = True
133
+ self.restart_episode = True
134
+ else:
135
+ print("Invalid key press, try again")
136
+
137
+ # self.keypress_input_thread.join() # wait for the keypress thread to finish
138
+
139
+ def signal_handler(self, signal, frame):
140
+ print("\nprogram exiting gracefully")
141
+ self.proceed_after_env_reset_event.set()
142
+ self.preempt_episode = True
143
+ self.exit_eval = True
144
+ self.continue_keypress_thread = False # will stop the keypress thread
145
+ self.keypress_input_thread.join() # wait for the keypress thread to finish
146
+ video_filepath = self.video_recorder.save()
147
+ # get the video file and convert to video tensor to log
148
+ self.logger.log_video('eval/video', video_filepath, self.global_step)
149
+ wandb.finish()
150
+ sys.exit(0)
151
+
152
+ def setup(self):
153
+ # create envs
154
+ self.eval_env = hydra.utils.call(self.cfg.suite.task_make_fn)
155
+ # expert_demo_config_path = os.path.join(os.path.dirname(self.cfg.expert_dataset), 'demo_config.yaml')
156
+ # self.expert_demo_config = yaml.load(open(expert_demo_config_path, 'r'), Loader=yaml.FullLoader)
157
+ # self.eval_env._env.action_trans_norm = expert_demo_config['max_translation_action_norm']
158
+ # self.eval_env._env.action_rot_norm = expert_demo_config['max_rotation_action_norm']
159
+ # self.eval_env._env.action_period = expert_demo_config['sample_period']
160
+ # print(f"setting max_translation_action_norm to {expert_demo_config['max_translation_action_norm']} and sample_period to {expert_demo_config['sample_period']}")
161
+ # print(f"setting max_rotation_action_norm to {expert_demo_config['max_rotation_action_norm']}")
162
+
163
+ # self.eval_env.set_demo_params(self.cfg.expert_dataset)
164
+
165
+ # Turn off random start
166
+ self.eval_env.random_start = False
167
+
168
+ # create replay buffer
169
+ # data_specs = [
170
+ # {
171
+ # 'observation': self.eval_env.observation_spec(),
172
+ # },
173
+ # # self.eval_env.observation_spec()['features'],
174
+ # self.eval_env.action_spec(),
175
+ # specs.Array(self.eval_env.action_spec().shape, self.eval_env.action_spec().dtype, 'vinn_action'),
176
+ # specs.Array((1, ), np.float32, 'reward'),
177
+ # specs.Array((1, ), np.float32, 'discount'),
178
+ # ]
179
+
180
+ # self.eval_replay_storage = ZarrEvalReplayBufferStorage(data_specs, self.work_dir / 'eval_buffer', debug_timestamps=self.cfg.debug_timestamps, save_buffer=self.cfg.save_buffer, debug_info_data_specs=self.eval_env.debug_info_data_specs, camera_info_dict=self.eval_env.get_camera_info_dict())
181
+ self.eval_replay_storage = RosbagEvalReplayBufferStorage(self.work_dir)
182
+
183
+ self.video_recorder = VideoRecorder(
184
+ self.work_dir if self.cfg.save_video else None,
185
+ ros_enabled=True,
186
+ fps=self.cfg.agent.config.policy_frequency,
187
+ )
188
+
189
+ print('workspace setup complete')
190
+
191
+ @property
192
+ def global_step(self):
193
+ # return self._global_step
194
+ return self.eval_env.get_global_step()
195
+
196
+ @property
197
+ def global_episode(self):
198
+ return self._global_episode
199
+
200
+ @property
201
+ def global_frame(self):
202
+ return self.global_step * self.cfg.action_repeat
203
+
204
+ @property
205
+ def global_epoch(self):
206
+ return self._global_epoch
207
+
208
+ def reset(self, eval_idx):
209
+ if not self.eval_env.enable_arm:
210
+ return np.array([0,0,0], dtype=np.float32)
211
+ self.eval_env.arm_refresh(reset=False)
212
+ # Set start position
213
+ try:
214
+ self.eval_env.set_position(self.start_pos[eval_idx])
215
+ except:
216
+ self.eval_env.arm.set_position(self.start_pos[eval_idx])
217
+ if self.eval_env.arm.keep_gripper_closed:
218
+ self.eval_env.arm.close_gripper_fully()
219
+ else:
220
+ self.eval_env.arm.open_gripper_fully()
221
+ time.sleep(0.1)
222
+ time_step = self.eval_env.step(np.zeros(self.eval_env.action_spec().shape[0], dtype=np.float32),
223
+ np.zeros(self.eval_env.action_spec().shape[0], dtype=np.float32))
224
+ return time_step
225
+
226
+ def eval(self):
227
+ # before evals start, prompt user for name of grasped object and the left book of the slot location
228
+ grasped_obj_name = input("Enter the name of the grasped object: ")
229
+ left_book_slot = input("Enter the left book slot location: ")
230
+ # update wandb config
231
+ if self.cfg.use_wandb:
232
+ wandb.config.update({'grasped_obj_name': grasped_obj_name, 'left_book_slot': left_book_slot})
233
+
234
+ self.preempt_episode = False
235
+ self.exit_eval = False
236
+ self.restart_episode = False
237
+
238
+ self.continue_keypress_thread = True
239
+ self.proceed_after_env_reset_event = threading.Event()
240
+ self.keypress_input_thread = threading.Thread(target=self.check_for_key_press)
241
+ self.keypress_input_thread.start()
242
+
243
+ # # Set model to eval mode
244
+ # self.agent.train(False)
245
+
246
+ eval_until_episode = utils.Until(self.cfg.num_eval)
247
+
248
+ self.use_action_history = False
249
+ # if "dp" in repr(self.agent) and "observation.action_history" in self.cfg.agent.config.policy_cfg.input_shapes:
250
+ if "observation.action_history" in self.cfg.agent.config.policy_cfg.input_shapes:
251
+ self.use_action_history = True
252
+
253
+ # self.eval_replay_storage._new_eval_step(0)
254
+
255
+ # if 'vinn' in repr(self.agent) or 'openloop' in repr(self.agent):
256
+ # with open(self.cfg.expert_dataset, 'rb') as f:
257
+ # if self.cfg.obs_type == 'pixels':
258
+ # self.expert_demo, _, self.expert_action, self.expert_reward = pickle.load(f)
259
+ # elif self.cfg.obs_type == 'features':
260
+ # _, self.expert_demo, self.expert_action, self.expert_reward = pickle.load(f)
261
+
262
+ # if self.cfg.action_trajectories:
263
+ # with open(self.cfg.expert_action_trajectories, 'rb') as f:
264
+ # self.expert_action = pickle.load(f)
265
+
266
+ # if isinstance(self.cfg.train_demo_idxs_list_or_num, int):
267
+ # if self.cfg.train_demo_idxs_list_or_num == -1:
268
+ # self.cfg.train_demo_idxs_list_or_num = len(self.expert_demo)
269
+ # train_demo_idxs_list_or_num = list(range(self.cfg.train_demo_idxs_list_or_num))
270
+
271
+ # self.expert_demo = self.expert_demo[train_demo_idxs_list_or_num]
272
+ # self.expert_action = self.expert_action[train_demo_idxs_list_or_num]
273
+ # self.expert_reward = self.expert_reward[train_demo_idxs_list_or_num]
274
+ # # if self.cfg.action_plans:
275
+ # # self.expert_action_plans = self.expert_action_plans[self.cfg.train_demo_idxs_list_or_num]
276
+ # # self.expert_demo = self.expert_demo[:self.cfg.num_demos]
277
+ # # self.expert_action = self.expert_action[:self.cfg.num_demos]
278
+ # # self.expert_reward = self.expert_reward[:self.cfg.num_demos]
279
+
280
+ # self.expert_demo = np.concatenate(self.expert_demo, axis=0)
281
+ # self.expert_rgb_obs = np.ascontiguousarray(np.transpose(self.expert_demo, (0,2,3,1))[:, :,:,:3].astype(np.uint8))
282
+ # self.expert_action = np.concatenate(self.expert_action, axis=0)
283
+
284
+ # self.agent.save_representations(self.expert_demo, self.expert_action, 128, config=self.expert_demo_config)
285
+
286
+ # Get start points
287
+ if self.cfg.random_start:
288
+ eval_starts = Path(self.cfg.eval_starts) / 'starts.pkl'
289
+ if eval_starts.exists():
290
+ with eval_starts.open('rb') as f:
291
+ self.start_pos = pickle.load(f)
292
+ else:
293
+ eval_starts = Path(self.cfg.eval_starts)
294
+ eval_starts.mkdir(parents=True, exist_ok=True)
295
+
296
+ # Generate start points
297
+ self.start_pos = []
298
+ try:
299
+ for _ in range(self.cfg.num_eval):
300
+ self.start_pos.append(self.eval_env.get_random_pos())
301
+ except:
302
+ for _ in range(self.cfg.num_eval):
303
+ self.start_pos.append(self.eval_env.arm.get_random_pos())
304
+
305
+ # Save start points for the task
306
+ eval_starts = eval_starts / 'starts.pkl'
307
+ with eval_starts.open('wb') as f:
308
+ pickle.dump(self.start_pos, f)
309
+
310
+ time_step = self.eval_env.reset()
311
+ # replay_thread = None
312
+ while eval_until_episode(self.global_episode) and not self.exit_eval:
313
+ # self.video_recorder.init(self.eval_env, video_filename=f'{self.global_episode}_eval.mp4')
314
+ print(f"Starting episode {self.global_episode}")
315
+ time_step = self.eval_env.reset() #Leon: need to call reset twice in case objects are trapped
316
+ self.video_recorder.init(self.eval_env, video_filename=f'{self.global_episode}_eval.mp4')
317
+ # x = input("Press Enter to continue... after reseting env")
318
+ print("Press Enter to continue... after reseting env. To rate prev episode, press 0 for failure and 1 for success")
319
+ self.proceed_after_env_reset_event.clear() # clear the event flag
320
+ self.proceed_after_env_reset_event.wait() # blocking wait for the event flag to be set
321
+ if self.global_episode > 0:
322
+ self.logger.log_metrics({'num_success': self.num_episode_successes}, self.global_step, 'eval', episode=self.global_episode)
323
+ self.logger.log_metrics({'success_rate': self.num_episode_successes/self.global_episode}, self.global_step, 'eval', episode=self.global_episode)
324
+
325
+ # log confidence intervals for success rate
326
+ k = self.num_episode_successes # number of successes
327
+ n = self.global_episode # number of trials
328
+
329
+ table_columns = []
330
+ table_data = []
331
+ for alpha in self.alpha_range:
332
+ lb = bc.binom_ci(k, n, alpha, 'lb')
333
+ ub = bc.binom_ci(k, n, alpha, 'ub')
334
+
335
+ self.logger.log_metrics({f'success_rate_lb_{alpha}': lb}, self.global_step, 'eval', episode=self.global_episode)
336
+ self.logger.log_metrics({f'success_rate_ub_{alpha}': ub}, self.global_step, 'eval', episode=self.global_episode)
337
+
338
+ time_step = self.eval_env.reset()
339
+ # debug_info_dict = self.eval_env.debug_info_dict
340
+ # if replay_thread is not None:
341
+ # # wait for the last replay thread to finish
342
+ # replay_thread.join()
343
+
344
+ # self.eval_replay_storage.add(time_step._replace(observation=time_step.observation[self.cfg.obs_type]), debug_info_dict)
345
+ # replay_thread = threading.Thread(target=self.eval_replay_storage.add, args=(time_step._replace(observation=time_step.observation[self.cfg.obs_type]), debug_info_dict))
346
+ # replay_thread = threading.Thread(target=self.eval_replay_storage.add, args=(time_step, debug_info_dict))
347
+
348
+ # replay_thread.start()
349
+ if self.cfg.random_start:
350
+ time_step = self.reset(self.global_episode)
351
+ time.sleep(2) #5)
352
+ # if 'vinn' in repr(self.agent):
353
+ # self.agent.reset()
354
+ # # self.agent.buffer.reset()
355
+ # # if self.cfg.open_loop:
356
+ # # self.agent.current_step = 0
357
+ # if 'openloop' in repr(self.agent):
358
+ # self.agent.curr_step = 0
359
+ # at start of each episode, provide zero action for policies that use action history
360
+ # shape should be (T_o, T_a, action_dim)
361
+
362
+ # while not time_step.last() and not self.preempt_episode:
363
+ self.video_recorder.ros_start_recording()
364
+ self.eval_replay_storage.start_episode()
365
+ self.eval_env.start_policy_timer()
366
+ while not self.eval_env.episode_done() and not self.preempt_episode:
367
+ # with torch.no_grad(), utils.eval_mode(self.agent):
368
+ # # if self.cfg.agent.provide_topk:
369
+ # # action, vinn_action, topk = self.agent.act(
370
+ # # time_step.observation['pixels'],
371
+ # # self.global_step,
372
+ # # eval_mode=True)
373
+ # # elif self.cfg.agent.provide_obs:
374
+ # # action, vinn_action, obs = self.agent.act(
375
+ # # time_step.observation['pixels'],
376
+ # # self.global_step,
377
+ # # eval_mode=True)
378
+ # # else:
379
+ # action, vinn_action = self.agent.act(
380
+ # time_step.observation,
381
+ # self.global_step,
382
+ # eval_mode=True,
383
+ # obs_timestamp=time_step.observation['timestamp'],
384
+ # obs_seq=time_step.observation['seq'],
385
+ # action_history=action_history,
386
+ # action_history_start_timestamp=action_history_start_timestamp,
387
+ # )
388
+ # DONT WAIT FOR POLICY TO GET AN ACTION
389
+ # we dont want to slow down grabbing obs and passing to sam/contact features
390
+
391
+ self.eval_env.run_policy_threads() # this just does a rospy sleep
392
+
393
+ # if self.use_action_history:
394
+ # action_history_start_timestamp = time_step.observation['timestamp']
395
+ # # action_history = action[:self.cfg.agent.config.policy_cfg.action_history_encoder_config.history_length, ...]
396
+ # # add n_obs_steps dimension to action_history, for now we assume n_obs_steps = 1
397
+ # # TODO: handle n_obs_steps > 1
398
+ # action_history = action[np.newaxis, ...]
399
+
400
+ # time_step = self.eval_env.step(action, vinn_action) # obs, reward after action has been taken
401
+ # debug_info_dict = self.eval_env.debug_info_dict
402
+
403
+ # time_step = self.eval_env.ros_step()
404
+
405
+ # replay_thread.join()
406
+
407
+ # time how long it takes to execute the step
408
+ # time_before_add = time.perf_counter()
409
+ # self.eval_replay_storage.add(time_step._replace(observation=time_step.observation[self.cfg.obs_type]), debug_info_dict)
410
+ # use thread to call the add function in a separate thread
411
+ # replay_thread = threading.Thread(target=self.eval_replay_storage.add, args=(time_step._replace(observation=time_step.observation[self.cfg.obs_type]), debug_info_dict))
412
+
413
+ # replay_thread = threading.Thread(target=self.eval_replay_storage.add, args=(time_step, debug_info_dict))
414
+ # replay_thread.start()
415
+
416
+ # print(f"Time to add to replay buffer: {time.perf_counter() - time_before_add}")
417
+
418
+ # self.video_recorder.record(self.eval_env)
419
+ # self._global_step += 1
420
+
421
+ self.eval_env.stop_policy_timer()
422
+
423
+ if self.restart_episode:
424
+ # means we should delete the current episode and start again
425
+ self.restart_episode = False
426
+ self.eval_replay_storage.reset_current_episode()
427
+ self.video_recorder.reset_current_episode()
428
+
429
+ else:
430
+ self.eval_replay_storage.store_current_episode()
431
+ video_filepath = self.video_recorder.save()
432
+ self.logger.log_video(f"eval/{video_filepath.name.rstrip('.mp4')}", video_filepath, self.global_step)
433
+ self._global_episode += 1
434
+
435
+ self.preempt_episode = False # reset preempt_episode flag
436
+
437
+ # self.video_recorder.save(f'{episode}_eval.mp4')
438
+ # get the video file and convert to video tensor to log
439
+
440
+ self.eval_env.reset()
441
+
442
+ print("Evaluation finished. To wrap up, rate prev episode, press 0 for failure and 1 for success")
443
+ self.proceed_after_env_reset_event.clear() # clear the event flag
444
+ self.proceed_after_env_reset_event.wait() # blocking wait for the event flag to be set
445
+ if self.global_episode > 0:
446
+ # self.logger.log_metrics({'num_success': self.num_episode_successes}, self.global_step, 'eval', episode=self.global_episode)
447
+ self.logger.log_metrics({'num_success': self.num_episode_successes}, self.global_step, 'eval', episode=self.global_episode)
448
+ self.logger.log_metrics({'success_rate': self.num_episode_successes/self.global_episode}, self.global_step, 'eval', episode=self.global_episode)
449
+
450
+ # log confidence intervals for success rate
451
+ k = self.num_episode_successes # number of successes
452
+ n = self.global_episode # number of trials
453
+
454
+ table_columns = ['success_rate']
455
+ table_data = [self.num_episode_successes/self.global_episode]
456
+ for alpha in self.alpha_range:
457
+ lb = bc.binom_ci(k, n, alpha, 'lb')
458
+ ub = bc.binom_ci(k, n, alpha, 'ub')
459
+
460
+ self.logger.log_metrics({f'success_rate_lb_{alpha}': lb}, self.global_step, 'eval', episode=self.global_episode)
461
+ self.logger.log_metrics({f'success_rate_ub_{alpha}': ub}, self.global_step, 'eval', episode=self.global_episode)
462
+
463
+ table_columns.extend([f'success_rate_lb_{alpha}', f'success_rate_ub_{alpha}'])
464
+ table_data.extend([lb, ub])
465
+
466
+ table_data = [table_data]
467
+
468
+ # seperately log as a table
469
+ wandb.log({
470
+ "eval/success_rate_ci": wandb.Table(data=table_data, columns=table_columns)
471
+ })
472
+
473
+ # also accumulate eval metrics across previous eval runs
474
+ # TODO: change wandb init to resume from an existing run!!!
475
+ run_filter={
476
+ "jobType": "eval",
477
+ "config.wandb_run_id": self.cfg.wandb_run_id,
478
+ "summary_metrics.episode": {"$gte": 5},
479
+ "config.checkpoint_epoch": self.cfg.checkpoint_epoch,
480
+ "state": "finished",
481
+ "config.grasped_obj_name": grasped_obj_name,
482
+ "config.left_book_slot": left_book_slot,
483
+ }
484
+
485
+ api = wandb.Api()
486
+ filtered_runs = api.runs("serialexperimentsleon/extrinsic_contact_downstream", filters=run_filter)
487
+ total_num_successes = self.num_episode_successes
488
+ total_num_episodes = self.global_episode
489
+ list_of_historical_run_ids = []
490
+ if len(filtered_runs) > 0:
491
+ for filtered_run in filtered_runs:
492
+ total_num_successes += filtered_run.summary_metrics['eval/num_success']
493
+ total_num_episodes += filtered_run.summary_metrics['episode']
494
+ list_of_historical_run_ids.append(filtered_run.id)
495
+
496
+ wandb.summary['total_num_successes'] = total_num_successes
497
+ wandb.summary['total_num_episodes'] = total_num_episodes
498
+ wandb.summary['total_success_rate'] = total_num_successes/total_num_episodes
499
+
500
+ # log the accumulated metrics as a table
501
+ total_table_columns = ['total_num_successes', 'total_num_episodes', 'total_success_rate']
502
+ total_table_data = [total_num_successes, total_num_episodes, total_num_successes/total_num_episodes]
503
+ for alpha in self.alpha_range:
504
+ lb = bc.binom_ci(total_num_successes, total_num_episodes, alpha, 'lb')
505
+ ub = bc.binom_ci(total_num_successes, total_num_episodes, alpha, 'ub')
506
+ total_table_columns.extend([f'total_success_rate_lb_{alpha}', f'total_success_rate_ub_{alpha}'])
507
+ total_table_data.extend([lb, ub])
508
+ wandb.summary[f'total_success_rate_lb_{alpha}'] = lb
509
+ wandb.summary[f'total_success_rate_ub_{alpha}'] = ub
510
+
511
+ total_table_data = [total_table_data]
512
+ wandb.log({
513
+ 'eval/total_success_rate_ci': wandb.Table(data=total_table_data, columns=total_table_columns)
514
+ })
515
+
516
+ self.continue_keypress_thread = False # will stop the keypress thread
517
+ self.keypress_input_thread.join() # wait for the keypress thread to finish
518
+
519
+ def load_checkpoint_conf(self, snapshot_path):
520
+ config_path = snapshot_path.parent / 'config.yaml'
521
+ if not config_path.exists():
522
+ raise FileNotFoundError(f'No snapshot conf found at {config_path}')
523
+ else:
524
+ # load the omegaconf config
525
+ hydra.core.global_hydra.GlobalHydra.instance().clear()
526
+ hydra.initialize(
527
+ str(_relative_path_between(Path(config_path).absolute().parent, Path(__file__).absolute().parent)),
528
+ )
529
+ cfg = hydra.compose(Path(config_path).stem)
530
+ from deepdiff import DeepDiff
531
+ from omegaconf import open_dict
532
+ diff = DeepDiff(OmegaConf.to_container(cfg), OmegaConf.to_container(self.cfg)) # old, new
533
+ # import re
534
+ overwriteable_keys = [f"root{overwritable_key}" for overwritable_key in ["['use_wandb']", "['path_to_depth_extrinsics']", "['eval']", "['root_dir']", "['wandb_notes']", "['agent']['config']['train_cfg']['use_amp']", "['agent']['config']['compile']", "['agent']['config']['policy_cfg']['num_inference_steps']"]]
535
+ if "values_changed" in diff:
536
+ # top_k_checkpoints, wandb_notes, agent.config.train_cfg.use_amp, save_snapshot_every_epochs_diffusion, check_topk_every_epochs_diffusion, validate_diffusion_on_action_loss_every_epochs, train_eval_diffusion_on_action_loss_every_epochs, validate_every_epochs_diffusion
537
+ # for keys above, overwrite the old config with the new config
538
+ for k, v in diff['values_changed'].items():
539
+ # replace any keys that are under "root['suite']"
540
+ if k in overwriteable_keys or k.startswith("root['suite']"):
541
+ print(f"Found changed key {k} with value {v}. Overwriting old checkpoint config")
542
+ if k == "root['agent']['config']['compile']":
543
+ if diff['values_changed'][k]['new_value']:
544
+ self.loading_uncompiled_checkpoint_with_compile = True
545
+ elif not diff['values_changed'][k]['new_value']:
546
+ # raise ValueError("Cannot load a compiled checkpoint without compile")
547
+ self.loading_compiled_checkpoint_with_no_compile = True
548
+ exec(f"{k.replace('root[', 'cfg[')} = {k.replace('root[', 'self.cfg[')}")
549
+ # for any new values, update the old checkpoint config
550
+ if "dictionary_item_added" in diff:
551
+ for new_key in diff['dictionary_item_added']: # this is a list
552
+ # if new_key == "root['suite']['task_make_fn']['observation_cfg']":
553
+ if new_key == "root['suite']['task_make_fn']['agent_policy_cfg']":
554
+ # pass the agents observation_cfg to the suite task_make_fn
555
+ with open_dict(cfg): # to allow addition of non-existing keys
556
+ # cfg.suite.task_make_fn.observation_cfg = cfg.agent.config.observation_cfg
557
+ cfg.suite.task_make_fn.agent_policy_cfg = cfg.agent.config
558
+ continue
559
+ elif "['agent']['config']['policy_cfg']['input_shapes']" in new_key:
560
+ # skip adding the new key if it is the input_shapes of the policy_cfg
561
+ continue
562
+ else:
563
+ print(f"Found new key {new_key} with value {eval(new_key.replace('root[', 'self.cfg['))}. Adding to checkpoint config")
564
+ # eval(new_key.replace('root', 'cfg')) = eval(new_key.replace('root', 'self.cfg'))
565
+ if new_key == "root['agent']['config']['compile']":
566
+ if self.cfg.agent.config.compile:
567
+ self.loading_uncompiled_checkpoint_with_compile = True
568
+
569
+ with open_dict(cfg):
570
+ exec(f"{new_key.replace('root[', 'cfg[')}={new_key.replace('root[', 'self.cfg[')}")
571
+ self.cfg = cfg
572
+
573
+ def load_checkpoint(self, snapshot_path, bc=False):
574
+ print(f'resuming {repr(self.agent)}: {snapshot_path}')
575
+ with snapshot_path.open('rb') as f:
576
+ payload = torch.load(f)
577
+ agent_payload = {}
578
+ for k, v in payload.items():
579
+ if k not in self.__dict__:
580
+ agent_payload[k] = v
581
+ elif k == '_global_epoch':
582
+ self._global_epoch = v
583
+ print(f'loaded epoch: {v}')
584
+ if self.cfg.use_wandb:
585
+ # add to config of wandb
586
+ wandb.config.update({'epoch': v})
587
+
588
+ # self.agent.load_snapshot_eval(agent_payload, bc)
589
+
590
+ @hydra.main(config_path='cfgs', config_name='config_eval')
591
+ def main(cfg):
592
+ from eval_robot import Workspace as W
593
+ root_dir = Path.cwd()
594
+ workspace = W(cfg)
595
+
596
+ workspace.eval()
597
+
598
+ if __name__ == '__main__':
599
+ main()