Skip to content

Commit fe5df69

Browse files
authored
Merging conditioning stuff first (#4)
* Seperates visualization code from the core drive.c Changes - Remove GIF generation code from drive.c - Improved load_weights to auto-detect file size - view flag - random maps if map-name not passed - policy-name flag to make videos for a particular policy. Saved in the policy directoryq * merging from main * fixed test_drive_render.py * added reward/entropy/discount conditioning * pre-commit fixes * some build fixes * removed output files * readd weights * fixed viz env * run pre-commit * pre-commit fix * fixed issues with conditioning and viz * changed respawn timstep to always be at obs[6] * merged main * fixed wrong jerk/classic obs values * was freeing twice
1 parent 269e7a7 commit fe5df69

12 files changed

Lines changed: 514 additions & 60 deletions

File tree

pufferlib/config/ocean/drive.ini

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -42,6 +42,7 @@ offroad_behavior = 0
4242
scenario_length = 91
4343
resample_frequency = 910
4444
num_maps = 1
45+
condition_type = "none" # Options: "none", "reward", "entropy", "discount", "all"
4546
; Determines which step of the trajectory to initialize the agents at upon reset
4647
init_steps = 0
4748
; Options: "control_vehicles", "control_agents", "control_tracks_to_predict"

pufferlib/extensions/cuda/pufferlib.cu

Lines changed: 12 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -21,7 +21,7 @@ __host__ __device__ void puff_advantage_row_cuda(float* values, float* rewards,
2121

2222
void vtrace_check_cuda(torch::Tensor values, torch::Tensor rewards,
2323
torch::Tensor dones, torch::Tensor importance, torch::Tensor advantages,
24-
int num_steps, int horizon) {
24+
torch::Tensor gammas, int num_steps, int horizon) {
2525

2626
// Validate input tensors
2727
torch::Device device = values.device();
@@ -35,27 +35,33 @@ void vtrace_check_cuda(torch::Tensor values, torch::Tensor rewards,
3535
t.contiguous();
3636
}
3737
}
38+
// Validate gammas tensor
39+
TORCH_CHECK(gammas.dim() == 1, "Gammas must be 1D");
40+
TORCH_CHECK(gammas.size(0) == num_steps, "Gammas size must match num_steps");
41+
TORCH_CHECK(gammas.dtype() == torch::kFloat32, "Gammas must be float32");
42+
TORCH_CHECK(gammas.is_cuda(), "Gammas must be on GPU");
43+
TORCH_CHECK(gammas.is_contiguous(), "Gammas must be contiguous");
3844
}
3945

4046
// [num_steps, horizon]
4147
__global__ void puff_advantage_kernel(float* values, float* rewards,
42-
float* dones, float* importance, float* advantages, float gamma,
48+
float* dones, float* importance, float* advantages, float* gammas,
4349
float lambda, float rho_clip, float c_clip, int num_steps, int horizon) {
4450
int row = blockIdx.x*blockDim.x + threadIdx.x;
4551
if (row >= num_steps) {
4652
return;
4753
}
4854
int offset = row*horizon;
4955
puff_advantage_row_cuda(values + offset, rewards + offset, dones + offset,
50-
importance + offset, advantages + offset, gamma, lambda, rho_clip, c_clip, horizon);
56+
importance + offset, advantages + offset, gammas[row], lambda, rho_clip, c_clip, horizon);
5157
}
5258

