From f6d0c4d2609d5efb55e220a51b99160aa3ab0874 Mon Sep 17 00:00:00 2001 From: Shangtong Zhang Date: Thu, 17 May 2018 23:37:50 -0600 Subject: [PATCH] Remove unused wrapper --- deep_rl/network/network_utils.py | 89 -------------------------------- 1 file changed, 89 deletions(-) diff --git a/deep_rl/network/network_utils.py b/deep_rl/network/network_utils.py index daadfe8..a0d303c 100644 --- a/deep_rl/network/network_utils.py +++ b/deep_rl/network/network_utils.py @@ -23,95 +23,6 @@ class BaseNet: x = torch.tensor(x, device=self.device, dtype=torch.float32) return x -# class DisjointActorCriticWrapper: -# def __init__(self, state_dim, action_dim, actor_network_fn, critic_network_fn): -# self.actor = actor_network_fn(state_dim, action_dim) -# self.critic = critic_network_fn(state_dim, action_dim) -# -# def state_dict(self): -# return [self.actor.state_dict(), self.critic.state_dict()] -# -# def load_state_dict(self, state_dicts): -# self.actor.load_state_dict(state_dicts[0]) -# self.critic.load_state_dict(state_dicts[1]) -# -# def parameters(self): -# return list(self.actor.parameters()) + list(self.critic.parameters()) -# -# def zero_grad(self): -# self.actor.zero_grad() -# self.critic.zero_grad() -# -# class GaussianActorCriticWrapper: -# def __init__(self, state_dim, action_dim, actor_fn, critic_fn, actor_opt_fn, critic_opt_fn): -# self.actor = actor_fn(state_dim, action_dim) -# self.critic = critic_fn(state_dim) -# self.actor_opt = actor_opt_fn(self.actor.parameters()) -# self.critic_opt = critic_opt_fn(self.critic.parameters()) -# -# def predict(self, state, actions=None): -# mean, std, log_std = self.actor.predict(state) -# values = self.critic.predict(state) -# dist = torch.distributions.Normal(mean, std) -# if actions is None: -# actions = dist.sample() -# log_probs = dist.log_prob(actions) -# log_probs = torch.sum(log_probs, dim=1, keepdim=True) -# return actions, log_probs, 0, values -# -# def tensor(self, x): -# return self.actor.tensor(x) -# -# def zero_grad(self): -# self.actor_opt.zero_grad() -# self.critic_opt.zero_grad() -# -# def parameters(self): -# return list(self.actor.parameters()) + list(self.critic.parameters()) -# -# def step(self): -# self.actor_opt.step() -# self.critic_opt.step() -# -# def state_dict(self): -# return [self.actor.state_dict(), self.critic.state_dict()] -# -# def load_state_dict(self, state_dicts): -# self.actor.load_state_dict(state_dicts[0]) -# self.critic.load_state_dict(state_dicts[1]) -# -# class CategoricalActorCriticWrapper: -# def __init__(self, state_dim, action_dim, network_fn, opt_fn): -# self.network = network_fn(state_dim, action_dim) -# self.opt = opt_fn(self.network.parameters()) -# -# def predict(self, state, action=None): -# prob, log_prob, value = self.network.predict(state) -# entropy_loss = torch.sum(prob * log_prob, dim=1, keepdim=True) -# dist = torch.distributions.Categorical(prob) -# if action is None: -# action = dist.sample() -# log_prob = dist.log_prob(action).unsqueeze(1) -# return action, log_prob, entropy_loss.mean(0), value -# -# def tensor(self, x): -# return self.network.tensor(x) -# -# def zero_grad(self): -# self.opt.zero_grad() -# -# def parameters(self): -# return self.network.parameters() -# -# def step(self): -# self.opt.step() -# -# def state_dict(self): -# return self.network.state_dict() -# -# def load_state_dict(self, state_dicts): -# self.network.load_state_dict(state_dicts) - def layer_init(layer, w_scale=1.0): nn.init.orthogonal_(layer.weight.data) layer.weight.data.mul_(w_scale)