From 51f125c01eb013b77ebf3a78795ae87dc9ffb9f9 Mon Sep 17 00:00:00 2001 From: Shangtong Zhang Date: Sat, 7 Oct 2017 21:23:02 -0600 Subject: [PATCH] Improve normalizer for scalar --- async_worker/continuous_actor_critic.py | 2 +- async_worker/ppo.py | 2 +- utils/normalizer.py | 13 +++++++++++-- 3 files changed, 13 insertions(+), 4 deletions(-) diff --git a/async_worker/continuous_actor_critic.py b/async_worker/continuous_actor_critic.py index 16a4152..dabb0dd 100644 --- a/async_worker/continuous_actor_critic.py +++ b/async_worker/continuous_actor_critic.py @@ -48,7 +48,7 @@ class ContinuousAdvantageActorCritic: steps += 1 total_reward += reward - reward = np.asscalar(self.reward_normalizer(np.array([reward]))) + reward = self.reward_normalizer(reward) if deterministic: if terminal: diff --git a/async_worker/ppo.py b/async_worker/ppo.py index 6adf27d..a6154b6 100644 --- a/async_worker/ppo.py +++ b/async_worker/ppo.py @@ -79,7 +79,7 @@ class ProximalPolicyOptimization: batched_steps += 1 episode_length += 1 - reward = np.asscalar(self.reward_normalizer(np.array([reward]))) + reward = self.reward_normalizer(reward) rewards.append(reward) if done: diff --git a/utils/normalizer.py b/utils/normalizer.py index b32ddd2..f7171cb 100644 --- a/utils/normalizer.py +++ b/utils/normalizer.py @@ -4,6 +4,7 @@ # declaration at the top # ####################################################################### import torch +import numpy as np class Normalizer: def __init__(self, o_size): @@ -22,13 +23,21 @@ class StaticNormalizer: self.online_stats = SharedStats(o_size) def __call__(self, o_): - o = torch.FloatTensor(o_) + if np.isscalar(o_): + o = torch.FloatTensor([o_]) + else: + o = torch.FloatTensor(o_) self.online_stats.feed(o) if self.offline_stats.n[0] == 0: return o_ std = (self.offline_stats.v + 1e-6) ** .5 o = (o - self.offline_stats.m) / std - return o.numpy().reshape(o_.shape) + o = o.numpy() + if np.isscalar(o_): + o = np.asscalar(o) + else: + o = o.reshape(o_.shape) + return o class SharedStats: def __init__(self, o_size):