5359
void compute_puff_advantage_cuda(torch::Tensor values, torch::Tensor rewards,
5460
torch::Tensor dones, torch::Tensor importance, torch::Tensor advantages,
55-
double gamma, double lambda, double rho_clip, double c_clip) {
61+
torch::Tensor gammas, double lambda, double rho_clip, double c_clip) {
5662
int num_steps = values.size(0);
5763
int horizon = values.size(1);
58-
vtrace_check_cuda(values, rewards, dones, importance, advantages, num_steps, horizon);
64+
vtrace_check_cuda(values, rewards, dones, importance, advantages, gammas, num_steps, horizon);
5965
TORCH_CHECK(values.is_cuda(), "All tensors must be on GPU");
6066

6167
int threads_per_block = 256;
@@ -67,7 +73,7 @@ void compute_puff_advantage_cuda(torch::Tensor values, torch::Tensor rewards,
6773
dones.data_ptr<float>(),
6874
importance.data_ptr<float>(),
6975
advantages.data_ptr<float>(),
70-
gamma,
76+
gammas.data_ptr<float>(),
7177
lambda,
7278
rho_clip,
7379
c_clip,

pufferlib/extensions/pufferlib.cpp

Lines changed: 14 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -42,7 +42,7 @@ void puff_advantage_row(float* values, float* rewards, float* dones,
4242

4343
void vtrace_check(torch::Tensor values, torch::Tensor rewards,
4444
torch::Tensor dones, torch::Tensor importance, torch::Tensor advantages,
45-
int num_steps, int horizon) {
45+
torch::Tensor gammas, int num_steps, int horizon) {
4646

4747
// Validate input tensors
4848
torch::Device device = values.device();
@@ -56,36 +56,42 @@ void vtrace_check(torch::Tensor values, torch::Tensor rewards,
5656
t.contiguous();
5757
}
5858
}
59+
// Validate gammas tensor
60+
TORCH_CHECK(gammas.dim() == 1, "Gammas must be 1D");
61+
TORCH_CHECK(gammas.size(0) == num_steps, "Gammas size must match num_steps");
62+
TORCH_CHECK(gammas.dtype() == torch::kFloat32, "Gammas must be float32");
63+
TORCH_CHECK(gammas.is_contiguous(), "Gammas must be contiguous");
5964
}
6065

6166

6267
// [num_steps, horizon]
6368
void puff_advantage(float* values, float* rewards, float* dones, float* importance,
64-
float* advantages, float gamma, float lambda, float rho_clip, float c_clip,
69+
float* advantages, float* gammas, float lambda, float rho_clip, float c_clip,
6570
int num_steps, const int horizon){
66-
for (int offset = 0; offset < num_steps*horizon; offset+=horizon) {
71+
for (int row = 0; row < num_steps; row++) {
72+
int offset = row * horizon;
6773
puff_advantage_row(values + offset, rewards + offset,
6874
dones + offset, importance + offset, advantages + offset,
69-
gamma, lambda, rho_clip, c_clip, horizon
75+
gammas[row], lambda, rho_clip, c_clip, horizon
7076
);
7177
}
7278
}
7379

7480

7581
void compute_puff_advantage_cpu(torch::Tensor values, torch::Tensor rewards,
7682
torch::Tensor dones, torch::Tensor importance, torch::Tensor advantages,
77-
double gamma, double lambda, double rho_clip, double c_clip) {
83+
torch::Tensor gammas, double lambda, double rho_clip, double c_clip) {
7884
int num_steps = values.size(0);
7985
int horizon = values.size(1);
80-
vtrace_check(values, rewards, dones, importance, advantages, num_steps, horizon);
86+
vtrace_check(values, rewards, dones, importance, advantages, gammas, num_steps, horizon);
8187
puff_advantage(values.data_ptr<float>(), rewards.data_ptr<float>(),
8288
dones.data_ptr<float>(), importance.data_ptr<float>(), advantages.data_ptr<float>(),
83-
gamma, lambda, rho_clip, c_clip, num_steps, horizon
89+
gammas.data_ptr<float>(), lambda, rho_clip, c_clip, num_steps, horizon
8490
);
8591
}
8692

8793
TORCH_LIBRARY(pufferlib, m) {
88-
m.def("compute_puff_advantage(Tensor(a!) values, Tensor(b!) rewards, Tensor(c!) dones, Tensor(d!) importance, Tensor(e!) advantages, float gamma, float lambda, float rho_clip, float c_clip) -> ()");
94+
m.def("compute_puff_advantage(Tensor(a!) values, Tensor(b!) rewards, Tensor(c!) dones, Tensor(d!) importance, Tensor(e!) advantages, Tensor gammas, float lambda, float rho_clip, float c_clip) -> ()");
8995
}
9096

9197
TORCH_LIBRARY_IMPL(pufferlib, CPU, m) {

pufferlib/ocean/drive/binding.c

Lines changed: 22 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -73,6 +73,7 @@ static PyObject* my_shared(PyObject* self, PyObject* args, PyObject* kwargs) {
7373
int init_mode = unpack(kwargs, "init_mode");
7474
int control_mode = unpack(kwargs, "control_mode");
7575
int init_steps = unpack(kwargs, "init_steps");
76+
int max_controlled_agents = unpack(kwargs, "max_controlled_agents");
7677
clock_gettime(CLOCK_REALTIME, &ts);
7778
srand(ts.tv_nsec);
7879
int total_agent_count = 0;
@@ -89,6 +90,7 @@ static PyObject* my_shared(PyObject* self, PyObject* args, PyObject* kwargs) {
8990
env->init_mode = init_mode;
9091
env->control_mode = control_mode;
9192
env->init_steps = init_steps;
93+
env->max_controlled_agents = max_controlled_agents;
9294
sprintf(map_file, "resources/drive/binaries/map_%03d.bin", map_id);
9395
env->entities = load_map_binary(map_file, env);
9496
set_active_agents(env);
@@ -176,6 +178,10 @@ static int my_init(Env* env, PyObject* args, PyObject* kwargs) {
176178
}
177179
env->action_type = conf.action_type;
178180
env->dynamics_model = conf.dynamics_model;
181+
if (PyDict_GetItemString(kwargs, "dynamics_model")) {
182+
char* dynamics_str = unpack_str(kwargs, "dynamics_model");
183+
env->dynamics_model = (strcmp(dynamics_str, "jerk") == 0) ? JERK : CLASSIC;
184+
}
179185
env->reward_vehicle_collision = conf.reward_vehicle_collision;
180186
env->reward_offroad_collision = conf.reward_offroad_collision;
181187
env->reward_goal = conf.reward_goal;
@@ -188,6 +194,22 @@ static int my_init(Env* env, PyObject* args, PyObject* kwargs) {
188194
env->offroad_behavior = conf.offroad_behavior;
189195
env->max_controlled_agents = unpack(kwargs, "max_controlled_agents");
190196
env->dt = conf.dt;
197+
198+
// Conditioning parameters
199+
env->use_rc = (bool)unpack(kwargs, "use_rc");
200+
env->use_ec = (bool)unpack(kwargs, "use_ec");
201+
env->use_dc = (bool)unpack(kwargs, "use_dc");
202+
env->collision_weight_lb = (float)unpack(kwargs, "collision_weight_lb");
203+
env->collision_weight_ub = (float)unpack(kwargs, "collision_weight_ub");
204+
env->offroad_weight_lb = (float)unpack(kwargs, "offroad_weight_lb");
205+
env->offroad_weight_ub = (float)unpack(kwargs, "offroad_weight_ub");
206+
env->goal_weight_lb = (float)unpack(kwargs, "goal_weight_lb");
207+
env->goal_weight_ub = (float)unpack(kwargs, "goal_weight_ub");
208+
env->entropy_weight_lb = (float)unpack(kwargs, "entropy_weight_lb");
209+
env->entropy_weight_ub = (float)unpack(kwargs, "entropy_weight_ub");
210+
env->discount_weight_lb = (float)unpack(kwargs, "discount_weight_lb");
211+
env->discount_weight_ub = (float)unpack(kwargs, "discount_weight_ub");
212+
191213
env->init_mode = (int)unpack(kwargs, "init_mode");
192214
env->control_mode = (int)unpack(kwargs, "control_mode");
193215
int map_id = unpack(kwargs, "map_id");

pufferlib/ocean/drive/drive.c

Lines changed: 3 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -19,7 +19,7 @@ void test_drivenet() {
1919

2020
//Weights* weights = load_weights("resources/drive/puffer_drive_weights.bin");
2121
Weights* weights = load_weights("puffer_drive_weights.bin");
22-
DriveNet* net = init_drivenet(weights, num_agents);
22+
DriveNet* net = init_drivenet(weights, num_agents, CLASSIC, false, false, false);
2323

2424
forward(net, observations, actions);
2525
for (int i = 0; i < num_agents*num_actions; i++) {
@@ -50,7 +50,7 @@ void demo() {
5050
.reward_ade = conf.reward_ade,
5151
.goal_radius = conf.goal_radius,
5252
.dt = conf.dt,
53-
.map_name = "resources/drive/binaries/map_000.bin",
53+
.map_name = "resources/drive/binaries/map_000.bin",
5454
.init_steps = conf.init_steps,
5555
.collision_behavior = conf.collision_behavior,
5656
.offroad_behavior = conf.offroad_behavior,
@@ -59,7 +59,7 @@ void demo() {
5959
c_reset(&env);
6060
c_render(&env);
6161
Weights* weights = load_weights("resources/drive/puffer_drive_weights.bin");
62-
DriveNet* net = init_drivenet(weights, env.active_agent_count, env.dynamics_model);
62+
DriveNet* net = init_drivenet(weights, env.active_agent_count, env.dynamics_model, false, false, false);
6363
//Client* client = make_client(&env);
6464
int accel_delta = 2;
6565
int steer_delta = 4;

0 commit comments

Comments
 (0)