Skip to content

Commit 623637f

Browse files
authored
[envpool] Fix explicit episodic-life reset for Atari and Vizdoom (#344)
## Summary - Problem: explicit `env.reset()` did not fully reset underlying state for episodic-life envs, so repeated resets could continue the same game/session instead of starting a fresh episode. - Scope: preserve the existing `force_reset` signal through the core async env path, make Atari and Vizdoom honor explicit resets under `episodic_life`, add regression coverage for both Python-facing APIs, and keep the branch green against current lint drift. - Outcome: explicit user resets now route to a real underlying reset for both env families instead of being treated like episodic-life continuation. Fixes #301 by restoring explicit reset semantics for episodic-life envs. ## Technical Details - Approach: thread the existing `force_reset` bit from `AsyncEnvPool` into `Env::PreProcess`, then consume it in env-specific reset logic. - Code pointers: - `envpool/core/env.h`: stores `force_reset_` on the env instance and threads it through `EnvStep`/`PreProcess`. - `envpool/core/async_envpool.h`: forwards `raw_action.force_reset` instead of collapsing explicit resets into generic reset handling. - `envpool/atari/atari_env.h`: treats explicit resets as a full `reset_game()` even when `episodic_life=True`. - `envpool/vizdoom/vizdoom_env.h`: treats explicit resets as a full `newEpisode()` even when `episodic_life=True`. - `envpool/atari/atari_env_test.cc`: adds a C++ regression test for repeated explicit resets. - `envpool/atari/atari_envpool_test.py`: adds a gymnasium regression test matching the user-facing Atari repro. - `envpool/vizdoom/vizdoom_test.py`: adds a Python regression test for repeated explicit resets under Vizdoom episodic life. - `envpool/core/xla_template.h`: fixes current `clang-tidy` / lint complaints encountered while validating this branch. - Notes: the Linux `make bazel-test` run on `dev` used `--distdir` with a locally downloaded `pretrain.tar.gz` and an Atari ROM override to avoid unrelated external download failures during validation. ## Test Plan ### Automated - `make bazel-test BAZELOPT="--distdir=/tmp/bazel-distdir-test --override_repository=atari_roms=/tmp/atari_roms_override_test"` on `dev`: passed (`30 / 30 tests pass`) before the Vizdoom follow-up. - issue #301 Atari repro script on `dev` against a patched wheel: passed (`info["lives"]` stayed at 5 across repeated `reset()`). - `ruff check envpool/vizdoom/vizdoom_test.py`: passed. - `ruff format --check envpool/vizdoom/vizdoom_test.py`: passed. - `clang-format --style=file -n --Werror envpool/vizdoom/vizdoom_env.h` on `dev`: passed. ### Known Verification Gap - Targeted `//envpool/vizdoom:vizdoom_test` execution on `dev` currently crashes in existing test setup with `v_video.cpp:1346 Assertion 'CleanWidth >= 320' failed`, including the pre-existing `test_timelimit` case. That prevented runtime confirmation of the new Vizdoom regression test on this box, but the failure reproduces independently of this patch.
1 parent a8c4620 commit 623637f

8 files changed

Lines changed: 122 additions & 19 deletions

File tree

envpool/atari/atari_env.h

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -151,7 +151,7 @@ class AtariEnv : public Env<AtariEnvSpec> {
151151
void Reset() override {
152152
int noop = dist_noop_(gen_) + 1 - static_cast<int>(fire_reset_);
153153
bool push_all = false;
154-
if (!episodic_life_ || env_->game_over() ||
154+
if (force_reset_ || !episodic_life_ || env_->game_over() ||
155155
elapsed_step_ >= max_episode_steps_) {
156156
env_->reset_game();
157157
elapsed_step_ = 0;

envpool/atari/atari_env_test.cc

Lines changed: 20 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -254,6 +254,26 @@ TEST(AtariEnvTest, EpisodicLife) {
254254
}
255255
}
256256

257+
TEST(AtariEnvTest, ExplicitResetWithEpisodicLife) {
258+
auto config = atari::AtariEnvSpec::kDefaultConfig;
259+
config["num_envs"_] = 1;
260+
config["batch_size"_] = 1;
261+
config["seed"_] = 42;
262+
config["episodic_life"_] = true;
263+
config["task"_] = "breakout";
264+
atari::AtariEnvSpec spec(config);
265+
atari::AtariEnvPool envpool(spec);
266+
TArray all_env_ids(Spec<int>({1}));
267+
all_env_ids[0] = 0;
268+
for (int i = 0; i < 20; ++i) {
269+
envpool.Reset(all_env_ids);
270+
AtariState state(envpool.Recv());
271+
auto lives = state["info:lives"_];
272+
int live = lives[0];
273+
EXPECT_EQ(live, 5) << "reset #" << i;
274+
}
275+
}
276+
257277
TEST(AtariEnvTest, ZeroDiscountOnLifeLoss) {
258278
std::srand(std::time(nullptr));
259279
int batch = 4;

envpool/atari/atari_envpool_test.py

Lines changed: 16 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -144,6 +144,22 @@ def test_reset_life(self) -> None:
144144
self.assertTrue(next_info["lives"][0] > 0)
145145
self.assertTrue(info["terminated"][0])
146146

147+
def test_explicit_reset_with_episodic_life_gymnasium(self) -> None:
148+
"""Issue 301."""
149+
env = make_gymnasium(
150+
"Breakout-v5",
151+
num_envs=1,
152+
seed=42,
153+
episodic_life=True,
154+
)
155+
for i in range(20):
156+
_, info = env.reset()
157+
self.assertEqual(
158+
int(info["lives"][0]),
159+
5,
160+
msg=f"reset #{i} returned lives={info['lives']}",
161+
)
162+
147163
def test_partial_step(self) -> None:
148164
num_envs = 5
149165
max_episode_steps = 10

envpool/core/async_envpool.h

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -124,7 +124,8 @@ class AsyncEnvPool : public EnvPool<typename Env::Spec> {
124124
int env_id = raw_action.env_id;
125125
int order = raw_action.order;
126126
bool reset = raw_action.force_reset || envs_[env_id]->IsDone();
127-
envs_[env_id]->EnvStep(state_buffer_queue_.get(), order, reset);
127+
envs_[env_id]->EnvStep(state_buffer_queue_.get(), order, reset,
128+
raw_action.force_reset);
128129
}
129130
});
130131
}

envpool/core/env.h

Lines changed: 6 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -65,6 +65,7 @@ class Env {
6565
EnvSpec spec_;
6666
int env_id_, seed_;
6767
std::mt19937 gen_;
68+
bool force_reset_{false};
6869

6970
private:
7071
StateBufferQueue* sbq_;
@@ -159,8 +160,8 @@ class Env {
159160
}
160161
}
161162

162-
void EnvStep(StateBufferQueue* sbq, int order, bool reset) {
163-
PreProcess(sbq, order, reset);
163+
void EnvStep(StateBufferQueue* sbq, int order, bool reset, bool force_reset) {
164+
PreProcess(sbq, order, reset, force_reset);
164165
if (reset) {
165166
Reset();
166167
} else {
@@ -178,9 +179,11 @@ class Env {
178179
virtual bool IsDone() { throw std::runtime_error("is_done not implemented"); }
179180

180181
protected:
181-
void PreProcess(StateBufferQueue* sbq, int order, bool reset) {
182+
void PreProcess(StateBufferQueue* sbq, int order, bool reset,
183+
bool force_reset) {
182184
sbq_ = sbq;
183185
order_ = order;
186+
force_reset_ = force_reset;
184187
if (reset) {
185188
current_step_ = 0;
186189
} else {

envpool/core/xla_template.h

Lines changed: 17 additions & 13 deletions
Original file line numberDiff line numberDiff line change
@@ -66,19 +66,19 @@ struct CustomCall {
6666

6767
static xla_ffi::Error ValidateArity(xla_ffi::RemainingArgs args,
6868
xla_ffi::RemainingRets rets) {
69-
constexpr std::size_t kExpectedArgs = std::tuple_size_v<In> + 1;
70-
constexpr std::size_t kExpectedRets = std::tuple_size_v<Out> + 1;
71-
if (args.size() != kExpectedArgs) {
69+
constexpr std::size_t expected_args = std::tuple_size_v<In> + 1;
70+
constexpr std::size_t expected_rets = std::tuple_size_v<Out> + 1;
71+
if (args.size() != expected_args) {
7272
return xla_ffi::Error::InvalidArgument(
73-
"Expected " + std::to_string(kExpectedArgs) + " buffers, got " +
73+
"Expected " + std::to_string(expected_args) + " buffers, got " +
7474
std::to_string(args.size()));
7575
}
76-
if (rets.size() != kExpectedRets) {
76+
if (rets.size() != expected_rets) {
7777
return xla_ffi::Error::InvalidArgument(
78-
"Expected " + std::to_string(kExpectedRets) + " results, got " +
78+
"Expected " + std::to_string(expected_rets) + " results, got " +
7979
std::to_string(rets.size()));
8080
}
81-
return xla_ffi::Error();
81+
return {};
8282
}
8383

8484
static xla_ffi::Error PopulateInBuffers(xla_ffi::RemainingArgs args,
@@ -90,7 +90,7 @@ struct CustomCall {
9090
}
9191
(*in_arr)[i] = (*buffer).untyped_data();
9292
}
93-
return xla_ffi::Error();
93+
return {};
9494
}
9595

9696
static xla_ffi::Error PopulateOutBuffers(xla_ffi::RemainingRets rets,
@@ -102,7 +102,7 @@ struct CustomCall {
102102
}
103103
(*out_arr)[i] = (*buffer)->untyped_data();
104104
}
105-
return xla_ffi::Error();
105+
return {};
106106
}
107107

108108
static xla_ffi::Error CpuExecute(xla_ffi::RemainingArgs args,
@@ -124,7 +124,7 @@ struct CustomCall {
124124
return err;
125125
}
126126
CC::Cpu(*obj, in_arr, out_arr);
127-
return xla_ffi::Error();
127+
return {};
128128
}
129129

130130
static xla_ffi::Error GpuExecute(cudaStream_t stream,
@@ -147,7 +147,7 @@ struct CustomCall {
147147
return err;
148148
}
149149
CC::Gpu(*obj, stream, in_arr, out_arr);
150-
return xla_ffi::Error();
150+
return {};
151151
}
152152

153153
static auto Specs(Class* obj) {
@@ -177,8 +177,12 @@ struct CustomCall {
177177
.RemainingArgs()
178178
.RemainingRets()
179179
.Attrs<xla_ffi::Dictionary>());
180-
return std::make_tuple(py::capsule(reinterpret_cast<void*>(cpu_handler)),
181-
py::capsule(reinterpret_cast<void*>(gpu_handler)));
180+
// NOLINTNEXTLINE(bugprone-casting-through-void)
181+
auto* cpu_handler_ptr = reinterpret_cast<void*>(cpu_handler);
182+
// NOLINTNEXTLINE(bugprone-casting-through-void)
183+
auto* gpu_handler_ptr = reinterpret_cast<void*>(gpu_handler);
184+
return std::make_tuple(py::capsule(cpu_handler_ptr),
185+
py::capsule(gpu_handler_ptr));
182186
}
183187

184188
static auto Xla(Class* obj) {

envpool/vizdoom/vizdoom_env.h

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -274,7 +274,8 @@ class VizdoomEnv : public Env<VizdoomEnvSpec> {
274274
bool IsDone() override { return done_; }
275275

276276
void Reset() override {
277-
if (dg_->isEpisodeFinished() || elapsed_step_ >= max_episode_steps_) {
277+
if (force_reset_ || dg_->isEpisodeFinished() ||
278+
elapsed_step_ >= max_episode_steps_) {
278279
elapsed_step_ = 0;
279280
if (episode_count_ > 0) { // NewEpisode at beginning may hang on MAEnv
280281
if (save_lmp_) {

envpool/vizdoom/vizdoom_test.py

Lines changed: 58 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -155,6 +155,64 @@ def test_obs_space(self) -> None:
155155
== 1 * 4
156156
)
157157

158+
def test_explicit_reset_with_episodic_life_gymnasium(self) -> None:
159+
env = make_gym(
160+
"D1Basic-v1",
161+
num_envs=1,
162+
seed=42,
163+
episodic_life=True,
164+
use_combined_action=True,
165+
)
166+
env.reset()
167+
tracked_keys = [
168+
"AMMO2",
169+
"HEALTH",
170+
"HITCOUNT",
171+
"KILLCOUNT",
172+
"SELECTED_WEAPON_AMMO",
173+
]
174+
175+
def scalar(info: dict, key: str) -> float:
176+
return float(np.asarray(info[key]).reshape(-1)[0])
177+
178+
action_id = None
179+
changed_key = None
180+
baseline_value = None
181+
changed_value = None
182+
183+
for candidate in range(env.action_space.n):
184+
_, baseline_info = env.reset()
185+
for _ in range(64):
186+
_, _, terminated, truncated, info = env.step(
187+
np.array([candidate], dtype=int)
188+
)
189+
for key in tracked_keys:
190+
current = scalar(info, key)
191+
baseline = scalar(baseline_info, key)
192+
if current != baseline:
193+
action_id = candidate
194+
changed_key = key
195+
baseline_value = baseline
196+
changed_value = current
197+
break
198+
if changed_key is not None or terminated[0] or truncated[0]:
199+
break
200+
if changed_key is not None:
201+
break
202+
203+
assert changed_key is not None
204+
assert baseline_value is not None
205+
assert changed_value is not None
206+
_, reset_info = env.reset()
207+
self.assertEqual(
208+
scalar(reset_info, changed_key),
209+
baseline_value,
210+
msg=(
211+
f"action={action_id}, key={changed_key}, "
212+
f"changed={changed_value}, baseline={baseline_value}"
213+
),
214+
)
215+
158216

159217
if __name__ == "__main__":
160218
absltest.main()

0 commit comments

Comments
 (0)