| import os |
| import torch |
| import torch.nn as nn |
| import numpy as np |
| import copy |
| import random |
| import time |
| from termcolor import cprint |
|
|
| from diffusion_policy.common.replay_buffer import ReplayBuffer |
| from diffusion_policy.dataset.base_dataset import BaseLowdimDataset |
| from diffusion_policy.common.pref_replay_buffer import PrefReplayBuffer |
| from diffusion_policy.common.pref_sampler import PrefSequenceSampler |
| from diffusion_policy.preference_labeling.preference_labeling import ( |
| load_or_create_indices, |
| load_or_compute_feats, |
| precompute_pair_rewards, |
| get_context_observations, |
| extract_segment_pseudo_reward, |
| ResNet, R3M, LIV, VIP |
| ) |
| from typing import Dict |
|
|
|
|
| class PbrlLowdimDataset(BaseLowdimDataset): |
| def __init__( |
| self, |
| replay_buffer_1: ReplayBuffer, |
| replay_buffer_2: ReplayBuffer, |
| abs_action=True, |
| sequence_length=1, |
| gamma=0.999, |
| num_queries=1, |
| seed=42, |
| gpu_device='cuda:0', |
| dense_reward=False, |
| val_ratio_data1=None, |
| val_ratio_data2=None, |
| dataset_1_path=None, |
| dataset_2_path=None, |
| task_name="can_lowdim", |
| pseudo_preference=False, |
| replay_buffer_expert=None, |
| dataset_expert_path=None, |
| feature_extractor="r3m_resnet18", |
| context_num=3, |
| seg_margin=0.6, |
| min_progress=0.0, |
| n_demos_for_preference=10, |
| ): |
| super().__init__() |
| assert abs_action is True, "Only absolute action is supported" |
| assert feature_extractor in ["imagenet_resnet18", "r3m_resnet18", "liv_resnet50", "vip_resnet50"] |
| self.pseudo_preference = pseudo_preference |
|
|
| |
| self.seg_margin = seg_margin |
| self.min_progress = min_progress |
| self.context_num = context_num |
| self.n_demos_for_preference = n_demos_for_preference |
|
|
| episode_ends_1 = replay_buffer_1.episode_ends |
| episode_ends_2 = replay_buffer_2.episode_ends |
| num_episodes_1 = int(len(episode_ends_1) * (1 - val_ratio_data1)) |
| num_episodes_2 = int(len(episode_ends_2) * (1 - val_ratio_data2)) |
| self.pref_replay_buffer = PrefReplayBuffer.create_empty_numpy() |
| random.seed(seed) |
| print(f"=====================> PbrlLowdimDataset: Num episodes (dataset_1): {num_episodes_1}, " |
| f"min_len={replay_buffer_1.episode_lengths.min()}, max_len={replay_buffer_1.episode_lengths.max()}") |
| print(f"=====================> PbrlLowdimDataset: Num episodes (dataset_2): {num_episodes_2}," |
| f"min_len={replay_buffer_2.episode_lengths.min()}, max_len={replay_buffer_2.episode_lengths.max()}") |
|
|
| |
| idx_path = f"logs/pbrl_indices/{task_name}/pair_{task_name}_nQ{num_queries}_L{sequence_length}_{num_episodes_1}_{num_episodes_2}.npz" |
| pair_indices = load_or_create_indices( |
| path=idx_path, |
| num_queries=num_queries, |
| num_episodes_1=num_episodes_1, |
| num_episodes_2=num_episodes_2, |
| episode_ends_1=episode_ends_1, |
| episode_ends_2=episode_ends_2, |
| sequence_length=sequence_length, |
| seed=seed, |
| use_cached=False, |
| save_cached=True |
| ) |
|
|
| if self.pseudo_preference: |
| if feature_extractor == "imagenet_resnet18": |
| encoder = ResNet().to(gpu_device).eval() |
| elif feature_extractor == "r3m_resnet18": |
| encoder = R3M().to(gpu_device).eval() |
| elif feature_extractor == "liv_resnet50": |
| encoder = LIV().to(gpu_device).eval() |
| elif feature_extractor == "vip_resnet50": |
| encoder = VIP().to(gpu_device).eval() |
| else: |
| raise ValueError(f"Unknown feature extractor: {feature_extractor}") |
|
|
| base_1 = os.path.join(os.path.dirname(dataset_1_path), "videos") |
| base_2 = os.path.join(os.path.dirname(dataset_2_path), "videos") |
| video_paths_1 = [f"{base_1}/episode_{i}.mp4" for i in range(num_episodes_1)] |
| video_paths_2 = [f"{base_2}/episode_{i}.mp4" for i in range(num_episodes_2)] |
|
|
| |
| os.makedirs("cache", exist_ok=True) |
| device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu") |
|
|
| start = time.time() |
| feats_1 = load_or_compute_feats( |
| f"cache/dataset_1_{task_name.replace('_lowdim', '')}_{feature_extractor}.npz", |
| video_paths_1, encoder, device, drop_last='datacollect_diffusion_transformer' in video_paths_1, |
| use_cached=True, save_cached=True) |
| feats_2 = load_or_compute_feats( |
| f"cache/dataset_2_{task_name.replace('_lowdim', '')}_{feature_extractor}.npz", |
| video_paths_2, encoder, device, drop_last='datacollect_diffusion_transformer' in video_paths_2, |
| use_cached=True, save_cached=True) |
| print(f"Total time to load/encode {len(video_paths_1) + len(video_paths_2)} videos: {time.time() - start:.2f}s") |
|
|
| if replay_buffer_expert is None: |
| replay_buffer_expert = replay_buffer_1 |
| dataset_expert_path = dataset_1_path |
|
|
| top_k = np.argpartition(replay_buffer_expert.episode_lengths, self.n_demos_for_preference)[:self.n_demos_for_preference] |
| base_3 = os.path.join(os.path.dirname(dataset_expert_path), "videos") |
| expert_paths = [f"{base_3}/episode_{i}.mp4" for i in top_k] |
| expert_feats = load_or_compute_feats(f"cache/experts_{task_name}_nD{self.n_demos_for_preference}_{feature_extractor}.npz", expert_paths, encoder, device, use_cached=False, save_cached=False) |
| expert_ctx = [get_context_observations(f, context_num=self.context_num) for f in expert_feats] |
|
|
| |
| print("precomputing trajectory's rewards for dataset 1...") |
| dataset_rewards_1 = precompute_pair_rewards(feats_1, expert_ctx, context_num=self.context_num) |
| print("precomputing trajectory's rewards for dataset 2...") |
| dataset_rewards_2 = precompute_pair_rewards(feats_2, expert_ctx, context_num=self.context_num) |
|
|
| assert len(dataset_rewards_1) == num_episodes_1 and len(dataset_rewards_2) == num_episodes_2 |
|
|
| |
| orca_match = 0 |
| retained_pairs = 0 |
| for i in range(num_queries): |
| ep_idx_1, ts_idx_1, ep_idx_2, ts_idx_2 = pair_indices[i] |
| ep_idx_1, ts_idx_1, ep_idx_2, ts_idx_2 = int(ep_idx_1), int(ts_idx_1), int(ep_idx_2), int(ts_idx_2) |
|
|
| episode_1 = replay_buffer_1.get_episode(ep_idx_1, keys=['obs', 'action', 'reward'], copy=False) |
| episode_2 = replay_buffer_2.get_episode(ep_idx_2, keys=['obs', 'action', 'reward'], copy=False) |
|
|
| |
| episode_1_len = len(episode_1['obs']) |
| if episode_1_len >= sequence_length: |
| start_1 = ts_idx_1 |
| length = sequence_length |
| for key in episode_1.keys(): |
| episode_1[key] = episode_1[key][start_1:start_1 + sequence_length] |
| else: |
| length = episode_1_len |
| for key in episode_1.keys(): |
| episode_1[key] = np.pad(episode_1[key], ((0, sequence_length - episode_1_len),) + ((0, 0),) * (episode_1[key].ndim - 1), mode='edge') |
|
|
| |
| episode_2_len = len(episode_2['obs']) |
| if episode_2_len >= sequence_length: |
| start_2 = ts_idx_2 |
| length_2 = sequence_length |
| for key in episode_2.keys(): |
| episode_2[key] = episode_2[key][start_2:start_2 + sequence_length] |
| else: |
| length_2 = episode_2_len |
| for key in episode_2.keys(): |
| episode_2[key] = np.pad(episode_2[key], ((0, sequence_length - episode_2_len),) + ((0, 0),) * (episode_2[key].ndim - 1), mode='edge') |
|
|
| |
| votes = np.sum([(gamma ** t) * reward for t, reward in enumerate(episode_1['reward'])]) |
| votes_2 = np.sum([(gamma ** t) * reward for t, reward in enumerate(episode_2['reward'])]) |
|
|
| if self.pseudo_preference: |
| gt = 1 if votes_2 > votes else 0 |
|
|
| orca_1_scores, orca_1_scores_all = extract_segment_pseudo_reward(dataset_rewards_1[ep_idx_1], ts_idx_1, sequence_length) |
| orca_2_scores, orca_2_scores_all = extract_segment_pseudo_reward(dataset_rewards_2[ep_idx_2], ts_idx_2, sequence_length) |
| |
| score_1 = orca_1_scores.max() |
| score_2 = orca_2_scores.max() |
|
|
| |
| if max(score_1, score_2) < self.min_progress: |
| pref_label = -1 |
| elif score_2 - score_1 > self.seg_margin: |
| pref_label = 1 |
| elif score_1 - score_2 > self.seg_margin: |
| pref_label = 0 |
| else: |
| pref_label = -1 |
|
|
| if pref_label != -1: |
| retained_pairs += 1 |
| orca_match += (pref_label == gt) |
|
|
| votes, votes_2 = score_1, score_2 |
| |
| self.pref_replay_buffer.add_pref_episode( |
| data={ |
| 'obs': episode_1['obs'], |
| 'action': episode_1['action'], |
| 'obs_2': episode_2['obs'], |
| 'action_2': episode_2['action'], |
| }, |
| meta_data={ |
| 'votes': votes, |
| 'votes_2': votes_2, |
| 'length': np.array([length]), |
| 'length_2': np.array([length_2]), |
| 'beta_priori': np.ones([2]), |
| 'beta_priori_2': np.ones([2]), |
| } |
| ) |
|
|
| |
| retention_rate = (retained_pairs / num_queries) * 100 |
| accuracy = (orca_match / retained_pairs * 100) if retained_pairs > 0 else 0.0 |
| |
| else: |
| |
| self.pref_replay_buffer.add_pref_episode( |
| data={ |
| 'obs': episode_1['obs'], |
| 'action': episode_1['action'], |
| 'obs_2': episode_2['obs'], |
| 'action_2': episode_2['action'], |
| }, |
| meta_data={ |
| 'votes': votes, |
| 'votes_2': votes_2, |
| 'length': np.array([length]), |
| 'length_2': np.array([length_2]), |
| 'beta_priori': np.ones([2]), |
| 'beta_priori_2': np.ones([2]), |
| } |
| ) |
|
|
| if self.pseudo_preference: |
| assert retained_pairs > 0, f"Margin ({self.seg_margin}) is too strict! 0 pairs retained out of {num_queries}." |
| cprint(f"Task={task_name.upper()}: n_expert={self.n_demos_for_preference}, n_queries={num_queries}, seq_len={sequence_length}, feat={feature_extractor}, min_progress={min_progress}, margin={seg_margin}", "green", attrs=["bold"]) |
| cprint(f" -> Pairs Retained: {retained_pairs} ({retention_rate:.1f}%)", "cyan") |
| cprint(f" -> ORCA Accuracy (on retained): {accuracy:.1f}%", "green", attrs=["bold"]) |
| train_mask = np.ones(retained_pairs, dtype=bool) |
| self.retained_pairs = retained_pairs |
| self.accuracy = accuracy |
| self.retention_rate = retention_rate |
| else: |
| train_mask = np.ones(num_queries, dtype=bool) |
| self.accuracy = 0.0 |
| self.retained_pairs = num_queries |
| self.retention_rate = 100 |
|
|
| self.sampler = PrefSequenceSampler( |
| replay_buffer=self.pref_replay_buffer, |
| sequence_length=sequence_length, |
| episode_mask=train_mask, |
| ) |
|
|
| self.gpu_device = gpu_device |
| self.train_mask = train_mask |
| self.sequence_length = sequence_length |
| self.dense_reward = dense_reward |
|
|
| def construct_pref_data(self): |
| data = self.pref_replay_buffer.data |
| pref_data = data.copy() |
| meta = self.pref_replay_buffer.meta |
| pref_data.update(meta) |
| if 'episode_ends' in pref_data.keys(): |
| del pref_data['episode_ends'] |
|
|
| return pref_data |
|
|
| def get_validation_dataset(self): |
| val_set = copy.copy(self) |
| val_set.sampler = PrefSequenceSampler( |
| replay_buffer=self.pref_replay_buffer, |
| sequence_length=self.sequence_length, |
| episode_mask=~self.train_mask, |
| ) |
| val_set.train_mask = ~self.train_mask |
| return val_set |
|
|
| def get_all_actions(self) -> torch.Tensor: |
| actions = np.concatenate(self.pref_replay_buffer.data['action'], self.pref_replay_buffer.data['action_2'], dim = 0) |
| return torch.from_numpy(actions) |
|
|
| def __len__(self) -> int: |
| return self.sampler.__len__() |
|
|
| def __getitem__(self, idx: int) -> Dict[str, torch.Tensor]: |
| torch_data = self.sampler.sample_sequence(idx) |
| return torch_data |