Files
InteractionImitation/scratch/etienne/pillbox/learners/sqil.py
Etienne Buehrle 7ae01f73a2 AdVIL tests
2021-08-04 16:41:18 +02:00

61 lines
2.7 KiB
Python

import warnings
from abc import ABC, abstractmethod
from typing import Dict, Generator, Optional, Union
import numpy as np
import torch as th
from gym import spaces
try:
# Check memory used by replay buffer when possible
import psutil
except ImportError:
psutil = None
from stable_baselines3.common.preprocessing import get_action_dim, get_obs_shape
from stable_baselines3.common.type_aliases import ReplayBufferSamples, RolloutBufferSamples
from stable_baselines3.common.vec_env import VecNormalize
from stable_baselines3.common.buffers import ReplayBuffer
class SQILReplayBuffer(ReplayBuffer):
def __init__(
self,
buffer_size: int,
observation_space: spaces.Space,
action_space: spaces.Space,
device: Union[th.device, str] = "cpu",
n_envs: int = 1,
optimize_memory_usage: bool = False,
expert_data: dict = dict(),
):
super(SQILReplayBuffer, self).__init__(buffer_size, observation_space, action_space, device, n_envs=n_envs, optimize_memory_usage=optimize_memory_usage)
self.expert_states = expert_data['obs']
self.expert_actions = expert_data['acts']
self.expert_next_states = expert_data['next_obs']
self.expert_dones = expert_data['dones']
def _get_samples(self, batch_inds: np.ndarray, env: Optional[VecNormalize] = None) -> ReplayBufferSamples:
num_samples = len(batch_inds)
num_expert_samples = int(num_samples / 2)
batch_inds = batch_inds[:num_expert_samples]
expert_inds = np.random.randint(0, len(self.expert_states), size=num_expert_samples)
# Balanced sampling
if self.optimize_memory_usage:
next_obs = self._normalize_obs(self.observations[(batch_inds + 1) % self.buffer_size, 0, :], env)
else:
next_obs = self._normalize_obs(self.next_observations[batch_inds, 0, :], env)
next_obs = np.concatenate((next_obs, self._normalize_obs(self.expert_next_states[expert_inds], env)), axis=0)
obs = self._normalize_obs(self.observations[batch_inds, 0, :], env)
obs = np.concatenate((obs, self._normalize_obs(self.expert_states[expert_inds], env)), axis=0)
actions = self.actions[batch_inds, 0, :]
actions = np.concatenate((actions, self.expert_actions[expert_inds].reshape(num_expert_samples, -1)), axis=0)
dones = self.dones[batch_inds]
dones = np.concatenate((dones, self.expert_dones[expert_inds].reshape(num_expert_samples, -1)), axis=0)
# SQIL Rewards
rewards = self.rewards[batch_inds] * 0.
rewards = np.concatenate((rewards, np.ones_like(rewards)), axis=0)
data = (obs, actions, next_obs, dones, rewards)
return ReplayBufferSamples(*tuple(map(self.to_torch, data)))