Commit c0a7489f authored by Ibrahim's avatar Ibrahim
Browse files

commonml.rl: Memory flush() truncates episodes instead of all memory

parent b702383c
Loading
Loading
Loading
Loading
+9 −0
Original line number Diff line number Diff line
from .ppo import (
    PPO,
    Policy,
    ActorCriticBox,
    ActorCriticDiscrete,
    ActorCriticMultiBinary,
    Memory,
    DEVICE
)
 No newline at end of file
+3 −0
Original line number Diff line number Diff line
"""
Hyperparameter search for RL algorithms.
"""
 No newline at end of file

commonml/rl/plot.py

0 → 100644
+3 −0
Original line number Diff line number Diff line
"""
Plotting functions for RL.
"""
 No newline at end of file
+92 −26
Original line number Diff line number Diff line
@@ -25,19 +25,38 @@ class Memory:
        self.states = []
        self.logprobs = []
        self.rewards = []
        self.returns = []
        self.is_terminals = []


    def add(self, state, action, logprob, reward, done):
        self.states.append(state)
        self.actions.append(action)
        self.logprobs.append(logprob)
        self.rewards.append(reward)
        self.is_terminals.append(done)


    def clear(self):
        del self.actions[:]
        del self.states[:]
        del self.logprobs[:]
        del self.rewards[:]
        del self.returns[:]
        del self.is_terminals[:]


    def flush(self):
        size = len(self.states)
        for i, done in enumerate(reversed(self.is_terminals)):
            if done:
                truncate = size - i
                del self.actions[:truncate]
                del self.states[:truncate]
                del self.logprobs[:truncate]
                del self.rewards[:truncate]
                del self.is_terminals[:truncate]
                break



class Policy(nn.Module):

@@ -45,11 +64,12 @@ class Policy(nn.Module):
    dist_kwargs = {}


    def __init__(self, state_dim, action_dim, n_latent_var):
    def __init__(self, state_dim, action_dim, n_latent_var, device=DEVICE):
        super().__init__()
        self.state_dim = state_dim
        self.action_dim = action_dim
        self.n_latent_var = n_latent_var
        self.device = device

        self.base = None

@@ -70,7 +90,7 @@ class Policy(nn.Module):

    def predict(self, state: np.ndarray):
        # pylint: disable=no-member
        state = torch.from_numpy(state).float().to(DEVICE)
        state = torch.from_numpy(state).float().to(self.device)
        action_probs = self.action_layer(state)  # Discrete
        dist = self.dist(action_probs, **self.dist_kwargs)
        action = dist.sample()
@@ -95,8 +115,8 @@ class ActorCriticDiscrete(Policy):
    dist = Categorical
    dist_kwargs = {}

    def __init__(self, state_dim, action_dim, n_latent_var):
        super().__init__(state_dim=state_dim, action_dim=action_dim, n_latent_var=n_latent_var)
    def __init__(self, state_dim, action_dim, n_latent_var, device=DEVICE):
        super().__init__(state_dim=state_dim, action_dim=action_dim, n_latent_var=n_latent_var, device=device)

        self.action_layer = nn.Sequential(
                nn.Linear(state_dim, n_latent_var),
@@ -119,8 +139,8 @@ class ActorCriticMultiBinary(Policy):
    dist_kwargs = {}


    def __init__(self, state_dim, action_dim, n_latent_var):
        super().__init__(state_dim, action_dim, n_latent_var)
    def __init__(self, state_dim, action_dim, n_latent_var, device=DEVICE):
        super().__init__(state_dim, action_dim, n_latent_var, device=device)

        self.action_layer = nn.Sequential(
                nn.Linear(state_dim, n_latent_var),
@@ -151,8 +171,8 @@ class ActorCriticBox(Policy):


    def __init__(self, state_dim, action_dim, n_latent_var, action_std=0.01,
            activation=nn.Tanh):
        super().__init__(state_dim=state_dim, action_dim=action_dim, n_latent_var=n_latent_var)
            activation=nn.Tanh, device=DEVICE):
        super().__init__(state_dim=state_dim, action_dim=action_dim, n_latent_var=n_latent_var, device=device)
        self.activation = activation
        self.action_layer = nn.Sequential(
                nn.Linear(state_dim, n_latent_var),
@@ -160,9 +180,9 @@ class ActorCriticBox(Policy):
                nn.Linear(n_latent_var, n_latent_var),
                nn.Tanh(),
                nn.Linear(n_latent_var, action_dim),
                self.activation()
                nn.Identity() if self.activation is None else self.activation()
                )
        self.dist_kwargs = dict(covariance_matrix=(torch.eye(action_dim) * action_std).to(DEVICE))
        self.dist_kwargs = dict(covariance_matrix=(torch.eye(action_dim) * action_std).to(self.device))


    def predict(self, state):
@@ -177,40 +197,48 @@ class PPO:

    def __init__(self, env, policy, state_dim, action_dim, n_latent_var=64, lr=0.02,
                 betas=(0.9, 0.999), gamma=0.99, epochs=5, eps_clip=0.2,
                 update_interval=2000, seed=None, summary: SummaryWriter=None,
                 **policy_kwargs):
                 truncate=False, update_interval=2000, seed=None,
                 device=DEVICE, summary: SummaryWriter=None, **policy_kwargs):
        self.random = np.random.RandomState(seed)
        self.seed = seed
        if seed is not None: torch.manual_seed(self.seed)
        self.env = env
        self.lr = lr
        self.betas = betas
        self.gamma = gamma
        self.eps_clip = eps_clip
        self.truncate = truncate
        self.epochs = epochs
        self.update_interval = update_interval
        self.device = device
        
        self.policy = policy(state_dim, action_dim, n_latent_var, **policy_kwargs).to(DEVICE)
        policy_kwargs['device'] = device
        self.policy = policy(state_dim, action_dim, n_latent_var, **policy_kwargs).to(self.device)
        self.optimizer = torch.optim.Adam(self.policy.parameters(), lr=lr, betas=betas)
        
        self.MseLoss = nn.MSELoss()

        self.random = np.random.RandomState(seed)
        self.seed = self.random.rand() if seed is None else seed
        self.summary = summary
        self.meta_policy = None
        if seed is not None: torch.manual_seed(self.seed)

    

    def update(self, policy, memory, epochs: int=1, optimizer=None, summary=None,
        grad_callback=None):

        rewards = returns(memory.rewards, memory.is_terminals, self.gamma)
        rewards = returns(memory.rewards, memory.is_terminals, self.gamma, truncate=self.truncate)
        truncate = len(rewards)
        # If the returns calculated are zero length, i.e. when memory does not
        # contain a single full episode, because returns() truncated incompleted
        # episodes, abort update:
        if truncate == 0:
            return np.nan
        # Casting to correct data type and DEVICE
        # pylint: disable=not-callable
        rewards = torch.tensor(rewards[:truncate]).float().to(DEVICE)
        old_states = torch.tensor(memory.states[:truncate]).float().to(DEVICE).detach()
        old_actions = torch.tensor(memory.actions[:truncate]).float().to(DEVICE).detach()
        old_logprobs = torch.tensor(memory.logprobs[:truncate]).float().to(DEVICE).detach()
        rewards = torch.tensor(rewards[:truncate]).float().to(self.device)
        old_states = torch.tensor(memory.states[:truncate]).float().to(self.device).detach()
        old_actions = torch.tensor(memory.actions[:truncate]).float().to(self.device).detach()
        old_logprobs = torch.tensor(memory.logprobs[:truncate]).float().to(self.device).detach()

        # If states/actions are 1D arrays of single number states/actions,
        # convert them to 2D matrix of 1 column where each row is one timestep.
