#1281·tianshou

使用 ShmemVectorEnv,缓冲观测值等于下一个观测值

作者: heitor57创建于 2025年10月28日更新于 2025年10月28日

def make_generic_env(): class GenericEnv(gym.Env): def init(self): super().init() self.num_features = 5 self.entity_id = 0 self.step_count = 0 self.observation_space = gym.spaces.Dict({ "id": gym.spaces.Box(low=0, high=100, shape=(1,), dtype=np.int32), "step": gym.spaces.Box(low=0, high=100, shape=(1,), dtype=np.int32), "features": gym.spaces.Box( low=0, high=9999, shape=(self.num_features,), dtype=np.int32 ), }) self.action_space = gym.spaces.Discrete(self.num_features) def _get_obs(self) -> Dict[str, Any]: return { "id": np.array([self.entity_id], dtype=np.int32), "step": np.array([self.step_count], dtype=np.int32), "features": np.arange(self.num_features, dtype=np.int32) + self.step_count, } def reset(self, *, seed=None, options=None): self.step_count = 0 obs = self._get_obs() return copy.deepcopy(obs), {} def step(self, action): self.step_count += 1 reward = float(action) terminated = self.step_count >= 3 truncated = False obs_next = self._get_obs() return copy.deepcopy(obs_next), reward, terminated, truncated, {} return GenericEnv()

def run_test(vector_env_type): print(f"Testing Vector Env Type: {vector_env_type.name}") num_envs = 2 env = vector_env_type([make_generic_env for _ in range(num_envs)]) initial_obs, _ = env.reset() policy = RandomActionPolicy(env.get_env_attr("action_space")[0]) buffer = VectorReplayBuffer(total_size=num_envs * 10, buffer_num=num_envs)

内容来源: thu-ml/tianshou