made a bit of progress

This commit is contained in:
2023-01-18 14:29:36 +01:00
parent f1982d1fe4
commit 8c36302edc
7 changed files with 27 additions and 14 deletions
+10 -3
View File
@@ -17,7 +17,7 @@ class BaseEconWrapper():
info=None
n_data_retrieved=0
def __init__(self, econ: base_env.BaseEnvironment):
def __init__(self, econ: base_env.BaseEnvironment,auto_reset):
self.env=econ
self.vote_lock=Lock()
@@ -26,6 +26,8 @@ class BaseEconWrapper():
self.action_edit_lock=Lock()
self.stop_edit_lock=Lock()
self.env_data_lock=Lock()
self.auto_reset=auto_reset
def register_vote(self):
"""Register reciever on base. Returns ID of Voter to pass on during blocking"""
@@ -79,7 +81,7 @@ class BaseEconWrapper():
#check for reset
self.vote_lock.acquire() # we might edit votes
if self.n_voters==self.n_votes_reset:
if self.n_voters<=self.n_votes_reset:
## perform reset
self.n_votes_reset=0
self._reset()
@@ -96,6 +98,7 @@ class BaseEconWrapper():
# release actions
# we are done
def stop_env(self):
"""Stops the wrapper"""
self.stop_edit_lock.acquire()
@@ -110,13 +113,17 @@ class BaseEconWrapper():
self.n_votes_reset=0
self.obs=self.env.reset() #Reset env
self.rew=None
self.done=None
self.done={'__all__': False}
self.info=None
self.env_data_lock.release() #Release lock
# Notify for reset
self.reset_notification.set()
for v in self.step_notifications:
v.clear() # unlock stepping
def force_reset(self):
"""Force a reset with no votes"""
self._reset()
#self.n_votes_reset=-self.n_voters
def _step(self):
"""Steping interaly"""
+2 -1
View File
@@ -21,6 +21,7 @@ class SB3EconConverter(VecEnv, gym.Env):
obs0=utils.package(obs[0],*self.packager)
obs0["flat"]
self.step_request_send=False
self.reset_request_send=False
self.auto_reset=auto_reset
self.observation_space=gym.spaces.Box(low=0,high=10,shape=(len(obs0["flat"]),),dtype=np.float32)
super().__init__(self.num_envs, self.observation_space, self.action_space)
@@ -73,7 +74,7 @@ class SB3EconConverter(VecEnv, gym.Env):
def reset(self) -> VecEnvObs:
obs=self.env.reset()
self.step_request_send=False
f_obs={}
self.curr_obs=obs
for k,v in obs.items():