Files
Etienne Buehrle 7ae01f73a2 AdVIL tests
2021-08-04 16:41:18 +02:00

130 lines
4.6 KiB
Python

import numpy as np
import torch
from torch.utils.data import Dataset, DataLoader
from itertools import chain
from typing import Callable, Union, Type, Optional, Dict, Any
# From https://github.com/DLR-RM/rl-baselines3-zoo/blob/8ea4f4a87afa548832ca17e575b351ec5928c1b0/utils/utils.py
def linear_schedule(initial_value: Union[float, str]) -> Callable[[float], float]:
"""
Linear learning rate schedule.
:param initial_value: (float or str)
:return: (function)
"""
if isinstance(initial_value, str):
initial_value = float(initial_value)
def func(progress_remaining: float) -> float:
"""
Progress will decrease from 1 (beginning) to 0
:param progress_remaining: (float)
:return: (float)
"""
return progress_remaining * initial_value
return func
class SADataset(torch.utils.data.Dataset):
def __init__(self, obs, acts, normalize):
if normalize:
obs = np.array(obs)
self.mean = obs.mean(axis=0)
self.std = obs.std(axis=0) + 1e-3
obs = (obs - self.mean) / (self.std)
self.is_normalized = True
else:
self.is_normalized = False
self.obs = torch.tensor(obs)
self.acts = torch.tensor(acts)
def __len__(self):
return len(self.obs)
def __getitem__(self, idx):
if torch.is_tensor(idx):
idx = idx.tolist()
obs = self.obs[idx]
acts = self.acts[idx]
sample = {'obs': obs, 'acts': acts}
return sample
def make_sa_dataloader(envname, max_trajs=None, normalize=False, batch_size=32):
demos = np.load(
"../experts/{0}/demos.npz".format(envname), allow_pickle=True)
num_trajs = demos["num_trajs"]
if max_trajs is None:
max_trajs = num_trajs
obs = []
acts = []
for traj in range(min(max_trajs, num_trajs)):
obs.extend(demos[str(traj)].item()['states'])
acts.extend(demos[str(traj)].item()['actions'])
dataset = SADataset(obs, acts, normalize)
dataloader = DataLoader(dataset, batch_size=batch_size,
shuffle=True, num_workers=0)
return dataloader
class SADSDataset(torch.utils.data.Dataset):
def __init__(self, obs, acts, next_obs, traj_lens):
self.obs = torch.tensor(obs)
self.acts = torch.tensor(acts)
self.next_obs = torch.tensor(next_obs)
dones = [[False for _ in range(l - 2)] + [True] for l in traj_lens]
self.dones = torch.tensor(list(chain.from_iterable(dones)))
def __len__(self):
return len(self.obs)
def __getitem__(self, idx):
if torch.is_tensor(idx):
idx = idx.tolist()
obs = self.obs[idx]
acts = self.acts[idx]
next_obs = self.next_obs[idx]
dones = self.dones[idx]
sample = {'obs': obs, 'acts': acts,
'next_obs': next_obs, 'dones': dones}
return sample
def make_sads_dataloader(envname, max_trajs=None):
demos = np.load(
"./experts/{0}/demos.npz".format(envname), allow_pickle=True)
num_trajs = demos["num_trajs"]
if max_trajs is None:
max_trajs = num_trajs
obs = []
next_obs = []
acts = []
lens = []
for traj in range(min(max_trajs, num_trajs)):
obs.extend(demos[str(traj)].item()['states'][:-1])
next_obs.extend(demos[str(traj)].item()['states'][1:])
acts.extend(demos[str(traj)].item()['actions'][:-1])
lens.append(len(demos[str(traj)].item()['states']))
dataset = SADSDataset(obs, acts, next_obs, lens)
dataloader = DataLoader(dataset, batch_size=32,
shuffle=False, num_workers=0, drop_last=True)
return dataloader
def make_sa_dataset(envname, max_trajs=None):
demos = np.load("../pillbox/experts/{0}/demos.npz".format(envname), allow_pickle=True)
num_trajs = demos["num_trajs"]
if max_trajs is None:
max_trajs = num_trajs
expert_states = []
expert_actions = []
expert_next_states = []
expert_dones = []
for traj in range(min(max_trajs, num_trajs)):
expert_states.extend(demos[str(traj)].item()['states'][:-1])
expert_next_states.extend(demos[str(traj)].item()['states'][1:])
expert_actions.extend(demos[str(traj)].item()['actions'][:-1])
l = len(demos[str(traj)].item()['states'])
expert_dones.extend([False for _ in range(l - 2)] + [True])
expert_data = dict()
expert_data['obs'] = np.array(expert_states)
expert_data['acts'] = np.array(expert_actions)
expert_data['next_obs'] = np.array(expert_next_states)
expert_data['dones'] = np.array(expert_dones)
return expert_data