Skip to content
This repository was archived by the owner on Apr 3, 2026. It is now read-only.
Open
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -113,13 +113,13 @@ class CrazyflieEnvCfg(DirectRLEnvCfg):
unsafe_velocity_reward_scale = -1.0

# random pose range
platform_spawn_range_xy = 3.0
platform_spawn_range_xy = 3
platform_spawn_z = 0.0
drone_min_height = 0.4
drone_max_height = 0.6
drone_min_height = 0.3
drone_max_height = 0.8

# Alpha bot movement parameters
platform_max_linear_velocity = 0.5 # m/s
platform_max_linear_velocity = 0.25 # m/s
platform_max_angular_velocity = 1.0 # rad/s


Expand All @@ -139,7 +139,7 @@ def __init__(self, cfg: CrazyflieEnvCfg, render_mode: str | None = None, **kwarg
self._platform_joint_indices = None

# Platform target velocities (linear and angular)
self._platform_target_lin_vel = torch.rand(self.num_envs, 3, device=self.device) / 2.0
self._platform_target_lin_vel = torch.rand(self.num_envs, 3, device=self.device)

# Goal position (platform center)
self._desired_pos_w = torch.zeros(self.num_envs, 3, device=self.device)
Expand Down Expand Up @@ -204,8 +204,6 @@ def _pre_physics_step(self, actions: torch.Tensor):
self._drone_target_ang_vel_b[:, 2] = self._actions[:, 3] * self.cfg.max_angular_velocity_z

def _apply_action(self):
dt = self.sim.cfg.dt * self.cfg.decimation

drone_target_lin_vel_w = quat_apply(
self._robot.data.root_quat_w,
self._drone_target_lin_vel_b
Expand All @@ -220,15 +218,7 @@ def _apply_action(self):
torch.cat([drone_target_lin_vel_w, drone_target_ang_vel_w], dim=-1)
)

# Apply platform control
self._platform.set_joint_velocity_target(
self._platform_wheel_vel,
joint_ids=self._platform_joint_indices
)
# Change the platform position based on its velocity
new_platform_pos = self._platform.data.root_pos_w + self._platform_target_lin_vel * dt
new_platform_quat = self._platform.data.root_quat_w
self._platform.write_root_pose_to_sim(torch.cat([new_platform_pos, new_platform_quat], dim=-1))
self._move_platform()

def _get_observations(self) -> dict:
self._desired_pos_w = self._platform.data.root_pos_w.clone()
Expand All @@ -247,6 +237,47 @@ def _get_observations(self) -> dict:
observations = {"policy": obs}
return observations

def _move_platform(self):
dt = self.sim.cfg.dt * self.cfg.decimation

# Apply platform control
self._platform.set_joint_velocity_target(
self._platform_wheel_vel,
joint_ids=self._platform_joint_indices
)
# Change the platform position based on its velocity
new_platform_pos = self._platform.data.root_pos_w + self._platform_target_lin_vel * dt
new_platform_quat = self._platform.data.root_quat_w
self._platform.write_root_pose_to_sim(torch.cat([new_platform_pos, new_platform_quat], dim=-1))

def _spawn_platform(self, env_ids: torch.Tensor):
num_envs_to_reset = len(env_ids)

# Reset platform position
random_xy = torch.zeros(num_envs_to_reset, 2, device=self.device).uniform_(
-self.cfg.platform_spawn_range_xy,
self.cfg.platform_spawn_range_xy
)
random_z = torch.full((num_envs_to_reset, 1), self.cfg.platform_spawn_z, device=self.device)
random_pos = torch.cat([random_xy, random_z], dim=-1)
random_pos += self._terrain.env_origins[env_ids]
default_root_state_platform = self._platform.data.default_root_state[env_ids]
default_root_state_platform[:, :3] = random_pos
random_yaw = torch.rand(num_envs_to_reset, device=self.device) * 3.14159
default_root_state_platform[:, 3] = torch.cos(random_yaw) # w
default_root_state_platform[:, 6] = torch.sin(random_yaw) # z
self._platform.write_root_pose_to_sim(default_root_state_platform[:, :7], env_ids)
self._platform_wheel_vel[env_ids] = 0.0

# Set random linear velocity for the platform
random_lin_vel = torch.zeros(num_envs_to_reset, 2, device=self.device).uniform_(
-self.cfg.platform_max_linear_velocity,
self.cfg.platform_max_linear_velocity
)
self._platform_target_lin_vel[env_ids, 0] = random_lin_vel[:, 0]
self._platform_target_lin_vel[env_ids, 1] = random_lin_vel[:, 1]
self._platform_target_lin_vel[env_ids, 2] = 0.0

def _get_rewards(self) -> torch.Tensor:
lin_vel = torch.sum(torch.square(self._robot.data.root_lin_vel_b), dim=1)
ang_vel = torch.sum(torch.square(self._robot.data.root_ang_vel_b), dim=1)
Expand Down Expand Up @@ -304,32 +335,7 @@ def _reset_idx(self, env_ids: torch.Tensor | None):
extras["Metrics/final_distance_to_goal"] = final_distance_to_goal.item()
self.extras["log"].update(extras)

num_envs_to_reset = len(env_ids)

# Reset platform position
random_xy = torch.zeros(num_envs_to_reset, 2, device=self.device).uniform_(
-self.cfg.platform_spawn_range_xy,
self.cfg.platform_spawn_range_xy
)
random_z = torch.full((num_envs_to_reset, 1), self.cfg.platform_spawn_z, device=self.device)
random_pos = torch.cat([random_xy, random_z], dim=-1)
random_pos += self._terrain.env_origins[env_ids]
default_root_state_platform = self._platform.data.default_root_state[env_ids]
default_root_state_platform[:, :3] = random_pos
random_yaw = torch.rand(num_envs_to_reset, device=self.device) * 3.14159
default_root_state_platform[:, 3] = torch.cos(random_yaw) # w
default_root_state_platform[:, 6] = torch.sin(random_yaw) # z
self._platform.write_root_pose_to_sim(default_root_state_platform[:, :7], env_ids)
self._platform_wheel_vel[env_ids] = 0.0

# Set random linear velocity for the platform
random_lin_vel = torch.zeros(num_envs_to_reset, 2, device=self.device).uniform_(
-self.cfg.platform_max_linear_velocity,
self.cfg.platform_max_linear_velocity
)
self._platform_target_lin_vel[env_ids, 0] = random_lin_vel[:, 0]
self._platform_target_lin_vel[env_ids, 1] = random_lin_vel[:, 1]
self._platform_target_lin_vel[env_ids, 2] = 0.0
self._spawn_platform(env_ids)

self._robot.reset(env_ids)
super()._reset_idx(env_ids)
Expand Down