@@ -282,13 +310,44 @@ class PPO:
    def learn(self, timesteps, update_interval=None, track_higher_grads=False,
              lr_scheduler=None,
              step_callback=None, interval_callback=None,
              reward_aggregation='episodic'):
              reward_aggregation='episodic') -> List[float]:
        """
        Run learning.

        Parameters
        ----------
        timesteps : int
            Number of steps to interact with environment.
        update_interval : int, optional
            Number of steps after which to update policy, by default None
        track_higher_grads : bool, optional
            Whether to track rate of change of parameters w.r.t themselves, by default False
        lr_scheduler : [type], optional
            A scheduler to adjust learning rate, by default None
        step_callback : Callable, optional
            A function to call after every step. It is passed a dictionary of
            local variables, by default None
        interval_callback : Callable, optional
            A function to call after every update interval. It is passed a 
            dictionary of local variables, by default None
        reward_aggregation : str, optional
            One of 'episodic', 'episodic.normalized', 'interval'.
            'episodic' returns total rewards per episode. If normalized
            returns reward divided by episode length. 'interval' returns
            rewards per update interval., by default 'episodic'

        Returns
        -------
        List[float]
            A list of aggregated rewards.
        """
        if update_interval is None:
            update_interval = self.update_interval
        state = self.env.reset()
        memory = Memory()
        episodic_rewards = [0.]
        interval_rewards = [0.]
        t_episode = 0
        # This context wraps the policy and optimizer to track parameter updates
        # over time such that d Params(time=t) / d Params(time=t-n) can be calculated.
        # If not tracking higher gradients, a dummy context is used which does
@@ -303,9 +362,13 @@ class PPO:
                state, reward, done, info = self.env.step(action)
                episodic_rewards[-1] += reward
                interval_rewards[-1] += reward
                t_episode += 1
                if done:
                    state = self.env.reset()
                    if reward_aggregation.endswith('normalized'):
                        episodic_rewards[-1] /= t_episode
                    episodic_rewards.append(0.)
                    t_episode = 0
                memory.actions.append(action)
                memory.logprobs.append(logprob)
                memory.rewards.append(reward)
@@ -323,13 +386,13 @@ class PPO:
                        interval_callback(locals())
                    if lr_scheduler is not None:
                        lr_scheduler()
                    memory.clear()
                    memory.flush()

        
            self.meta_policy = policy if track_higher_grads else None
        self.policy.load_state_dict(policy.state_dict())

        if reward_aggregation == 'episodic':
        if reward_aggregation.startswith('episodic'):
            return episodic_rewards[:-1 if len(episodic_rewards) > 1 else None]
        else:
            return interval_rewards
@@ -348,10 +411,13 @@ def returns(rewards, is_terminals, gamma, truncate=False):
    # where they should be 5,4,3
    if truncate:
        if True in is_terminals:
            # TODO: optimize by using reversed() instead of copying array using slicing[::-1]
            idx_from_end = is_terminals[::-1].index(True)
            if idx_from_end > 0:
                rewards = rewards[:-idx_from_end]
                is_terminals = is_terminals[:-idx_from_end]
        else:
            is_terminals = []
    returns = []
    discounted_reward = 0
    for reward, is_terminal in zip(reversed(rewards), reversed(is_terminals)):