Skip to content

Commit 19b5eb6

Browse files
Early environment resets based on agents' respawn status. (#167)
* Added early termination parameter based on respawn status of all agents in an episode * pre-commit fix * fix test * Apply precommit. * Reduce variance in aggregate metrics by logging only if we have data for at least num_agents. --------- Co-authored-by: Daphne Cornelisse <cor.daphne@gmail.com>
1 parent a8bce58 commit 19b5eb6

7 files changed

Lines changed: 43 additions & 5 deletions

File tree

pufferlib/config/ocean/drive.ini

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -41,6 +41,7 @@ offroad_behavior = 0
4141
; Number of steps before reset
4242
scenario_length = 91
4343
resample_frequency = 910
44+
termination_mode = 1 # 0 - terminate at scenario_length, 1 - terminate after all agents have been reset
4445
map_dir = "resources/drive/binaries/training"
4546
num_maps = 10000
4647
; Determines which step of the trajectory to initialize the agents at upon reset

pufferlib/ocean/drive/binding.c

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -186,6 +186,7 @@ static int my_init(Env *env, PyObject *args, PyObject *kwargs) {
186186
env->reward_goal_post_respawn = conf.reward_goal_post_respawn;
187187
env->reward_ade = conf.reward_ade;
188188
env->scenario_length = conf.scenario_length;
189+
env->termination_mode = conf.termination_mode;
189190
env->collision_behavior = conf.collision_behavior;
190191
env->offroad_behavior = conf.offroad_behavior;
191192
env->max_controlled_agents = unpack(kwargs, "max_controlled_agents");

pufferlib/ocean/drive/drive.h

Lines changed: 13 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -324,6 +324,7 @@ struct Drive {
324324
GridMap *grid_map;
325325
int *neighbor_offsets;
326326
int scenario_length;
327+
int termination_mode;
327328
float reward_vehicle_collision;
328329
float reward_offroad_collision;
329330
float reward_ade;
@@ -2215,7 +2216,18 @@ void c_step(Drive *env) {
22152216
memset(env->rewards, 0, env->active_agent_count * sizeof(float));
22162217
memset(env->terminals, 0, env->active_agent_count * sizeof(unsigned char));
22172218
env->timestep++;
2218-
if (env->timestep == env->scenario_length) {
2219+
2220+
int originals_remaining = 0;
2221+
for (int i = 0; i < env->active_agent_count; i++) {
2222+
int agent_idx = env->active_agent_indices[i];
2223+
// Keep flag true if there is at least one agent that has not been respawned yet
2224+
if (env->entities[agent_idx].respawn_count == 0) {
2225+
originals_remaining = 1;
2226+
break;
2227+
}
2228+
}
2229+
2230+
if (env->timestep == env->scenario_length || (!originals_remaining && env->termination_mode == 1)) {
22192231
add_log(env);
22202232
c_reset(env);
22212233
return;

pufferlib/ocean/drive/drive.py

Lines changed: 4 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -28,6 +28,7 @@ def __init__(
2828
offroad_behavior=0,
2929
dt=0.1,
3030
scenario_length=None,
31+
termination_mode=None,
3132
resample_frequency=91,
3233
num_maps=100,
3334
num_agents=512,
@@ -57,6 +58,7 @@ def __init__(
5758
self.reward_ade = reward_ade
5859
self.human_agent_idx = human_agent_idx
5960
self.scenario_length = scenario_length
61+
self.termination_mode = termination_mode
6062
self.resample_frequency = resample_frequency
6163
self.dynamics_model = dynamics_model
6264

@@ -179,6 +181,7 @@ def __init__(
179181
offroad_behavior=self.offroad_behavior,
180182
dt=dt,
181183
scenario_length=(int(scenario_length) if scenario_length is not None else None),
184+
termination_mode=(int(self.termination_mode) if self.termination_mode is not None else 0),
182185
max_controlled_agents=self.max_controlled_agents,
183186
map_id=map_ids[i],
184187
max_agents=nxt - cur,
@@ -204,7 +207,7 @@ def step(self, actions):
204207
self.tick += 1
205208
info = []
206209
if self.tick % self.report_interval == 0:
207-
log = binding.vec_log(self.c_envs)
210+
log = binding.vec_log(self.c_envs, self.num_agents)
208211
if log:
209212
info.append(log)
210213
# print(log)

pufferlib/ocean/env_binding.h

Lines changed: 21 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -556,30 +556,48 @@ static int assign_to_dict(PyObject *dict, char *key, float value) {
556556
}
557557

558558
static PyObject *vec_log(PyObject *self, PyObject *args) {
559+
if (PyTuple_Size(args) != 2) {
560+
PyErr_SetString(PyExc_TypeError, "vec_log requires 2 arguments");
561+
return NULL;
562+
}
563+
559564
VecEnv *vec = unpack_vecenv(args);
560565
if (!vec) {
561566
return NULL;
562567
}
563568

564569
// Iterates over logs one float at a time. Will break
565570
// horribly if Log has non-float data.
571+
PyObject *num_agents_arg = PyTuple_GetItem(args, 1);
572+
float num_agents = (float)PyLong_AsLong(num_agents_arg);
573+
566574
Log aggregate = {0};
567575
int num_keys = sizeof(Log) / sizeof(float);
568576
for (int i = 0; i < vec->num_envs; i++) {
569577
Env *env = vec->envs[i];
570578
for (int j = 0; j < num_keys; j++) {
571579
((float *)&aggregate)[j] += ((float *)&env->log)[j];
572-
((float *)&env->log)[j] = 0.0f;
573580
}
574581
}
575582

576583
PyObject *dict = PyDict_New();
577-
if (aggregate.n == 0.0f) {
584+
585+
// Only log if we have at least num_agents worth of data
586+
if (aggregate.n < num_agents) {
578587
return dict;
579588
}
580589

581-
// Average
590+
// Got enough data. Reset logs and return metrics
591+
for (int i = 0; i < vec->num_envs; i++) {
592+
Env *env = vec->envs[i];
593+
for (int j = 0; j < num_keys; j++) {
594+
((float *)&env->log)[j] = 0.0f;
595+
}
596+
}
597+
582598
float n = aggregate.n;
599+
600+
// Average across agents
583601
for (int i = 0; i < num_keys; i++) {
584602
((float *)&aggregate)[i] /= n;
585603
}

pufferlib/ocean/env_config.h

Lines changed: 3 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -23,6 +23,7 @@ typedef struct {
2323
float dt;
2424
int goal_behavior;
2525
int scenario_length;
26+
int termination_mode;
2627
int init_steps;
2728
int init_mode;
2829
int control_mode;
@@ -78,6 +79,8 @@ static int handler(void *config, const char *section, const char *name, const ch
7879
env_config->dt = atof(value);
7980
} else if (MATCH("env", "scenario_length")) {
8081
env_config->scenario_length = atoi(value);
82+
} else if (MATCH("env", "termination_mode")) {
83+
env_config->termination_mode = atoi(value);
8184
} else if (MATCH("env", "init_steps")) {
8285
env_config->init_steps = atoi(value);
8386
} else if (MATCH("env", "init_mode")) {
-2.34 MB
Binary file not shown.

0 commit comments

Comments
 (0)