mirror of
https://github.com/saymrwulf/stable-baselines3.git
synced 2026-09-15 22:10:25 +00:00
Attempt to fix loss of perf because of VecEnvs
This commit is contained in:
parent
0e727a5f72
commit
a9b8276efb
2 changed files with 11 additions and 11 deletions
|
|
@ -7,7 +7,7 @@ from torchy_baselines import TD3, CEMRL, PPO
|
||||||
def test_td3():
|
def test_td3():
|
||||||
model = TD3('MlpPolicy', 'Pendulum-v0', policy_kwargs=dict(net_arch=[64, 64]),
|
model = TD3('MlpPolicy', 'Pendulum-v0', policy_kwargs=dict(net_arch=[64, 64]),
|
||||||
start_timesteps=100, verbose=1, create_eval_env=True)
|
start_timesteps=100, verbose=1, create_eval_env=True)
|
||||||
model.learn(total_timesteps=500, eval_freq=100)
|
model.learn(total_timesteps=20000, eval_freq=1000)
|
||||||
model.save("test_save")
|
model.save("test_save")
|
||||||
model.load("test_save")
|
model.load("test_save")
|
||||||
os.remove("test_save.pth")
|
os.remove("test_save.pth")
|
||||||
|
|
@ -15,7 +15,7 @@ def test_td3():
|
||||||
def test_cemrl():
|
def test_cemrl():
|
||||||
model = CEMRL('MlpPolicy', 'Pendulum-v0', policy_kwargs=dict(net_arch=[16]), pop_size=2, n_grad=1,
|
model = CEMRL('MlpPolicy', 'Pendulum-v0', policy_kwargs=dict(net_arch=[16]), pop_size=2, n_grad=1,
|
||||||
start_timesteps=100, verbose=1, create_eval_env=True)
|
start_timesteps=100, verbose=1, create_eval_env=True)
|
||||||
model.learn(total_timesteps=1000, eval_freq=500)
|
model.learn(total_timesteps=20000, eval_freq=1000)
|
||||||
model.save("test_save")
|
model.save("test_save")
|
||||||
model.load("test_save")
|
model.load("test_save")
|
||||||
os.remove("test_save.pth")
|
os.remove("test_save.pth")
|
||||||
|
|
|
||||||
|
|
@ -68,11 +68,11 @@ class ReplayBuffer(BaseBuffer):
|
||||||
|
|
||||||
def add(self, state, next_state, action, reward, done):
|
def add(self, state, next_state, action, reward, done):
|
||||||
# Copy to avoid modification by reference
|
# Copy to avoid modification by reference
|
||||||
self.states[self.pos] = th.FloatTensor(np.array(state))
|
self.states[self.pos] = th.FloatTensor(np.array(state).copy())
|
||||||
self.next_states[self.pos] = th.FloatTensor(np.array(next_state))
|
self.next_states[self.pos] = th.FloatTensor(np.array(next_state).copy())
|
||||||
self.actions[self.pos] = th.FloatTensor(np.array(action))
|
self.actions[self.pos] = th.FloatTensor(np.array(action).copy())
|
||||||
self.rewards[self.pos] = th.FloatTensor(np.array(reward))
|
self.rewards[self.pos] = th.FloatTensor(np.array(reward).copy())
|
||||||
self.dones[self.pos] = th.FloatTensor(np.array(done))
|
self.dones[self.pos] = th.FloatTensor(np.array(done).copy())
|
||||||
|
|
||||||
self.pos += 1
|
self.pos += 1
|
||||||
if self.pos == self.buffer_size:
|
if self.pos == self.buffer_size:
|
||||||
|
|
@ -131,10 +131,10 @@ class RolloutBuffer(BaseBuffer):
|
||||||
def add(self, state, action, reward, done, value, log_prob):
|
def add(self, state, action, reward, done, value, log_prob):
|
||||||
self.values[self.pos] = th.FloatTensor(value.clone().cpu().flatten())
|
self.values[self.pos] = th.FloatTensor(value.clone().cpu().flatten())
|
||||||
self.log_probs[self.pos] = th.FloatTensor(log_prob.cpu().clone())
|
self.log_probs[self.pos] = th.FloatTensor(log_prob.cpu().clone())
|
||||||
self.states[self.pos] = th.FloatTensor(np.array(state))
|
self.states[self.pos] = th.FloatTensor(np.array(state).copy())
|
||||||
self.actions[self.pos] = th.FloatTensor(np.array(action))
|
self.actions[self.pos] = th.FloatTensor(np.array(action).copy())
|
||||||
self.rewards[self.pos] = th.FloatTensor(np.array(reward))
|
self.rewards[self.pos] = th.FloatTensor(np.array(reward).copy())
|
||||||
self.dones[self.pos] = th.FloatTensor(np.array(done))
|
self.dones[self.pos] = th.FloatTensor(np.array(done).copy())
|
||||||
self.pos += 1
|
self.pos += 1
|
||||||
if self.pos == self.buffer_size:
|
if self.pos == self.buffer_size:
|
||||||
self.full = True
|
self.full = True
|
||||||
|
|
|
||||||
Loading…
Reference in a new issue